From 021f523431ec5697baa9104d68e90926b97c9f71 Mon Sep 17 00:00:00 2001 From: KKRainbow <443152178@qq.com> Date: Sun, 26 Jul 2026 15:41:55 +0800 Subject: [PATCH] refactor(core): separate portable core from native runtime (#2451) Create easytier-core as the portable owner of configuration, connectivity, tunnels, peer and routing state, gateways, management, the data plane, and instance lifecycle. Keep operating-system integration, native protocol engines, process startup, and presentation in easytier behind explicit Host capability adapters. Create easytier-proto to own schemas, generated RPC types, descriptors, and feature-scoped protocol slices. Remove runtime protobuf reflection from core while preserving unknown route-peer fields across forwarding. Normalize instance construction through CoreInstance, CoreHostAdapters, CoreProcessRuntime, and InstanceManager. Make the runtime config store the only authoritative mutable configuration after startup. Move the portable TCP/UDP data plane into core and extract a generic OperationBroker for completion, cancellation, disposal, and capacity accounting. Expose the session-based FFI v2 completion API and keep the WASI guest ABI, wire schemas, and adapters with core. Migrate CLI, GUI, web, FFI, Android JNI, OHOS, uptime, and mobile consumers to the shared manager and core state. Add explicit user/web config ownership and revision-aware web reconciliation. Preserve configuration, wire, and management behavior while fixing regressions discovered by the full platform and integration matrix: - inherit advertised relay capabilities in foreign networks; - refresh OSPF peer state immediately after runtime config changes; - restore CLI GlobalCtx event output without forcing GUI logging; - retain legacy encryption names and standalone RPC tunnel metadata; - restore ICMP host composition and fragmented UDP handling; - use portable 64-bit atomics on 32-bit MIPS targets; and - retain discarded operations until late cancellation completes. Validate the refactor across 45 GitHub checks, including Linux, macOS, Windows, FreeBSD, web, GUI, Android, OHOS, feature profiles, and three-node and subnet-proxy integration tests. BREAKING CHANGE: internal Rust module paths are not preserved. Legacy native data-plane APIs are replaced by the session-based FFI v2 API. The dedicated Android data-plane wrapper is removed. --- .github/workflows/core.yml | 2 +- .github/workflows/test.yml | 6 +- CONTEXT.md | 23 + Cargo.lock | 329 +- Cargo.toml | 2 + docs/core-architecture.md | 522 ++ docs/data-plane-runtime-plan.md | 1133 ++++ .../easytier-android-jni/Cargo.toml | 2 +- .../easytier-android-jni/exports.map | 1 - .../com/easytier/jni/EasyTierDataPlaneJNI.kt | 451 -- .../src/data_plane_api.rs | 673 --- .../easytier-android-jni/src/lib.rs | 528 +- easytier-contrib/easytier-ffi/Cargo.toml | 12 +- .../easytier-ffi/DATA_PLANE_ABI.md | 100 + .../examples/example_data_plane_async.c | 429 -- .../easytier-ffi/examples/go/README.md | 138 - .../easytier-ffi/examples/go/easytier.go | 593 -- .../examples/go/easytier_async.go | 1018 ---- .../examples/go/easytier_async_test.go | 360 -- .../easytier-ffi/examples/go/easytier_test.go | 140 - .../easytier-ffi/examples/go/go.mod | 5 - .../easytier-ffi/src/config_server.rs | 191 +- .../easytier-ffi/src/data_plane.rs | 928 ---- .../easytier-ffi/src/data_plane/abi.rs | 662 +++ .../easytier-ffi/src/data_plane/mod.rs | 16 + .../easytier-ffi/src/data_plane/session.rs | 636 +++ .../easytier-ffi/src/data_plane_async.rs | 1162 ---- .../easytier-ffi/src/instance_api.rs | 163 +- easytier-contrib/easytier-ffi/src/json_rpc.rs | 37 +- easytier-contrib/easytier-ffi/src/lib.rs | 665 +-- easytier-contrib/easytier-ffi/src/state.rs | 94 +- easytier-contrib/easytier-ffi/src/tests.rs | 362 +- easytier-contrib/easytier-ffi/src/types.rs | 20 + easytier-contrib/easytier-ohrs/Cargo.lock | 415 +- .../src/config_repo/import_export.rs | 1 + .../src/config_repo/validation.rs | 1 + .../easytier-ohrs/src/exports/runtime_api.rs | 8 +- .../src/kernel_bridge/socket_server.rs | 40 +- easytier-contrib/easytier-ohrs/src/lib.rs | 33 +- .../easytier-uptime/src/health_checker.rs | 93 +- easytier-core/Cargo.toml | 137 + easytier-core/src/config/api.rs | 166 + easytier-core/src/config/api_input.rs | 652 +++ easytier-core/src/config/encryption.rs | 73 + easytier-core/src/config/gateway.rs | 138 + easytier-core/src/config/mod.rs | 804 +++ easytier-core/src/config/peers.rs | 315 ++ easytier-core/src/config/runtime.rs | 253 + easytier-core/src/config/toml.rs | 1632 ++++++ easytier-core/src/config/toml/snapshot.rs | 15 + easytier-core/src/connectivity/composite.rs | 474 ++ .../src/connectivity/connector_host.rs | 708 +++ easytier-core/src/connectivity/direct/mod.rs | 1302 +++++ easytier-core/src/connectivity/direct/udp.rs | 21 + .../src/connectivity/hole_punch/mod.rs | 33 + .../connectivity/hole_punch/peer_adapters.rs | 188 + .../src/connectivity/hole_punch/policy.rs | 212 + .../connectivity/hole_punch/port_mapping.rs | 439 ++ .../src/connectivity/hole_punch/tcp.rs | 1104 ++++ .../connectivity/hole_punch/udp/binding.rs | 802 +++ .../src/connectivity/hole_punch/udp/client.rs | 995 ++++ .../src/connectivity/hole_punch/udp/common.rs | 249 + .../connectivity/hole_punch/udp/connector.rs | 527 ++ .../src/connectivity/hole_punch/udp/mod.rs | 39 + .../hole_punch/udp/punch_listener.rs | 195 + .../src/connectivity/hole_punch/udp/rpc.rs | 828 +++ .../connectivity/hole_punch/udp/runtime.rs | 509 ++ .../src/connectivity/hole_punch/udp/server.rs | 1386 +++++ .../hole_punch/udp/socket_array.rs | 376 ++ .../src/connectivity/hole_punch/udp/task.rs | 217 + .../src/connectivity/manual/discovery.rs | 73 + .../manual/discovery/implementation.rs | 924 +++ easytier-core/src/connectivity/manual/mod.rs | 1336 +++++ easytier-core/src/connectivity/mod.rs | 37 + .../src/connectivity/protocol/mod.rs | 811 +++ .../src/connectivity/protocol/raw.rs | 867 +++ easytier-core/src/connectivity/stun/client.rs | 1033 ++++ .../src/connectivity/stun/collector.rs | 768 +++ easytier-core/src/connectivity/stun/mod.rs | 9 + .../src/connectivity/stun/responder.rs | 247 + .../src/connectivity/transport/mod.rs | 70 + .../src/connectivity/transport/tcp.rs | 53 + .../src/connectivity/transport/udp.rs | 102 + easytier-core/src/events.rs | 104 + easytier-core/src/foundation/mod.rs | 11 + .../src/foundation/operation_broker.rs | 471 ++ .../src/foundation/stats.rs | 479 +- easytier-core/src/foundation/task.rs | 310 ++ easytier-core/src/foundation/time.rs | 14 + .../src/foundation}/token_bucket.rs | 164 +- .../src/gateway/dataplane/deadline.rs | 52 + easytier-core/src/gateway/dataplane/error.rs | 117 + easytier-core/src/gateway/dataplane/flow.rs | 415 ++ easytier-core/src/gateway/dataplane/mod.rs | 976 ++++ .../src/gateway/dataplane/operation.rs | 145 + easytier-core/src/gateway/dataplane/packet.rs | 377 ++ .../src/gateway/dataplane/resource.rs | 165 + easytier-core/src/gateway/dataplane/route.rs | 166 + .../src/gateway/dataplane/session.rs | 1540 +++++ easytier-core/src/gateway/dataplane/stack.rs | 130 + easytier-core/src/gateway/dataplane/tcp.rs | 160 + easytier-core/src/gateway/dataplane/tests.rs | 839 +++ easytier-core/src/gateway/dataplane/udp.rs | 80 + easytier-core/src/gateway/dhcp.rs | 564 ++ easytier-core/src/gateway/magic_dns.rs | 17 + easytier-core/src/gateway/magic_dns/packet.rs | 468 ++ .../src/gateway/magic_dns/records.rs | 429 ++ easytier-core/src/gateway/mod.rs | 32 + easytier-core/src/gateway/port_forward.rs | 453 ++ .../src/gateway/proxy/cidr_monitor.rs | 276 + easytier-core/src/gateway/proxy/cidr_table.rs | 164 + easytier-core/src/gateway/proxy/icmp_host.rs | 33 + .../src/gateway/proxy/icmp_proxy_engine.rs | 516 ++ .../src/gateway/proxy/icmp_proxy_service.rs | 337 ++ .../src/gateway/proxy/ip_reassembler.rs | 321 ++ easytier-core/src/gateway/proxy/mod.rs | 34 + easytier-core/src/gateway/proxy/proxy_acl.rs | 49 + easytier-core/src/gateway/proxy/service.rs | 526 ++ .../src/gateway/proxy/tcp_proxy_engine.rs | 603 ++ .../src/gateway/proxy/tcp_proxy_service.rs | 631 +++ .../src/gateway/proxy/tcp_socket_connector.rs | 62 + easytier-core/src/gateway/proxy/traits.rs | 98 + .../src/gateway/proxy/udp_proxy_engine.rs | 507 ++ .../src/gateway/proxy/udp_proxy_service.rs | 217 + .../src/gateway/proxy/udp_socket_runtime.rs | 647 +++ .../src/gateway/proxy/wrapped_tcp_proxy.rs | 450 ++ .../src/gateway/proxy/wrapped_transport.rs | 1142 ++++ .../proxy/wrapped_transport/connect_api.rs | 56 + .../proxy/wrapped_transport/engine_api.rs | 56 + .../proxy/wrapped_transport/packet_api.rs | 43 + .../proxy/wrapped_transport/packet_plane.rs | 531 ++ .../packet_plane_disabled.rs | 97 + .../proxy/wrapped_transport_destination.rs | 375 ++ easytier-core/src/gateway/smoltcp/mod.rs | 7 + easytier-core/src/gateway/smoltcp/stack.rs | 158 + .../smoltcp}/tokio_smoltcp/channel_device.rs | 0 .../gateway/smoltcp}/tokio_smoltcp/device.rs | 0 .../src/gateway/smoltcp}/tokio_smoltcp/mod.rs | 56 +- .../gateway/smoltcp}/tokio_smoltcp/reactor.rs | 11 +- .../gateway/smoltcp}/tokio_smoltcp/socket.rs | 78 +- .../tokio_smoltcp/socket_allocator.rs | 0 easytier-core/src/gateway/socks5/adapter.rs | 215 + .../src/gateway/socks5/codec.rs | 184 +- .../src/gateway/socks5/codec}/target_addr.rs | 98 +- easytier-core/src/gateway/socks5/host.rs | 366 ++ easytier-core/src/gateway/socks5/mod.rs | 13 + .../src/gateway/socks5}/server.rs | 496 +- easytier-core/src/gateway/udp_broadcast.rs | 740 +++ easytier-core/src/gateway/vpn_portal.rs | 713 +++ easytier-core/src/host/dns.rs | 425 ++ easytier-core/src/host/environment.rs | 231 + easytier-core/src/host/mod.rs | 15 + easytier-core/src/host/packet.rs | 281 + easytier-core/src/host/socket/factory.rs | 503 ++ easytier-core/src/host/socket/listener.rs | 462 ++ easytier-core/src/host/socket/mod.rs | 816 +++ easytier-core/src/host/socket/udp.rs | 522 ++ easytier-core/src/host/testkit.rs | 206 + .../src/instance/build_capabilities.rs | 79 + easytier-core/src/instance/config.rs | 386 ++ .../src/instance/data_plane_extension.rs | 46 + easytier-core/src/instance/lifecycle.rs | 249 + easytier-core/src/instance/management.rs | 211 + .../src/instance/management_extension.rs | 22 + .../src/instance/management_state.rs | 34 + easytier-core/src/instance/manager.rs | 664 +++ easytier-core/src/instance/mod.rs | 897 +++ easytier-core/src/instance/packet_io.rs | 190 + easytier-core/src/instance/packet_plane.rs | 143 + .../src/instance/packet_proxy_extension.rs | 41 + .../src/instance/public_ipv6_extension.rs | 10 + easytier-core/src/instance/test_utils.rs | 93 + easytier-core/src/instance/tests.rs | 2364 ++++++++ .../src/instance/vpn_portal_extension.rs | 13 + easytier-core/src/lib.rs | 21 + easytier-core/src/listener/mod.rs | 1016 ++++ easytier-core/src/listener/plan.rs | 475 ++ easytier-core/src/listener/transport.rs | 1581 ++++++ easytier-core/src/management/full/compiled.rs | 45 + .../src/management/full/config_patch.rs | 355 ++ .../src/management/full/instance_info.rs | 79 + .../src/management/full/logger_rpc.rs | 126 + easytier-core/src/management/full/mod.rs | 110 + .../src/management/full/packet_proxy.rs | 156 + .../src/management/full/process_rpc.rs | 746 +++ .../src/management/full}/remote_client.rs | 39 +- .../src/management/full/web_client.rs | 403 ++ .../src/management/instance_rpc/full.rs | 406 ++ .../src/management/instance_rpc/mod.rs | 327 ++ .../management/instance_rpc/packet_proxy.rs | 135 + .../src/management/instance_rpc/projection.rs | 39 + easytier-core/src/management/mod.rs | 52 + .../src/management/rpc_server_hook.rs | 59 + easytier-core/src/management/selector.rs | 227 + easytier-core/src/management/server.rs | 176 + .../src/packet}/compressor.rs | 72 +- easytier-core/src/packet/compressor/zstd.rs | 76 + easytier-core/src/packet/hole_punch.rs | 79 + .../src/packet/mod.rs | 140 +- .../src/packet/stun.rs | 29 +- .../src/peers/acl/filter.rs | 338 +- easytier-core/src/peers/acl/mod.rs | 7 + .../src/peers/acl/processor.rs | 258 +- easytier-core/src/peers/admission.rs | 131 + easytier-core/src/peers/conn/mod.rs | 8 + easytier-core/src/peers/conn/peer.rs | 290 + .../src/peers/conn}/peer_conn.rs | 1202 +--- .../src/peers/conn}/peer_conn_ping.rs | 41 +- .../src/peers/conn}/peer_map.rs | 116 +- .../src/peers/conn}/peer_session.rs | 71 +- easytier-core/src/peers/context.rs | 1775 ++++++ easytier-core/src/peers/credential_manager.rs | 497 ++ easytier-core/src/peers/error.rs | 31 + .../src/peers/foreign_network/client.rs | 35 +- .../src/peers/foreign_network/mod.rs | 1693 ++++++ easytier-core/src/peers/mod.rs | 60 + .../src/peers}/peer_center/instance.rs | 204 +- easytier-core/src/peers/peer_center/mod.rs | 21 + .../src/peers}/peer_center/server.rs | 8 +- easytier-core/src/peers/peer_manager.rs | 3933 +++++++++++++ easytier-core/src/peers/peer_rpc.rs | 113 + easytier-core/src/peers/public_ipv6/mod.rs | 738 +++ .../src/peers/public_ipv6/provider.rs | 601 ++ .../src/peers/public_ipv6/service.rs | 313 +- .../src/peers/relay_peer_map.rs | 195 +- .../src/peers/route}/graph_algo.rs | 0 .../src/peers/route/mod.rs | 29 +- .../src/peers/route}/peer_ospf_route.rs | 4933 ++++------------- .../src/peers/route/route_peer_wire.rs | 399 ++ easytier-core/src/peers/test_support.rs | 100 + easytier-core/src/peers/tests.rs | 113 + .../src/peers/traffic_metrics.rs | 114 +- easytier-core/src/peers/util.rs | 10 + easytier-core/src/peers/whitelist.rs | 18 + easytier-core/src/process_runtime.rs | 415 ++ .../src/rpc}/bidirect.rs | 89 +- .../src/rpc}/client.rs | 270 +- .../rpc_impl => easytier-core/src/rpc}/mod.rs | 2 +- .../src/rpc}/packet.rs | 169 +- .../src/rpc}/server.rs | 202 +- .../src/rpc}/service_registry.rs | 0 easytier-core/src/rpc/standalone.rs | 733 +++ easytier-core/src/socket/mod.rs | 104 + easytier-core/src/socket/ring.rs | 299 + easytier-core/src/socket/tcp.rs | 683 +++ easytier-core/src/socket/udp/layer.rs | 1144 ++++ easytier-core/src/socket/udp/listener.rs | 205 + easytier-core/src/socket/udp/mod.rs | 35 + easytier-core/src/socket/udp/packet.rs | 407 ++ easytier-core/src/socket/udp/session.rs | 785 +++ easytier-core/src/socket/udp/tests.rs | 2091 +++++++ .../src/socket/udp/virtual_socket.rs | 265 + .../src/tunnel}/encrypt/aes_gcm.rs | 6 +- easytier-core/src/tunnel/encrypt/chacha20.rs | 150 + easytier-core/src/tunnel/encrypt/mod.rs | 272 + .../src/tunnel}/encrypt/xor.rs | 9 +- .../src/tunnel/filter.rs | 96 +- easytier-core/src/tunnel/framed.rs | 375 ++ easytier-core/src/tunnel/mod.rs | 98 + easytier-core/src/tunnel/mpsc.rs | 139 + easytier-core/src/tunnel/ring.rs | 591 ++ .../src/tunnel}/secure_datagram.rs | 52 +- .../src/tunnel/stats.rs | 1 - easytier-core/src/tunnel/tcp.rs | 198 + easytier-core/src/tunnel/udp.rs | 422 ++ .../src/tunnel/web_security.rs | 32 +- easytier-core/src/tunnel/wrapper.rs | 52 + easytier-core/src/wasi/abi.rs | 85 + easytier-core/src/wasi/adapter/dns.rs | 121 + easytier-core/src/wasi/adapter/environment.rs | 76 + easytier-core/src/wasi/adapter/mod.rs | 6 + easytier-core/src/wasi/adapter/packet.rs | 59 + .../src/wasi/adapter/socket/backend.rs | 266 + easytier-core/src/wasi/adapter/socket/mod.rs | 124 + easytier-core/src/wasi/adapter/socket/udp.rs | 179 + easytier-core/src/wasi/imports.rs | 139 + easytier-core/src/wasi/mod.rs | 24 + easytier-core/src/wasi/runtime.rs | 778 +++ .../src/wasi/runtime/abi/data_plane.rs | 737 +++ easytier-core/src/wasi/runtime_driver.rs | 151 + easytier-core/src/wasi/schema.rs | 34 + easytier-core/src/wasi/time.rs | 318 ++ easytier-core/src/wasi/wire/common.rs | 53 + easytier-core/src/wasi/wire/data_plane.rs | 161 + easytier-core/src/wasi/wire/dns.rs | 248 + easytier-core/src/wasi/wire/mod.rs | 8 + easytier-core/src/wasi/wire/options.rs | 394 ++ easytier-core/src/wasi/wire/socket.rs | 207 + .../testdata/wasi_core_instance_create.json | 14 + easytier-core/tests/feature_profiles.rs | 21 + easytier-gui/src-tauri/Cargo.toml | 1 + easytier-gui/src-tauri/src/lib.rs | 119 +- easytier-gui/src/auto-imports.d.ts | 2 + easytier-proto/Cargo.toml | 101 + easytier-proto/build/main.rs | 142 + {easytier => easytier-proto}/build/rpc.rs | 1 + .../src => easytier-proto}/proto/acl.proto | 0 .../proto/api_config.proto | 0 .../proto/api_instance.proto | 0 .../proto/api_logger.proto | 0 .../proto/api_manage.proto | 0 .../src => easytier-proto}/proto/common.proto | 0 easytier-proto/proto/core_config.proto | 56 + easytier-proto/proto/core_peer.proto | 77 + .../src => easytier-proto}/proto/error.proto | 0 .../proto/magic_dns.proto | 0 .../proto/peer_rpc.proto | 0 .../src => easytier-proto}/proto/tests.proto | 0 .../src => easytier-proto}/proto/web.proto | 0 .../src/proto => easytier-proto/src}/acl.rs | 2 + .../src/proto => easytier-proto/src}/api.rs | 114 +- .../proto => easytier-proto/src}/common.rs | 50 +- easytier-proto/src/core_config.rs | 3 + easytier-proto/src/core_peer.rs | 5 + .../src/proto => easytier-proto/src}/error.rs | 0 easytier-proto/src/lib.rs | 33 + .../proto => easytier-proto/src}/magic_dns.rs | 1 + .../proto => easytier-proto/src}/peer_rpc.rs | 140 +- .../src}/rpc_types/__rt.rs | 1 - .../src}/rpc_types/controller.rs | 0 .../src}/rpc_types/descriptor.rs | 0 .../src}/rpc_types/error.rs | 2 +- .../src}/rpc_types/handler.rs | 0 .../src}/rpc_types/mod.rs | 0 easytier-proto/src/tests.rs | 2 + .../src/proto => easytier-proto/src}/web.rs | 0 easytier-web/Cargo.toml | 3 +- .../frontend-lib/scripts/codegen-proto.mjs | 2 +- .../src/client_manager/managed_config.rs | 13 +- easytier-web/src/client_manager/mod.rs | 97 +- .../src/client_manager/runtime_reconcile.rs | 6 +- easytier-web/src/client_manager/session.rs | 20 +- .../session/runtime_revision.rs | 16 +- .../db/entity/user_running_network_configs.rs | 6 +- easytier-web/src/db/mod.rs | 14 +- easytier-web/src/main.rs | 25 +- easytier-web/src/restful/mod.rs | 3 +- easytier-web/src/restful/network.rs | 9 +- easytier/Cargo.toml | 207 +- easytier/benches/README.md | 129 +- easytier/benches/packet_bytes_extraction.rs | 2 +- easytier/benches/tx_throughput.rs | 472 -- easytier/build/main.rs | 120 +- easytier/src/common/config.rs | 1849 +----- easytier/src/common/constants.rs | 33 +- easytier/src/common/credential_manager.rs | 48 + easytier/src/common/dns.rs | 435 +- easytier/src/common/error.rs | 34 +- easytier/src/common/global_ctx.rs | 825 +-- easytier/src/common/idn.rs | 70 - easytier/src/common/ifcfg/mod.rs | 49 +- easytier/src/common/ifcfg/netlink.rs | 652 +-- easytier/src/common/ifcfg/netlink_wire.rs | 682 +++ easytier/src/common/ifcfg/route.rs | 89 +- easytier/src/common/ifcfg/win/luid.rs | 382 +- easytier/src/common/ifcfg/win/mod.rs | 1 - easytier/src/common/ifcfg/win/netsh.rs | 118 - easytier/src/common/ifcfg/win/types.rs | 36 - easytier/src/common/log.rs | 501 -- easytier/src/common/log/file.rs | 221 + easytier/src/common/log/management.rs | 19 + easytier/src/common/log/mod.rs | 638 +++ easytier/src/common/log/tracing_backend.rs | 121 + easytier/src/common/mod.rs | 128 +- easytier/src/common/netns.rs | 10 + easytier/src/common/network.rs | 424 +- easytier/src/common/stun.rs | 1711 +----- .../common/tracing_rolling_appender/mod.rs | 23 +- easytier/src/common/upnp.rs | 412 +- easytier/src/connector/direct.rs | 1115 ---- easytier/src/connector/dns_connector.rs | 260 - easytier/src/connector/http_connector.rs | 361 -- easytier/src/connector/manual.rs | 525 -- easytier/src/connector/mod.rs | 462 -- easytier/src/connector/tcp_hole_punch.rs | 778 --- .../connector/udp_hole_punch/both_easy_sym.rs | 422 -- .../src/connector/udp_hole_punch/common.rs | 850 --- easytier/src/connector/udp_hole_punch/cone.rs | 308 - easytier/src/connector/udp_hole_punch/mod.rs | 702 --- .../connector/udp_hole_punch/sym_to_cone.rs | 723 --- easytier/src/core.rs | 25 +- easytier/src/easytier-cli.rs | 27 +- easytier/src/gateway/fast_socks5/util/mod.rs | 2 - .../src/gateway/fast_socks5/util/stream.rs | 65 - easytier/src/gateway/hedge.rs | 95 + easytier/src/gateway/icmp_proxy.rs | 523 +- easytier/src/gateway/ip_reassembler.rs | 325 -- easytier/src/gateway/kcp_proxy.rs | 607 +- easytier/src/gateway/mod.rs | 96 +- easytier/src/gateway/quic_proxy.rs | 658 +-- easytier/src/gateway/socks5.rs | 1773 ------ easytier/src/gateway/socks5/dataplane.rs | 541 -- easytier/src/gateway/tcp_proxy.rs | 1036 ---- easytier/src/gateway/udp_proxy.rs | 817 --- easytier/src/gateway/wrapped_proxy.rs | 153 - easytier/src/host_runtime.rs | 235 + easytier/src/instance/cli_event_logger.rs | 272 + easytier/src/instance/composition.rs | 529 ++ easytier/src/instance/config.rs | 184 + easytier/src/instance/config_storage.rs | 42 + .../instance/dns_server/client_instance.rs | 133 +- easytier/src/instance/dns_server/mod.rs | 15 +- easytier/src/instance/dns_server/runner.rs | 23 +- easytier/src/instance/dns_server/server.rs | 54 +- .../instance/dns_server/server_instance.rs | 305 +- .../dns_server/system_config/linux.rs | 362 -- .../instance/dns_server/system_config/mod.rs | 3 - .../dns_server/system_config/windows.rs | 12 +- easytier/src/instance/dns_server/tests.rs | 137 +- easytier/src/instance/factory.rs | 200 + easytier/src/instance/host.rs | 59 + easytier/src/instance/instance.rs | 1880 ------- easytier/src/instance/listeners.rs | 553 +- easytier/src/instance/mod.rs | 24 +- easytier/src/instance/proxy_cidrs_monitor.rs | 95 - easytier/src/instance/public_ipv6_provider.rs | 2039 +------ .../instance/public_ipv6_provider/linux.rs | 1487 +++++ .../public_ipv6_provider/unsupported.rs | 79 + easytier/src/instance/runtime_host.rs | 114 + .../instance/runtime_host/event_journal.rs | 132 + .../instance/runtime_host/implementation.rs | 44 + .../src/instance/runtime_host/magic_dns.rs | 74 + .../src/instance/runtime_host/tun_common.rs | 78 + .../src/instance/runtime_host/tun_desktop.rs | 253 + .../src/instance/runtime_host/tun_disabled.rs | 109 + .../src/instance/runtime_host/tun_mobile.rs | 193 + easytier/src/instance/test_instance.rs | 135 + easytier/src/instance/udp_hole_punch.rs | 41 + easytier/src/instance/virtual_nic.rs | 246 +- .../src/instance/windows_udp_broadcast.rs | 1042 +--- .../windows_udp_broadcast/capture_raw.rs | 21 + .../capture_windivert.rs | 173 + .../instance/windows_udp_broadcast/runtime.rs | 359 ++ easytier/src/instance_manager.rs | 848 --- easytier/src/launcher.rs | 1674 ------ easytier/src/lib.rs | 18 +- easytier/src/peer_center/mod.rs | 52 - easytier/src/peers/credential_manager.rs | 667 --- easytier/src/peers/encrypt/mod.rs | 107 - easytier/src/peers/encrypt/openssl.rs | 201 - easytier/src/peers/encrypt/ring.rs | 252 - easytier/src/peers/foreign_network_manager.rs | 2597 --------- easytier/src/peers/mod.rs | 73 - easytier/src/peers/peer.rs | 566 -- easytier/src/peers/peer_manager.rs | 3756 ------------- easytier/src/peers/peer_rpc.rs | 347 -- easytier/src/peers/peer_rpc_service.rs | 288 - easytier/src/peers/peer_task.rs | 205 - easytier/src/peers/rpc_service.rs | 300 - easytier/src/peers/tests.rs | 1627 ------ easytier/src/proto/mod.rs | 23 +- easytier/src/proto/rpc/mod.rs | 3 + easytier/src/proto/rpc/standalone.rs | 111 + easytier/src/proto/rpc_impl/standalone.rs | 224 - easytier/src/proto/tests.rs | 76 +- easytier/src/proto/utils.rs | 101 - easytier/src/rpc_service/acl_manage.rs | 50 - easytier/src/rpc_service/api.rs | 306 +- easytier/src/rpc_service/config.rs | 49 - easytier/src/rpc_service/connector_manage.rs | 36 - easytier/src/rpc_service/credential_manage.rs | 62 - easytier/src/rpc_service/instance_manage.rs | 1301 ----- easytier/src/rpc_service/json_rpc.rs | 232 - easytier/src/rpc_service/logger.rs | 105 +- .../src/rpc_service/mapped_listener_manage.rs | 38 - easytier/src/rpc_service/mod.rs | 131 +- easytier/src/rpc_service/peer_center.rs | 108 - easytier/src/rpc_service/peer_manage.rs | 116 - .../src/rpc_service/port_forward_manage.rs | 36 - easytier/src/rpc_service/protected_port.rs | 61 - easytier/src/rpc_service/proxy.rs | 41 - easytier/src/rpc_service/stats.rs | 50 - easytier/src/rpc_service/vpn_portal.rs | 36 - .../src/{tunnel => socket}/fake_tcp/LICENSE | 0 easytier/src/socket/fake_tcp/mod.rs | 533 ++ .../fake_tcp/netfilter/linux_bpf.rs | 6 +- .../fake_tcp/netfilter/macos_bpf.rs | 2 +- .../fake_tcp/netfilter/mod.rs | 0 .../fake_tcp/netfilter/pnet.rs | 53 +- .../fake_tcp/netfilter/windivert.rs | 2 +- .../src/{tunnel => socket}/fake_tcp/packet.rs | 2 - .../src/{tunnel => socket}/fake_tcp/stack.rs | 40 +- easytier/src/socket/mod.rs | 5 + easytier/src/socket/tcp.rs | 374 ++ easytier/src/socket/udp.rs | 343 ++ easytier/src/socket/udp_src.rs | 13 + easytier/src/socket/udp_src/fallback.rs | 81 + easytier/src/socket/udp_src/unix.rs | 570 ++ easytier/src/socket/udp_src/windows.rs | 716 +++ easytier/src/tests/credential_tests.rs | 643 +-- easytier/src/tests/ipv6_test.rs | 34 +- easytier/src/tests/mod.rs | 107 +- easytier/src/tests/three_node.rs | 1285 ++--- easytier/src/tests/upnp_test.rs | 900 +-- easytier/src/tunnel/buf.rs | 92 - easytier/src/tunnel/common.rs | 496 +- easytier/src/tunnel/fake_tcp/mod.rs | 594 -- easytier/src/tunnel/insecure_tls.rs | 86 - easytier/src/tunnel/mod.rs | 272 +- easytier/src/tunnel/mpsc.rs | 256 - easytier/src/tunnel/protocol.rs | 481 ++ easytier/src/tunnel/protocol/adapters/mod.rs | 43 + easytier/src/tunnel/protocol/adapters/quic.rs | 88 + .../src/tunnel/protocol/adapters/websocket.rs | 104 + .../src/tunnel/protocol/adapters/wireguard.rs | 100 + easytier/src/tunnel/quic.rs | 1064 ++-- easytier/src/tunnel/quic/session_socket.rs | 279 + easytier/src/tunnel/ring.rs | 391 -- easytier/src/tunnel/tcp.rs | 378 -- easytier/src/tunnel/udp.rs | 1517 ----- easytier/src/tunnel/udp_src.rs | 210 - easytier/src/tunnel/unix.rs | 218 - easytier/src/tunnel/websocket.rs | 609 +- easytier/src/tunnel/wireguard.rs | 769 ++- easytier/src/utils/error.rs | 58 - easytier/src/utils/mod.rs | 4 +- easytier/src/utils/string.rs | 4 - easytier/src/utils/task.rs | 142 - easytier/src/vpn_portal/mod.rs | 48 - easytier/src/vpn_portal/wireguard.rs | 392 +- easytier/src/web_client/controller.rs | 65 - easytier/src/web_client/mod.rs | 345 +- easytier/src/web_client/session.rs | 170 - 523 files changed, 102813 insertions(+), 67095 deletions(-) create mode 100644 CONTEXT.md create mode 100644 docs/core-architecture.md create mode 100644 docs/data-plane-runtime-plan.md delete mode 100644 easytier-contrib/easytier-android-jni/kotlin/com/easytier/jni/EasyTierDataPlaneJNI.kt delete mode 100644 easytier-contrib/easytier-android-jni/src/data_plane_api.rs create mode 100644 easytier-contrib/easytier-ffi/DATA_PLANE_ABI.md delete mode 100644 easytier-contrib/easytier-ffi/examples/example_data_plane_async.c delete mode 100644 easytier-contrib/easytier-ffi/examples/go/README.md delete mode 100644 easytier-contrib/easytier-ffi/examples/go/easytier.go delete mode 100644 easytier-contrib/easytier-ffi/examples/go/easytier_async.go delete mode 100644 easytier-contrib/easytier-ffi/examples/go/easytier_async_test.go delete mode 100644 easytier-contrib/easytier-ffi/examples/go/easytier_test.go delete mode 100644 easytier-contrib/easytier-ffi/examples/go/go.mod delete mode 100644 easytier-contrib/easytier-ffi/src/data_plane.rs create mode 100644 easytier-contrib/easytier-ffi/src/data_plane/abi.rs create mode 100644 easytier-contrib/easytier-ffi/src/data_plane/mod.rs create mode 100644 easytier-contrib/easytier-ffi/src/data_plane/session.rs delete mode 100644 easytier-contrib/easytier-ffi/src/data_plane_async.rs create mode 100644 easytier-core/Cargo.toml create mode 100644 easytier-core/src/config/api.rs create mode 100644 easytier-core/src/config/api_input.rs create mode 100644 easytier-core/src/config/encryption.rs create mode 100644 easytier-core/src/config/gateway.rs create mode 100644 easytier-core/src/config/mod.rs create mode 100644 easytier-core/src/config/peers.rs create mode 100644 easytier-core/src/config/runtime.rs create mode 100644 easytier-core/src/config/toml.rs create mode 100644 easytier-core/src/config/toml/snapshot.rs create mode 100644 easytier-core/src/connectivity/composite.rs create mode 100644 easytier-core/src/connectivity/connector_host.rs create mode 100644 easytier-core/src/connectivity/direct/mod.rs create mode 100644 easytier-core/src/connectivity/direct/udp.rs create mode 100644 easytier-core/src/connectivity/hole_punch/mod.rs create mode 100644 easytier-core/src/connectivity/hole_punch/peer_adapters.rs create mode 100644 easytier-core/src/connectivity/hole_punch/policy.rs create mode 100644 easytier-core/src/connectivity/hole_punch/port_mapping.rs create mode 100644 easytier-core/src/connectivity/hole_punch/tcp.rs create mode 100644 easytier-core/src/connectivity/hole_punch/udp/binding.rs create mode 100644 easytier-core/src/connectivity/hole_punch/udp/client.rs create mode 100644 easytier-core/src/connectivity/hole_punch/udp/common.rs create mode 100644 easytier-core/src/connectivity/hole_punch/udp/connector.rs create mode 100644 easytier-core/src/connectivity/hole_punch/udp/mod.rs create mode 100644 easytier-core/src/connectivity/hole_punch/udp/punch_listener.rs create mode 100644 easytier-core/src/connectivity/hole_punch/udp/rpc.rs create mode 100644 easytier-core/src/connectivity/hole_punch/udp/runtime.rs create mode 100644 easytier-core/src/connectivity/hole_punch/udp/server.rs create mode 100644 easytier-core/src/connectivity/hole_punch/udp/socket_array.rs create mode 100644 easytier-core/src/connectivity/hole_punch/udp/task.rs create mode 100644 easytier-core/src/connectivity/manual/discovery.rs create mode 100644 easytier-core/src/connectivity/manual/discovery/implementation.rs create mode 100644 easytier-core/src/connectivity/manual/mod.rs create mode 100644 easytier-core/src/connectivity/mod.rs create mode 100644 easytier-core/src/connectivity/protocol/mod.rs create mode 100644 easytier-core/src/connectivity/protocol/raw.rs create mode 100644 easytier-core/src/connectivity/stun/client.rs create mode 100644 easytier-core/src/connectivity/stun/collector.rs create mode 100644 easytier-core/src/connectivity/stun/mod.rs create mode 100644 easytier-core/src/connectivity/stun/responder.rs create mode 100644 easytier-core/src/connectivity/transport/mod.rs create mode 100644 easytier-core/src/connectivity/transport/tcp.rs create mode 100644 easytier-core/src/connectivity/transport/udp.rs create mode 100644 easytier-core/src/events.rs create mode 100644 easytier-core/src/foundation/mod.rs create mode 100644 easytier-core/src/foundation/operation_broker.rs rename easytier/src/common/stats_manager.rs => easytier-core/src/foundation/stats.rs (75%) create mode 100644 easytier-core/src/foundation/task.rs create mode 100644 easytier-core/src/foundation/time.rs rename {easytier/src/common => easytier-core/src/foundation}/token_bucket.rs (79%) create mode 100644 easytier-core/src/gateway/dataplane/deadline.rs create mode 100644 easytier-core/src/gateway/dataplane/error.rs create mode 100644 easytier-core/src/gateway/dataplane/flow.rs create mode 100644 easytier-core/src/gateway/dataplane/mod.rs create mode 100644 easytier-core/src/gateway/dataplane/operation.rs create mode 100644 easytier-core/src/gateway/dataplane/packet.rs create mode 100644 easytier-core/src/gateway/dataplane/resource.rs create mode 100644 easytier-core/src/gateway/dataplane/route.rs create mode 100644 easytier-core/src/gateway/dataplane/session.rs create mode 100644 easytier-core/src/gateway/dataplane/stack.rs create mode 100644 easytier-core/src/gateway/dataplane/tcp.rs create mode 100644 easytier-core/src/gateway/dataplane/tests.rs create mode 100644 easytier-core/src/gateway/dataplane/udp.rs create mode 100644 easytier-core/src/gateway/dhcp.rs create mode 100644 easytier-core/src/gateway/magic_dns.rs create mode 100644 easytier-core/src/gateway/magic_dns/packet.rs create mode 100644 easytier-core/src/gateway/magic_dns/records.rs create mode 100644 easytier-core/src/gateway/mod.rs create mode 100644 easytier-core/src/gateway/port_forward.rs create mode 100644 easytier-core/src/gateway/proxy/cidr_monitor.rs create mode 100644 easytier-core/src/gateway/proxy/cidr_table.rs create mode 100644 easytier-core/src/gateway/proxy/icmp_host.rs create mode 100644 easytier-core/src/gateway/proxy/icmp_proxy_engine.rs create mode 100644 easytier-core/src/gateway/proxy/icmp_proxy_service.rs create mode 100644 easytier-core/src/gateway/proxy/ip_reassembler.rs create mode 100644 easytier-core/src/gateway/proxy/mod.rs create mode 100644 easytier-core/src/gateway/proxy/proxy_acl.rs create mode 100644 easytier-core/src/gateway/proxy/service.rs create mode 100644 easytier-core/src/gateway/proxy/tcp_proxy_engine.rs create mode 100644 easytier-core/src/gateway/proxy/tcp_proxy_service.rs create mode 100644 easytier-core/src/gateway/proxy/tcp_socket_connector.rs create mode 100644 easytier-core/src/gateway/proxy/traits.rs create mode 100644 easytier-core/src/gateway/proxy/udp_proxy_engine.rs create mode 100644 easytier-core/src/gateway/proxy/udp_proxy_service.rs create mode 100644 easytier-core/src/gateway/proxy/udp_socket_runtime.rs create mode 100644 easytier-core/src/gateway/proxy/wrapped_tcp_proxy.rs create mode 100644 easytier-core/src/gateway/proxy/wrapped_transport.rs create mode 100644 easytier-core/src/gateway/proxy/wrapped_transport/connect_api.rs create mode 100644 easytier-core/src/gateway/proxy/wrapped_transport/engine_api.rs create mode 100644 easytier-core/src/gateway/proxy/wrapped_transport/packet_api.rs create mode 100644 easytier-core/src/gateway/proxy/wrapped_transport/packet_plane.rs create mode 100644 easytier-core/src/gateway/proxy/wrapped_transport/packet_plane_disabled.rs create mode 100644 easytier-core/src/gateway/proxy/wrapped_transport_destination.rs create mode 100644 easytier-core/src/gateway/smoltcp/mod.rs create mode 100644 easytier-core/src/gateway/smoltcp/stack.rs rename {easytier/src/gateway => easytier-core/src/gateway/smoltcp}/tokio_smoltcp/channel_device.rs (100%) rename {easytier/src/gateway => easytier-core/src/gateway/smoltcp}/tokio_smoltcp/device.rs (100%) rename {easytier/src/gateway => easytier-core/src/gateway/smoltcp}/tokio_smoltcp/mod.rs (77%) rename {easytier/src/gateway => easytier-core/src/gateway/smoltcp}/tokio_smoltcp/reactor.rs (91%) rename {easytier/src/gateway => easytier-core/src/gateway/smoltcp}/tokio_smoltcp/socket.rs (88%) rename {easytier/src/gateway => easytier-core/src/gateway/smoltcp}/tokio_smoltcp/socket_allocator.rs (100%) create mode 100644 easytier-core/src/gateway/socks5/adapter.rs rename easytier/src/gateway/fast_socks5/mod.rs => easytier-core/src/gateway/socks5/codec.rs (56%) rename {easytier/src/gateway/fast_socks5/util => easytier-core/src/gateway/socks5/codec}/target_addr.rs (75%) create mode 100644 easytier-core/src/gateway/socks5/host.rs create mode 100644 easytier-core/src/gateway/socks5/mod.rs rename {easytier/src/gateway/fast_socks5 => easytier-core/src/gateway/socks5}/server.rs (69%) create mode 100644 easytier-core/src/gateway/udp_broadcast.rs create mode 100644 easytier-core/src/gateway/vpn_portal.rs create mode 100644 easytier-core/src/host/dns.rs create mode 100644 easytier-core/src/host/environment.rs create mode 100644 easytier-core/src/host/mod.rs create mode 100644 easytier-core/src/host/packet.rs create mode 100644 easytier-core/src/host/socket/factory.rs create mode 100644 easytier-core/src/host/socket/listener.rs create mode 100644 easytier-core/src/host/socket/mod.rs create mode 100644 easytier-core/src/host/socket/udp.rs create mode 100644 easytier-core/src/host/testkit.rs create mode 100644 easytier-core/src/instance/build_capabilities.rs create mode 100644 easytier-core/src/instance/config.rs create mode 100644 easytier-core/src/instance/data_plane_extension.rs create mode 100644 easytier-core/src/instance/lifecycle.rs create mode 100644 easytier-core/src/instance/management.rs create mode 100644 easytier-core/src/instance/management_extension.rs create mode 100644 easytier-core/src/instance/management_state.rs create mode 100644 easytier-core/src/instance/manager.rs create mode 100644 easytier-core/src/instance/mod.rs create mode 100644 easytier-core/src/instance/packet_io.rs create mode 100644 easytier-core/src/instance/packet_plane.rs create mode 100644 easytier-core/src/instance/packet_proxy_extension.rs create mode 100644 easytier-core/src/instance/public_ipv6_extension.rs create mode 100644 easytier-core/src/instance/test_utils.rs create mode 100644 easytier-core/src/instance/tests.rs create mode 100644 easytier-core/src/instance/vpn_portal_extension.rs create mode 100644 easytier-core/src/lib.rs create mode 100644 easytier-core/src/listener/mod.rs create mode 100644 easytier-core/src/listener/plan.rs create mode 100644 easytier-core/src/listener/transport.rs create mode 100644 easytier-core/src/management/full/compiled.rs create mode 100644 easytier-core/src/management/full/config_patch.rs create mode 100644 easytier-core/src/management/full/instance_info.rs create mode 100644 easytier-core/src/management/full/logger_rpc.rs create mode 100644 easytier-core/src/management/full/mod.rs create mode 100644 easytier-core/src/management/full/packet_proxy.rs create mode 100644 easytier-core/src/management/full/process_rpc.rs rename {easytier/src/rpc_service => easytier-core/src/management/full}/remote_client.rs (92%) create mode 100644 easytier-core/src/management/full/web_client.rs create mode 100644 easytier-core/src/management/instance_rpc/full.rs create mode 100644 easytier-core/src/management/instance_rpc/mod.rs create mode 100644 easytier-core/src/management/instance_rpc/packet_proxy.rs create mode 100644 easytier-core/src/management/instance_rpc/projection.rs create mode 100644 easytier-core/src/management/mod.rs create mode 100644 easytier-core/src/management/rpc_server_hook.rs create mode 100644 easytier-core/src/management/selector.rs create mode 100644 easytier-core/src/management/server.rs rename {easytier/src/common => easytier-core/src/packet}/compressor.rs (71%) create mode 100644 easytier-core/src/packet/compressor/zstd.rs create mode 100644 easytier-core/src/packet/hole_punch.rs rename easytier/src/tunnel/packet_def.rs => easytier-core/src/packet/mod.rs (87%) rename easytier/src/common/stun_codec_ext.rs => easytier-core/src/packet/stun.rs (91%) rename easytier/src/peers/acl_filter.rs => easytier-core/src/peers/acl/filter.rs (59%) create mode 100644 easytier-core/src/peers/acl/mod.rs rename easytier/src/common/acl_processor.rs => easytier-core/src/peers/acl/processor.rs (86%) create mode 100644 easytier-core/src/peers/admission.rs create mode 100644 easytier-core/src/peers/conn/mod.rs create mode 100644 easytier-core/src/peers/conn/peer.rs rename {easytier/src/peers => easytier-core/src/peers/conn}/peer_conn.rs (54%) rename {easytier/src/peers => easytier-core/src/peers/conn}/peer_conn_ping.rs (92%) rename {easytier/src/peers => easytier-core/src/peers/conn}/peer_map.rs (79%) rename {easytier/src/peers => easytier-core/src/peers/conn}/peer_session.rs (91%) create mode 100644 easytier-core/src/peers/context.rs create mode 100644 easytier-core/src/peers/credential_manager.rs create mode 100644 easytier-core/src/peers/error.rs rename easytier/src/peers/foreign_network_client.rs => easytier-core/src/peers/foreign_network/client.rs (81%) create mode 100644 easytier-core/src/peers/foreign_network/mod.rs create mode 100644 easytier-core/src/peers/mod.rs rename {easytier/src => easytier-core/src/peers}/peer_center/instance.rs (67%) create mode 100644 easytier-core/src/peers/peer_center/mod.rs rename {easytier/src => easytier-core/src/peers}/peer_center/server.rs (97%) create mode 100644 easytier-core/src/peers/peer_manager.rs create mode 100644 easytier-core/src/peers/peer_rpc.rs create mode 100644 easytier-core/src/peers/public_ipv6/mod.rs create mode 100644 easytier-core/src/peers/public_ipv6/provider.rs rename easytier/src/peers/public_ipv6.rs => easytier-core/src/peers/public_ipv6/service.rs (72%) rename {easytier => easytier-core}/src/peers/relay_peer_map.rs (87%) rename {easytier/src/peers => easytier-core/src/peers/route}/graph_algo.rs (100%) rename easytier/src/peers/route_trait.rs => easytier-core/src/peers/route/mod.rs (90%) rename {easytier/src/peers => easytier-core/src/peers/route}/peer_ospf_route.rs (53%) create mode 100644 easytier-core/src/peers/route/route_peer_wire.rs create mode 100644 easytier-core/src/peers/test_support.rs create mode 100644 easytier-core/src/peers/tests.rs rename {easytier => easytier-core}/src/peers/traffic_metrics.rs (83%) create mode 100644 easytier-core/src/peers/util.rs create mode 100644 easytier-core/src/peers/whitelist.rs create mode 100644 easytier-core/src/process_runtime.rs rename {easytier/src/proto/rpc_impl => easytier-core/src/rpc}/bidirect.rs (67%) rename {easytier/src/proto/rpc_impl => easytier-core/src/rpc}/client.rs (59%) rename {easytier/src/proto/rpc_impl => easytier-core/src/rpc}/mod.rs (75%) rename {easytier/src/proto/rpc_impl => easytier-core/src/rpc}/packet.rs (51%) rename {easytier/src/proto/rpc_impl => easytier-core/src/rpc}/server.rs (61%) rename {easytier/src/proto/rpc_impl => easytier-core/src/rpc}/service_registry.rs (100%) create mode 100644 easytier-core/src/rpc/standalone.rs create mode 100644 easytier-core/src/socket/mod.rs create mode 100644 easytier-core/src/socket/ring.rs create mode 100644 easytier-core/src/socket/tcp.rs create mode 100644 easytier-core/src/socket/udp/layer.rs create mode 100644 easytier-core/src/socket/udp/listener.rs create mode 100644 easytier-core/src/socket/udp/mod.rs create mode 100644 easytier-core/src/socket/udp/packet.rs create mode 100644 easytier-core/src/socket/udp/session.rs create mode 100644 easytier-core/src/socket/udp/tests.rs create mode 100644 easytier-core/src/socket/udp/virtual_socket.rs rename {easytier/src/peers => easytier-core/src/tunnel}/encrypt/aes_gcm.rs (96%) create mode 100644 easytier-core/src/tunnel/encrypt/chacha20.rs create mode 100644 easytier-core/src/tunnel/encrypt/mod.rs rename {easytier/src/peers => easytier-core/src/tunnel}/encrypt/xor.rs (89%) rename {easytier => easytier-core}/src/tunnel/filter.rs (78%) create mode 100644 easytier-core/src/tunnel/framed.rs create mode 100644 easytier-core/src/tunnel/mod.rs create mode 100644 easytier-core/src/tunnel/mpsc.rs create mode 100644 easytier-core/src/tunnel/ring.rs rename {easytier/src/peers => easytier-core/src/tunnel}/secure_datagram.rs (97%) rename {easytier => easytier-core}/src/tunnel/stats.rs (98%) create mode 100644 easytier-core/src/tunnel/tcp.rs create mode 100644 easytier-core/src/tunnel/udp.rs rename easytier/src/web_client/security.rs => easytier-core/src/tunnel/web_security.rs (94%) create mode 100644 easytier-core/src/tunnel/wrapper.rs create mode 100644 easytier-core/src/wasi/abi.rs create mode 100644 easytier-core/src/wasi/adapter/dns.rs create mode 100644 easytier-core/src/wasi/adapter/environment.rs create mode 100644 easytier-core/src/wasi/adapter/mod.rs create mode 100644 easytier-core/src/wasi/adapter/packet.rs create mode 100644 easytier-core/src/wasi/adapter/socket/backend.rs create mode 100644 easytier-core/src/wasi/adapter/socket/mod.rs create mode 100644 easytier-core/src/wasi/adapter/socket/udp.rs create mode 100644 easytier-core/src/wasi/imports.rs create mode 100644 easytier-core/src/wasi/mod.rs create mode 100644 easytier-core/src/wasi/runtime.rs create mode 100644 easytier-core/src/wasi/runtime/abi/data_plane.rs create mode 100644 easytier-core/src/wasi/runtime_driver.rs create mode 100644 easytier-core/src/wasi/schema.rs create mode 100644 easytier-core/src/wasi/time.rs create mode 100644 easytier-core/src/wasi/wire/common.rs create mode 100644 easytier-core/src/wasi/wire/data_plane.rs create mode 100644 easytier-core/src/wasi/wire/dns.rs create mode 100644 easytier-core/src/wasi/wire/mod.rs create mode 100644 easytier-core/src/wasi/wire/options.rs create mode 100644 easytier-core/src/wasi/wire/socket.rs create mode 100644 easytier-core/testdata/wasi_core_instance_create.json create mode 100644 easytier-core/tests/feature_profiles.rs create mode 100644 easytier-proto/Cargo.toml create mode 100644 easytier-proto/build/main.rs rename {easytier => easytier-proto}/build/rpc.rs (99%) rename {easytier/src => easytier-proto}/proto/acl.proto (100%) rename {easytier/src => easytier-proto}/proto/api_config.proto (100%) rename {easytier/src => easytier-proto}/proto/api_instance.proto (100%) rename {easytier/src => easytier-proto}/proto/api_logger.proto (100%) rename {easytier/src => easytier-proto}/proto/api_manage.proto (100%) rename {easytier/src => easytier-proto}/proto/common.proto (100%) create mode 100644 easytier-proto/proto/core_config.proto create mode 100644 easytier-proto/proto/core_peer.proto rename {easytier/src => easytier-proto}/proto/error.proto (100%) rename {easytier/src => easytier-proto}/proto/magic_dns.proto (100%) rename {easytier/src => easytier-proto}/proto/peer_rpc.proto (100%) rename {easytier/src => easytier-proto}/proto/tests.proto (100%) rename {easytier/src => easytier-proto}/proto/web.proto (100%) rename {easytier/src/proto => easytier-proto/src}/acl.rs (98%) rename {easytier/src/proto => easytier-proto/src}/api.rs (74%) rename {easytier/src/proto => easytier-proto/src}/common.rs (93%) create mode 100644 easytier-proto/src/core_config.rs create mode 100644 easytier-proto/src/core_peer.rs rename {easytier/src/proto => easytier-proto/src}/error.rs (100%) create mode 100644 easytier-proto/src/lib.rs rename {easytier/src/proto => easytier-proto/src}/magic_dns.rs (79%) rename {easytier/src/proto => easytier-proto/src}/peer_rpc.rs (72%) rename {easytier/src/proto => easytier-proto/src}/rpc_types/__rt.rs (97%) rename {easytier/src/proto => easytier-proto/src}/rpc_types/controller.rs (100%) rename {easytier/src/proto => easytier-proto/src}/rpc_types/descriptor.rs (100%) rename {easytier/src/proto => easytier-proto/src}/rpc_types/error.rs (95%) rename {easytier/src/proto => easytier-proto/src}/rpc_types/handler.rs (100%) rename {easytier/src/proto => easytier-proto/src}/rpc_types/mod.rs (100%) create mode 100644 easytier-proto/src/tests.rs rename {easytier/src/proto => easytier-proto/src}/web.rs (100%) delete mode 100644 easytier/benches/tx_throughput.rs create mode 100644 easytier/src/common/credential_manager.rs delete mode 100644 easytier/src/common/idn.rs create mode 100644 easytier/src/common/ifcfg/netlink_wire.rs delete mode 100644 easytier/src/common/ifcfg/win/netsh.rs delete mode 100644 easytier/src/common/log.rs create mode 100644 easytier/src/common/log/file.rs create mode 100644 easytier/src/common/log/management.rs create mode 100644 easytier/src/common/log/mod.rs create mode 100644 easytier/src/common/log/tracing_backend.rs delete mode 100644 easytier/src/connector/direct.rs delete mode 100644 easytier/src/connector/dns_connector.rs delete mode 100644 easytier/src/connector/http_connector.rs delete mode 100644 easytier/src/connector/manual.rs delete mode 100644 easytier/src/connector/mod.rs delete mode 100644 easytier/src/connector/tcp_hole_punch.rs delete mode 100644 easytier/src/connector/udp_hole_punch/both_easy_sym.rs delete mode 100644 easytier/src/connector/udp_hole_punch/common.rs delete mode 100644 easytier/src/connector/udp_hole_punch/cone.rs delete mode 100644 easytier/src/connector/udp_hole_punch/mod.rs delete mode 100644 easytier/src/connector/udp_hole_punch/sym_to_cone.rs delete mode 100644 easytier/src/gateway/fast_socks5/util/mod.rs delete mode 100644 easytier/src/gateway/fast_socks5/util/stream.rs create mode 100644 easytier/src/gateway/hedge.rs delete mode 100644 easytier/src/gateway/ip_reassembler.rs delete mode 100644 easytier/src/gateway/socks5.rs delete mode 100644 easytier/src/gateway/socks5/dataplane.rs delete mode 100644 easytier/src/gateway/tcp_proxy.rs delete mode 100644 easytier/src/gateway/udp_proxy.rs delete mode 100644 easytier/src/gateway/wrapped_proxy.rs create mode 100644 easytier/src/host_runtime.rs create mode 100644 easytier/src/instance/cli_event_logger.rs create mode 100644 easytier/src/instance/composition.rs create mode 100644 easytier/src/instance/config.rs create mode 100644 easytier/src/instance/config_storage.rs delete mode 100644 easytier/src/instance/dns_server/system_config/linux.rs create mode 100644 easytier/src/instance/factory.rs create mode 100644 easytier/src/instance/host.rs delete mode 100644 easytier/src/instance/instance.rs delete mode 100644 easytier/src/instance/proxy_cidrs_monitor.rs create mode 100644 easytier/src/instance/public_ipv6_provider/linux.rs create mode 100644 easytier/src/instance/public_ipv6_provider/unsupported.rs create mode 100644 easytier/src/instance/runtime_host.rs create mode 100644 easytier/src/instance/runtime_host/event_journal.rs create mode 100644 easytier/src/instance/runtime_host/implementation.rs create mode 100644 easytier/src/instance/runtime_host/magic_dns.rs create mode 100644 easytier/src/instance/runtime_host/tun_common.rs create mode 100644 easytier/src/instance/runtime_host/tun_desktop.rs create mode 100644 easytier/src/instance/runtime_host/tun_disabled.rs create mode 100644 easytier/src/instance/runtime_host/tun_mobile.rs create mode 100644 easytier/src/instance/test_instance.rs create mode 100644 easytier/src/instance/udp_hole_punch.rs create mode 100644 easytier/src/instance/windows_udp_broadcast/capture_raw.rs create mode 100644 easytier/src/instance/windows_udp_broadcast/capture_windivert.rs create mode 100644 easytier/src/instance/windows_udp_broadcast/runtime.rs delete mode 100644 easytier/src/instance_manager.rs delete mode 100644 easytier/src/launcher.rs delete mode 100644 easytier/src/peer_center/mod.rs delete mode 100644 easytier/src/peers/credential_manager.rs delete mode 100644 easytier/src/peers/encrypt/mod.rs delete mode 100644 easytier/src/peers/encrypt/openssl.rs delete mode 100644 easytier/src/peers/encrypt/ring.rs delete mode 100644 easytier/src/peers/foreign_network_manager.rs delete mode 100644 easytier/src/peers/mod.rs delete mode 100644 easytier/src/peers/peer.rs delete mode 100644 easytier/src/peers/peer_manager.rs delete mode 100644 easytier/src/peers/peer_rpc.rs delete mode 100644 easytier/src/peers/peer_rpc_service.rs delete mode 100644 easytier/src/peers/peer_task.rs delete mode 100644 easytier/src/peers/rpc_service.rs delete mode 100644 easytier/src/peers/tests.rs create mode 100644 easytier/src/proto/rpc/mod.rs create mode 100644 easytier/src/proto/rpc/standalone.rs delete mode 100644 easytier/src/proto/rpc_impl/standalone.rs delete mode 100644 easytier/src/proto/utils.rs delete mode 100644 easytier/src/rpc_service/acl_manage.rs delete mode 100644 easytier/src/rpc_service/config.rs delete mode 100644 easytier/src/rpc_service/connector_manage.rs delete mode 100644 easytier/src/rpc_service/credential_manage.rs delete mode 100644 easytier/src/rpc_service/instance_manage.rs delete mode 100644 easytier/src/rpc_service/json_rpc.rs delete mode 100644 easytier/src/rpc_service/mapped_listener_manage.rs delete mode 100644 easytier/src/rpc_service/peer_center.rs delete mode 100644 easytier/src/rpc_service/peer_manage.rs delete mode 100644 easytier/src/rpc_service/port_forward_manage.rs delete mode 100644 easytier/src/rpc_service/protected_port.rs delete mode 100644 easytier/src/rpc_service/proxy.rs delete mode 100644 easytier/src/rpc_service/stats.rs delete mode 100644 easytier/src/rpc_service/vpn_portal.rs rename easytier/src/{tunnel => socket}/fake_tcp/LICENSE (100%) create mode 100644 easytier/src/socket/fake_tcp/mod.rs rename easytier/src/{tunnel => socket}/fake_tcp/netfilter/linux_bpf.rs (99%) rename easytier/src/{tunnel => socket}/fake_tcp/netfilter/macos_bpf.rs (99%) rename easytier/src/{tunnel => socket}/fake_tcp/netfilter/mod.rs (100%) rename easytier/src/{tunnel => socket}/fake_tcp/netfilter/pnet.rs (86%) rename easytier/src/{tunnel => socket}/fake_tcp/netfilter/windivert.rs (99%) rename easytier/src/{tunnel => socket}/fake_tcp/packet.rs (99%) rename easytier/src/{tunnel => socket}/fake_tcp/stack.rs (94%) create mode 100644 easytier/src/socket/mod.rs create mode 100644 easytier/src/socket/tcp.rs create mode 100644 easytier/src/socket/udp.rs create mode 100644 easytier/src/socket/udp_src.rs create mode 100644 easytier/src/socket/udp_src/fallback.rs create mode 100644 easytier/src/socket/udp_src/unix.rs create mode 100644 easytier/src/socket/udp_src/windows.rs delete mode 100644 easytier/src/tunnel/buf.rs delete mode 100644 easytier/src/tunnel/fake_tcp/mod.rs delete mode 100644 easytier/src/tunnel/insecure_tls.rs delete mode 100644 easytier/src/tunnel/mpsc.rs create mode 100644 easytier/src/tunnel/protocol.rs create mode 100644 easytier/src/tunnel/protocol/adapters/mod.rs create mode 100644 easytier/src/tunnel/protocol/adapters/quic.rs create mode 100644 easytier/src/tunnel/protocol/adapters/websocket.rs create mode 100644 easytier/src/tunnel/protocol/adapters/wireguard.rs create mode 100644 easytier/src/tunnel/quic/session_socket.rs delete mode 100644 easytier/src/tunnel/ring.rs delete mode 100644 easytier/src/tunnel/tcp.rs delete mode 100644 easytier/src/tunnel/udp.rs delete mode 100644 easytier/src/tunnel/udp_src.rs delete mode 100644 easytier/src/tunnel/unix.rs delete mode 100644 easytier/src/utils/error.rs delete mode 100644 easytier/src/utils/task.rs delete mode 100644 easytier/src/web_client/controller.rs delete mode 100644 easytier/src/web_client/session.rs diff --git a/.github/workflows/core.yml b/.github/workflows/core.yml index 033952c9..185fd2cd 100644 --- a/.github/workflows/core.yml +++ b/.github/workflows/core.yml @@ -36,7 +36,7 @@ jobs: concurrent_skipping: 'same_content_newer' skip_after_successful_duplicate: 'true' cancel_others: 'true' - paths: '["Cargo.toml", "Cargo.lock", "easytier/**", ".github/workflows/core.yml", ".github/actions/**", "easytier-web/**"]' + paths: '["Cargo.toml", "Cargo.lock", "easytier/**", "easytier-core/**", "easytier-proto/**", ".github/workflows/core.yml", ".github/actions/**", "easytier-web/**"]' build_web: runs-on: ubuntu-latest needs: pre_job diff --git a/.github/workflows/test.yml b/.github/workflows/test.yml index 599f11ad..e28d6751 100644 --- a/.github/workflows/test.yml +++ b/.github/workflows/test.yml @@ -34,7 +34,7 @@ jobs: # All of these options are optional, so you can remove them if you are happy with the defaults concurrent_skipping: 'never' skip_after_successful_duplicate: 'true' - paths: '["Cargo.toml", "Cargo.lock", "easytier/**", ".github/workflows/test.yml", ".github/actions/**"]' + paths: '["Cargo.toml", "Cargo.lock", "easytier/**", "easytier-core/**", "easytier-proto/**", "easytier-web/**", "easytier-gui/src-tauri/**", "easytier-contrib/**", ".github/workflows/test.yml", ".github/actions/**"]' check: name: Run linters & check @@ -98,7 +98,9 @@ jobs: - uses: taiki-e/install-action@nextest - name: Archive test - run: cargo nextest archive --archive-file tests.tar.zst --package easytier --features full + run: >- + cargo nextest archive --archive-file tests.tar.zst + --package easytier --package easytier-core --features full - uses: actions/upload-artifact@v5 with: diff --git a/CONTEXT.md b/CONTEXT.md new file mode 100644 index 00000000..350566bc --- /dev/null +++ b/CONTEXT.md @@ -0,0 +1,23 @@ +# EasyTier Domain Context + +## Module layers + +`easytier-core` layers dependencies from `foundation` upward through the +portable networking domains. `foundation` contains infrastructure Modules +that have no dependency on a networking domain and may be used by any higher +layer. + +## Operation broker + +An operation broker owns the lifecycle of asynchronous work submitted by an +external caller to core. It allocates opaque operation IDs, arbitrates +completion, cancellation, and disposal, retains terminal outcomes, and +publishes a batch-drainable completion queue. + +The broker does not interpret operation kinds, outcomes, resources, wire +formats, or domain errors. Each domain Module owns those semantics and composes +the broker under the same lock as any state that must change atomically with an +operation transition. + +Host capability operations use a separate seam. They turn Host readiness into +Rust task wakeups and do not share the caller-to-core broker state machine. diff --git a/Cargo.lock b/Cargo.lock index ef1042b4..ce04bd43 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -1943,17 +1943,6 @@ version = "2.8.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "575f75dfd25738df5b91b8e43e14d44bda14637a58fae779fd2b064f8bf3e010" -[[package]] -name = "dbus" -version = "0.9.7" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "1bb21987b9fb1613058ba3843121dd18b163b254d8a6e797e144cbac14d96d1b" -dependencies = [ - "libc", - "libdbus-sys", - "winapi", -] - [[package]] name = "defguard_wireguard_rs" version = "0.4.2" @@ -1965,7 +1954,7 @@ dependencies = [ "log", "netlink-packet-core", "netlink-packet-generic", - "netlink-packet-route 0.17.1", + "netlink-packet-route", "netlink-packet-utils", "netlink-packet-wireguard", "netlink-sys", @@ -2012,17 +2001,6 @@ dependencies = [ "thiserror 1.0.63", ] -[[package]] -name = "delegate" -version = "0.13.5" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "780eb241654bf097afb00fc5f054a09b687dad862e485fdcf8399bb056565370" -dependencies = [ - "proc-macro2", - "quote", - "syn 2.0.117", -] - [[package]] name = "der" version = "0.7.10" @@ -2044,17 +2022,6 @@ dependencies = [ "serde_core", ] -[[package]] -name = "derivative" -version = "2.2.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "fcc3dd5e9e9c0b295d6e1e4d811fb6f157d5ffd784b8d202fc62eac8035a770b" -dependencies = [ - "proc-macro2", - "quote", - "syn 1.0.109", -] - [[package]] name = "derive_arbitrary" version = "1.4.1" @@ -2316,19 +2283,14 @@ checksum = "d0881ea181b1df73ff77ffaaf9c7544ecc11e82fba9b5f27b262a3c73a332555" name = "easytier" version = "2.6.4" dependencies = [ - "aes-gcm", "anyhow", "arc-swap", - "ariadne", "async-recursion", - "async-ringbuf", - "async-stream", "async-trait", "atomic-shim", "atomic_refcell", "auto_impl", "base64 0.22.1", - "bitflags 2.8.0", "bon", "boringtun-easytier", "bytecodec", @@ -2345,12 +2307,11 @@ dependencies = [ "crossbeam", "ctor 0.8.0", "dashmap", - "dbus", "defguard_wireguard_rs", - "delegate", - "derivative", "derive_builder", "derive_more 2.1.1", + "easytier-core", + "easytier-proto", "encoding", "flume 0.12.0", "forwarded-header-value", @@ -2364,57 +2325,35 @@ dependencies = [ "hickory-proto", "hickory-resolver", "hickory-server", - "hmac", "http", - "http_req", "humansize", "humantime-serde", - "idna 1.0.3", "igd-next", "indoc", - "itertools 0.14.0", "kcp-sys", + "log", "machine-uid", "maplit", "mimalloc", "moka", - "multimap", "natpmp", - "netlink-packet-core", - "netlink-packet-route 0.21.0", - "netlink-packet-utils", "netlink-sys", "network-interface", "nix 0.29.0", "once_cell", - "openssl", - "ordered_hash_map", "parking_lot", "paste", - "pbjson", - "pbjson-build", "percent-encoding", - "petgraph", "pin-project-lite", "pnet", - "prefix-trie", - "proc-macro2", "prost 0.14.3", - "prost-build", - "prost-reflect", - "prost-reflect-build", - "prost-wkt-types", "quanta", "quinn", "quinn-proto", - "quote", "rand 0.8.5", "rcgen", "regex", - "reqwest 0.12.12", - "resolv-conf", "ring", - "ringbuf", "rstest", "rust-i18n", "rustls", @@ -2423,10 +2362,7 @@ dependencies = [ "serde_json", "serial_test", "service-manager", - "sha2", "shellexpand", - "smoltcp", - "snow", "socket2 0.5.10", "strum 0.27.2", "stun_codec", @@ -2439,12 +2375,9 @@ dependencies = [ "tikv-jemalloc-ctl", "tikv-jemalloc-sys", "tikv-jemallocator", - "time", - "timedmap", "tokio", "tokio-rustls", "tokio-socks", - "tokio-stream", "tokio-util", "tokio-websockets", "toml 0.8.19", @@ -2454,9 +2387,6 @@ dependencies = [ "unicode-width 0.1.11", "url", "uuid", - "version-compare", - "which 7.0.3", - "wildmatch", "winapi", "windivert", "windows 0.62.2", @@ -2464,8 +2394,6 @@ dependencies = [ "winreg 0.52.0", "x25519-dalek", "zerocopy 0.7.35", - "zip", - "zstd", ] [[package]] @@ -2482,21 +2410,82 @@ dependencies = [ "serde_json", ] +[[package]] +name = "easytier-core" +version = "2.6.4" +dependencies = [ + "aes-gcm", + "anyhow", + "arc-swap", + "ariadne", + "async-ringbuf", + "async-trait", + "atomic-shim", + "auto_impl", + "base64 0.22.1", + "bitflags 2.8.0", + "bytecodec", + "bytes", + "chacha20poly1305", + "chrono", + "cidr", + "crossbeam", + "dashmap", + "derive_builder", + "easytier-proto", + "futures", + "guarden 0.2.0", + "hmac", + "http-body-util", + "hyper", + "hyper-util", + "idna 1.0.3", + "ordered_hash_map", + "parking_lot", + "percent-encoding", + "petgraph", + "pin-project-lite", + "pnet_packet", + "prefix-trie", + "prost 0.14.3", + "prost-types 0.14.3", + "quanta", + "rand 0.8.5", + "rustls", + "serde", + "serde_json", + "sha2", + "smoltcp", + "snow", + "strum 0.27.2", + "stun_codec", + "thiserror 1.0.63", + "tokio", + "tokio-rustls", + "tokio-util", + "toml 0.8.19", + "tracing", + "url", + "uuid", + "webpki-roots 0.26.3", + "wildmatch", + "x25519-dalek", + "zerocopy 0.7.35", + "zstd", +] + [[package]] name = "easytier-ffi" version = "0.1.0" dependencies = [ "async-trait", - "dashmap", "easytier", + "easytier-core", "log", "once_cell", - "percent-encoding", "serde", "serde_json", "tokio", - "tokio-util", - "url", "uuid", ] @@ -2510,6 +2499,7 @@ dependencies = [ "dashmap", "dunce", "easytier", + "easytier-core", "gethostname 1.1.0", "libc", "once_cell", @@ -2533,6 +2523,39 @@ dependencies = [ "windows 0.52.0", ] +[[package]] +name = "easytier-proto" +version = "2.6.4" +dependencies = [ + "anyhow", + "async-trait", + "auto_impl", + "base64 0.22.1", + "bytes", + "chrono", + "cidr", + "hmac", + "indoc", + "pbjson", + "pbjson-build", + "proc-macro2", + "prost 0.14.3", + "prost-build", + "prost-types 0.14.3", + "prost-wkt-types", + "quote", + "reqwest 0.12.12", + "serde", + "serde_json", + "sha2", + "thiserror 1.0.63", + "tokio", + "url", + "uuid", + "x25519-dalek", + "zip", +] + [[package]] name = "easytier-uptime" version = "0.1.0" @@ -2588,6 +2611,7 @@ dependencies = [ "clap", "dashmap", "easytier", + "easytier-core", "image 0.24.9", "imageproc", "maxminddb", @@ -2836,12 +2860,6 @@ dependencies = [ "syn 2.0.117", ] -[[package]] -name = "env_home" -version = "0.1.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "c7f84e12ccf0a7ddc17a6c41c93326024c42920d7ee630d04950e6926645c0fe" - [[package]] name = "env_logger" version = "0.10.2" @@ -4011,22 +4029,6 @@ dependencies = [ "pin-project-lite", ] -[[package]] -name = "http_req" -version = "0.13.1" -source = "git+https://github.com/EasyTier/http_req.git#b10aa9fc0db3067cc3d2174683a87250b80a1ea9" -dependencies = [ - "base64 0.22.1", - "rand 0.8.5", - "rustls", - "rustls-pemfile", - "rustls-pki-types", - "unicase", - "webpki", - "webpki-roots 0.26.3", - "zeroize", -] - [[package]] name = "httparse" version = "1.9.4" @@ -4834,16 +4836,6 @@ version = "0.2.186" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "68ab91017fe16c622486840e4c83c9a37afeff978bd239b5293d61ece587de66" -[[package]] -name = "libdbus-sys" -version = "0.2.5" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "06085512b750d640299b79be4bad3d2fa90a9c00b1fd9e1b46364f66f0485c72" -dependencies = [ - "cc", - "pkg-config", -] - [[package]] name = "libloading" version = "0.7.4" @@ -5272,9 +5264,6 @@ name = "multimap" version = "0.10.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "1d87ecb2933e8aeadb3e3a02b828fed80a7528047e68b4f424523a0981a3a084" -dependencies = [ - "serde", -] [[package]] name = "nalgebra" @@ -5360,7 +5349,7 @@ dependencies = [ "ipnet", "libc", "netlink-packet-core", - "netlink-packet-route 0.17.1", + "netlink-packet-route", "netlink-sys", "once_cell", "system-configuration", @@ -5404,21 +5393,6 @@ dependencies = [ "netlink-packet-utils", ] -[[package]] -name = "netlink-packet-route" -version = "0.21.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "483325d4bfef65699214858f097d504eb812c38ce7077d165f301ec406c3066e" -dependencies = [ - "anyhow", - "bitflags 2.8.0", - "byteorder", - "libc", - "log", - "netlink-packet-core", - "netlink-packet-utils", -] - [[package]] name = "netlink-packet-utils" version = "0.5.2" @@ -5707,15 +5681,6 @@ dependencies = [ "syn 2.0.117", ] -[[package]] -name = "num_threads" -version = "0.1.7" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "5c7398b9c8b70908f6371f47ed36737907c87c52af34c268fed0bf0ceb92ead9" -dependencies = [ - "libc", -] - [[package]] name = "oauth2" version = "5.0.0" @@ -6030,15 +5995,6 @@ version = "0.1.5" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "ff011a302c396a5197692431fc1948019154afc178baf7d8e37367442a4601cf" -[[package]] -name = "openssl-src" -version = "300.5.2+3.5.2" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "d270b79e2926f5150189d475bc7e9d2c69f9c4697b185fa917d5a32b792d21b4" -dependencies = [ - "cc", -] - [[package]] name = "openssl-sys" version = "0.9.103" @@ -6047,7 +6003,6 @@ checksum = "7f9e8deee91df40a943c71b917e5874b951d32a802526c85721ce3b776c929d6" dependencies = [ "cc", "libc", - "openssl-src", "pkg-config", "vcpkg", ] @@ -7013,41 +6968,6 @@ dependencies = [ "syn 2.0.117", ] -[[package]] -name = "prost-reflect" -version = "0.16.4" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "590aa145fee8f7a26b5a6055365e7c5e89a5c1caae9869de76ec0ee73181a2f9" -dependencies = [ - "base64 0.22.1", - "prost 0.14.3", - "prost-reflect-derive", - "prost-types 0.14.3", - "serde", - "serde-value", -] - -[[package]] -name = "prost-reflect-build" -version = "0.16.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "8214ae2c30bbac390db0134d08300e770ef89b6d4e5abf855e8d300eded87e28" -dependencies = [ - "prost-build", - "prost-reflect", -] - -[[package]] -name = "prost-reflect-derive" -version = "0.16.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "7b6d90e29fa6c0d13c2c19ba5e4b3fb0efbf5975d27bcf4e260b7b15455bcabe" -dependencies = [ - "proc-macro2", - "quote", - "syn 2.0.117", -] - [[package]] name = "prost-types" version = "0.13.5" @@ -8658,7 +8578,7 @@ dependencies = [ "encoding_rs", "plist", "sys-info", - "which 4.4.2", + "which", "xml-rs", ] @@ -9944,9 +9864,7 @@ checksum = "743bd48c283afc0388f9b8827b976905fb217ad9e647fae3a379a9283c4def2c" dependencies = [ "deranged", "itoa", - "libc", "num-conv", - "num_threads", "powerfmt", "serde_core", "time-core", @@ -9969,12 +9887,6 @@ dependencies = [ "time-core", ] -[[package]] -name = "timedmap" -version = "1.0.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "825f6c8a18bc36d56a62f66af7296385b628c9c5543a8663d4c217fc920bfefd" - [[package]] name = "tinystr" version = "0.7.6" @@ -10501,7 +10413,6 @@ dependencies = [ "sharded-slab", "smallvec", "thread_local", - "time", "tracing", "tracing-core", "tracing-log", @@ -11233,16 +11144,6 @@ dependencies = [ "system-deps", ] -[[package]] -name = "webpki" -version = "0.22.4" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "ed63aea5ce73d0ff405984102c42de94fc55a6b75765d621c65262469b3c9b53" -dependencies = [ - "ring", - "untrusted", -] - [[package]] name = "webpki-root-certs" version = "0.26.11" @@ -11333,18 +11234,6 @@ dependencies = [ "rustix 0.38.34", ] -[[package]] -name = "which" -version = "7.0.3" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "24d643ce3fd3e5b54854602a080f34fb10ab75e0b813ee32d00ca2b44fa74762" -dependencies = [ - "either", - "env_home", - "rustix 1.0.7", - "winsafe", -] - [[package]] name = "whoami" version = "1.6.1" @@ -12128,12 +12017,6 @@ dependencies = [ "windows-sys 0.59.0", ] -[[package]] -name = "winsafe" -version = "0.0.19" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "d135d17ab770252ad95e9a872d365cf3090e3be864a34ab46f48555993efc904" - [[package]] name = "wintun" version = "0.5.0" diff --git a/Cargo.toml b/Cargo.toml index 2a2eb144..9e76b9b7 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -1,6 +1,8 @@ [workspace] resolver = "2" members = [ + "easytier-core", + "easytier-proto", "easytier", "easytier-gui/src-tauri", "easytier-web", diff --git a/docs/core-architecture.md b/docs/core-architecture.md new file mode 100644 index 00000000..83c5ff5d --- /dev/null +++ b/docs/core-architecture.md @@ -0,0 +1,522 @@ +# EasyTier Core Architecture + +## Status and scope + +This document describes the current architecture after the portable-core +refactor. It is the source of truth for ownership, dependency direction, +feature boundaries, and validation. It intentionally records the resulting +design rather than the migration history. + +The refactor has three principal crate roles: + +- `easytier-core` owns portable EasyTier configuration, protocol state, + routing, peer state, connectivity orchestration, packet processing, and + instance lifecycle. +- `easytier` is the native composition root. It owns operating-system + resources, native protocol engines, process integration, CLI and native + presentation. +- `easytier-proto` owns generated protobuf and RPC types, descriptor data, and + the feature slices needed by core and presentation users. + +`easytier-core` is designed to compile without direct operating-system network +access. It supports native hosts through Rust traits and has a target-only WASI +adapter and ABI implementation under `easytier-core/src/wasi`. + +This architecture does not require compatibility with old internal module +paths. Wire compatibility, configuration compatibility, management semantics, +and externally used application behaviour remain compatibility requirements. + +## Architectural vocabulary + +The following terms have specific meanings in this document: + +- **Module**: an interface and the implementation hidden behind it. +- **Host**: the process or runtime embedding core and owning platform + resources. +- **Host capability**: an operation core may request but must not implement + with direct OS calls. +- **Adapter**: a concrete implementation of a Host capability or protocol + extension. +- **Composition root**: code that creates core configuration, Host Adapters, + instances, and process-level services. +- **Runtime configuration**: the authoritative normalized state used after an + instance starts. +- **Packet plane**: portable packet classification, routing, transformation, + proxy/NAT state, and forwarding decisions. + +New abstractions should pass a deletion test: deleting a useful deep Module +should force non-trivial policy or lifecycle logic to reappear in multiple +callers. A pass-through wrapper with no independent invariant is not an +architectural boundary. + +## Crate dependency direction + +The principal dependency direction is: + +```text +easytier-proto <- easytier-core <- easytier +``` + +Presentation crates and platform integrations consume these crates. Portable +policy must not move outward merely because one current consumer is native. +Conversely, core must not absorb an OS mechanism or a protocol engine whose +dependencies cannot satisfy the core target contract. + +### `easytier-proto` + +The protobuf crate is split by public Cargo features: + +- `core` provides the common wire messages, peer RPC messages, generated RPC + runtime, and descriptor bytes needed by core. +- `api` adds management API messages. +- protocol-specific features add only their generated message modules. +- `json-rpc` enables the well-known protobuf JSON types used by the management + plane. +- `full` is the compatibility aggregate used by complete products. + +The core crate depends on `easytier-proto` with default features disabled and +enables only `core`, adding API or JSON-RPC types through its own management +features. + +The main core/native path has no `prost-reflect` dependency. OSPF route +reflection uses the focused wire editor in +`peers/route/route_peer_wire.rs`. It retains the original encoded +`RoutePeerInfo`, replaces only the fields credential filtering is allowed to +change, and leaves all other top-level and nested fields intact. This is +required so unknown fields survive mixed-version, multi-hop propagation. +Generated Rust types remain responsible for normal message construction and +validation. + +Descriptor sets are still generated and embedded by `easytier-proto`; removing +runtime reflection did not remove descriptor data used by configuration and +RPC tooling. The OHOS integration has its own schema service and dependency +policy and is outside this replacement. + +### `easytier-core` + +Core owns portable behaviour and exposes capability seams. Its normal +dependencies use Tokio runtime, time, synchronization, and I/O traits without +requiring the full Tokio feature set. + +Core may depend on optional portable engines when their owning feature is +enabled. It does not create real native TCP/UDP sockets, alter routes, open a +TUN device, enter a network namespace, configure system DNS, manage a service, +or invoke UPnP/NAT-PMP directly. + +### `easytier` + +The native crate owns: + +- process startup, shutdown, signals, service management, and allocators; +- filesystem configuration input and persistence; +- real TCP/UDP, DNS, TUN, raw-socket, route, interface, namespace, and socket + option operations; +- UPnP and NAT-PMP operations; +- Unix and FakeTCP resources; +- WebSocket/WSS, QUIC, WireGuard, and KCP concrete engines; +- native Magic DNS serving and system DNS integration; +- CLI, web, GUI, FFI, and native management presentation. + +Native code may translate values and assemble Adapters. It must not maintain a +second peer graph, reproduce core routing or hole-punch policy, or invent an +alternative instance lifecycle. + +## Internal core layers + +The physical module layout follows this downward order: + +```text +foundation + <- config / packet + <- socket + <- host + <- tunnel + <- listener / connectivity + <- peers / rpc + <- gateway + <- instance + <- management +``` + +`process_runtime` is a process- or module-scoped owner shared by instances. +`wasi` is target integration and is compiled only for tests or the WASI target; +it is not an additional portable domain layer. + +### Foundation + +`foundation/` contains task supervision, the time facade, rate limiting, +statistics primitives, and the domain-neutral external operation broker. The +broker owns asynchronous operation lifecycle and completion storage while the +calling domain owns operation kinds, outcomes, resources, and errors. +Foundation must not depend on a domain layer. + +### Configuration and packets + +`config/` owns: + +- the complete `TomlConfig` model; +- parsing, serialization, and validation; +- OS-independent defaults; +- peer, encryption, gateway, and API input models; +- normalized runtime snapshots and the live runtime configuration store. + +The Host supplies platform facts through `CoreInstanceHostConfig`. Core applies +the policy that combines those facts with TOML input. This is especially +important for a WASI build: the compile-time guest target cannot be used as a +proxy for the Host operating system. + +`packet/` owns EasyTier packet structures, compression, STUN and hole-punch +wire codecs. It does not own socket I/O or connection policy. + +### Socket and Host seams + +`socket/` contains transport-neutral primitives: + +- `SocketContext`, including IP-family policy, optional socket mark, and an + opaque network-namespace token; +- virtual TCP socket, listener, and factory traits; +- virtual UDP socket and factory traits; +- UDP session multiplexing, classification, and lifecycle; +- in-process Ring sockets. + +`host/` is the single home of Host capability seams: + +- DNS and DNS record resolution; +- connector environment observations; +- packet ingress and egress; +- Host socket operation bridges and handle-based TCP/UDP/listener adapters. + +Core owns scheduling, backpressure, cancellation, UDP session state, and +protocol state even when each actual operation crosses a Host Adapter. A Host +Adapter owns the real resource and performs the OS operation. + +The native `NativeHostRuntime` is process-wide and does not retain an instance +`GlobalCtx`, namespace guard, socket mark, or connectivity state. Differences +between instances travel in each request's `SocketContext`. A narrow +instance-host projection may expose listener and interface facts, but it does +not become another socket factory. + +### Tunnel and listener + +A socket is a raw communication endpoint. A Tunnel is an EasyTier connection +created by adding framing, metadata, handshakes, and protocol lifecycle. + +Core owns: + +- raw TCP framing and upgrade; +- UDP tunnel/session framing and classification; +- Ring Tunnel identity and registry state; +- encryption and secure-datagram policy that is portable; +- client/server protocol selection interfaces; +- listener planning, optional/required listener policy, retry, accept + scheduling, running-listener registry, and orderly shutdown. + +Native protocol Adapters own WebSocket/WSS, QUIC, WireGuard, and KCP engines. +Unix and FakeTCP are socket resources that feed a core protocol upgrader; they +are not independent owners of EasyTier peer state. + +Each protocol registration must provide a coherent client/server Adapter. +Unavailable configured transports must be rejected during validation or +protocol selection, rather than silently falling back to another transport. + +### Connectivity + +`connectivity/` owns: + +- manual connection and endpoint discovery policy; +- direct candidate selection; +- retry, backoff, blacklists, and listener reuse; +- STUN requests, responses, probing, NAT inference, and published endpoint + state; +- TCP and UDP hole-punch state machines; +- UDP port-mapping policy and lease lifecycle; +- conversion of successful sockets into protocol-upgrade requests. + +The Host owns DNS execution, socket syscalls, interface enumeration, bind +device/mark/namespace operations, and concrete UPnP/NAT-PMP calls. STUN-only +hole punching remains available when the Host does not supply a port-mapping +Adapter. + +Some connectivity files intentionally implement peer-facing adapter traits for +`PeerManagerCore`. These are localized integration edges between adjacent +domains, not permission for lower socket or Host layers to depend on peers. + +### Peers and RPC + +`peers/` is the authoritative owner of: + +- admission and connection sessions; +- peer maps and connection lifecycle; +- ACL and whitelist decisions; +- OSPF route calculation and graph algorithms; +- peer and credential RPC registration; +- foreign-network admission, identity, relay, and lifecycle; +- peer-center state and public IPv6 policy; +- traffic metrics and peer snapshots. + +Submodules progress from kernel types and utilities, through ACL/context, +connection state, route state, manager services, and finally foreign-network +and peer-center composition. Callers consume the public surface declared by +the domain rather than reaching into a parallel native peer owner. + +`rpc/` owns the peer-flavoured RPC transport, packet fragmentation, client and +server lifecycle, handler registry, and standalone listener/client lifecycle. +Generated service descriptors and message types remain in `easytier-proto`. + +### Gateway + +`gateway/` owns portable packet-plane features: + +- proxy CIDR state and monitoring policy; +- packet parsing, reassembly, NAT/proxy state, and TCP/UDP/ICMP decisions; +- the smoltcp-backed portable dataplane selected by its feature; +- SOCKS5 framing, authentication, association, routing, and session state; +- wrapped-transport planning and session state used by KCP and QUIC Adapters; +- DHCP allocation policy; +- Magic DNS route and response policy; +- VPN portal client/session policy; +- UDP broadcast classification and rewrite policy. + +TUN, raw sockets, transparent-destination lookup, concrete protocol engines, +native DNS servers, namespace operations, and route application stay in native +Adapters. + +Optional gateway capabilities are selected by cohesive Modules. Disabled +implementations retain stable lifecycle calls and report unsupported +configuration where a stable interface is required; they do not duplicate +portable policy. + +The instance-scoped `DataPlaneSession` composes the foundation operation broker +under the same session lock as its resource and quota state. The broker owns +generic completion, cancellation, free, drain, and take transitions. The data +plane retains TCP/UDP resource ownership, operation metadata, route deadlines, +and error semantics. + +The proposed restructuring of the smoltcp data plane, SOCKS5 and port-forward +Adapters, portable KCP engine, event-driven FFI/WASI completion model, and Go +Host integration is tracked in +[`data-plane-runtime-plan.md`](data-plane-runtime-plan.md). That document is a +future implementation plan; this document remains the source of truth for the +currently implemented architecture until the plan is completed. + +### Instance and management + +`CoreInstance::new(CoreInstanceConfig, CoreHostAdapters)` is the sole direct +construction path for a normalized instance. `CoreInstance::from_toml` uses +the same normalization and construction path. Core constructs the peer graph, +runtime store, STUN collector, connectivity managers, listener runtime, packet +plane, gateway runtimes, and lifecycle owners. + +A core instance: + +- owns all mutable portable state for one network; +- is one-shot after `stop`; +- exposes one complete `start` and one `stop` lifecycle interface; +- starts Modules in a fixed serial composition order without cross-Module + started flags or staged activation; +- installs initial ACL, proxy CIDR, and manual-peer inputs before startup; +- serializes lifecycle operations with one instance-level operation lock; +- owns cooperative cancellation and component shutdown order; +- exposes `CorePacketPlane` as the narrow packet/route projection used by Host + dataplane Adapters; +- treats its normalized runtime store as authoritative after construction. + +`CoreHostAdapters` contains the required Host, DNS, packet sink, and +`CoreProcessRuntime`, plus optional protocol and platform capabilities. The +bundle carries capabilities, not preconstructed portable managers. + +Each Module owns partial-start cleanup for its internal resources. +`CoreInstance` has one outer cancellation and recovery path for the complete +serial startup. `Running` therefore means the Host runtime and every enabled +portable Module have started successfully; there is no separate post-Host +activation state. Host packet tasks stop before PeerManager resources are +cleared. + +`InstanceManager` is the canonical UUID-indexed instance collection for one +Host composition. Its `InstanceFactory` constructs one complete record before +the manager performs an atomic uniqueness check. The manager owns collection +membership; it does not own startup order, persistence, daemon policy, cached +errors, ABI handles, or RPC projections. + +`management/` consumes the canonical manager and instances. It owns: + +- stable UUID/name selection; +- read-only instance and peer management RPC; +- full process mutation and configuration transactions when enabled; +- persistence and logger-control capability interfaces; +- management listener/client lifecycle and JSON-RPC presentation. + +There is one process-level management entry. Instances and the manager do not +depend on management response projections. + +## Process-scoped state + +`CoreProcessRuntime` owns portable resources shared across instances in one +process or instantiated module: + +- the Ring Tunnel registry and namespace; +- a reference-counted protected TCP-port registry. + +The composition root creates and shares one runtime. Management listener ports +are protected before bind and held by leases after the concrete port is known. +Native and target adapters supply bound resources but do not implement a +second protected-port registry. + +Process-global capability objects may contain stateless or shared platform +mechanisms. They must not contain instance-specific peer, route, +configuration, or connectivity state. + +## Runtime configuration authority + +`TomlConfig` is an owned construction input. After startup, it is not a second +mutable source of truth. + +The normalized core runtime store is authoritative for: + +- peer feature flags and routing policy; +- listeners and initial peers; +- ACL and whitelist inputs; +- manual and VPN portal CIDRs; +- gateway and connectivity settings; +- runtime configuration patches. + +Host persistence is an effect following a successful core transaction. A Host +Adapter must not call back into an instance to obtain a hidden configuration +snapshot while core is applying an operation. + +Non-serializable resources such as TUN descriptors, packet sinks, execution +domains, and native protocol engines are construction context, not TOML +fields. + +## Logging + +The main native runtime uses a small logger implemented in +`easytier/src/common/log`: + +- `log` records and `tracing` events share console and file sinks; +- timestamps, compact formatting, optional terminal colours, `NO_COLOR`, and + basic `RUST_LOG` target/level filters are implemented directly; +- file rotation uses the existing EasyTier rolling appender; +- management RPC can reload the file level; +- an atomic maximum-level gate rejects disabled events before target matching + or file-filter locking; +- concurrent file-level reload serializes the filter and atomic-level update. + +File logging and no-file logging are separate selected backends. The default +tracing backend records events and deliberately ignores span trees. The +optional `tracing` feature selects the tokio-console subscriber integration; +only that diagnostic profile pulls the main crate's `tracing-subscriber` and +`console-subscriber` dependencies. + +Contrib applications and platform integrations may have independent logging +requirements and are not implicitly wired to the native process logger. + +## Feature model + +Features represent coherent capabilities, not arbitrary source fragments. +Important core feature relationships are: + +- `management-rpc` enables generated management API types and read-only + management services. +- `management` adds configuration writes, full management composition, rich + errors, and JSON-RPC. +- `proxy-packet` enables portable packet parsing/proxy machinery and the + required smoltcp packet features. +- `proxy-smoltcp-stack` adds the async TCP/UDP smoltcp stack. +- `dns-resolver` is the shared Hickory resolver leaf used by endpoint + discovery and Magic DNS without coupling either capability to the other. +- `endpoint-discovery` adds HTTPS endpoint discovery dependencies. +- `magic-dns` enables its DNS server, management wire messages, and portable + packet-query integration. +- `tcp-hole-punch` enables the TCP hole-punch runtime. +- `dhcp-ipv4`, `public-ipv6-provider`, `vpn-portal`, + `wrapped-transport`, and `proxy-cidr-monitor` are independent gateway or + platform-policy leaves. +- `extended-services` is the compatibility aggregate for those leaves. +- encryption and compression engines remain independently selectable. + +The native crate maps product features to the core and protocol features it +actually consumes. A protocol feature must not accidentally enable unrelated +gateway or management capabilities. + +Production feature and platform selection belongs at Module or Adapter +boundaries rather than inside shared implementations. The logger demonstrates +the intended pattern: file and tracing variants are complete backend modules +with one stable interface, so shared event processing contains no feature +branches. + +## Module boundaries + +The dependency directions in this document define the intended module +boundaries. Changes that require a new upward edge must first define a stable +lower-layer interface or explicitly revise this architecture. + +Modules are `pub(crate)` by default. Each domain's `mod.rs` declares its +outward surface. Public visibility is used for real cross-crate Host, +configuration, management, packet-plane, or test-support interfaces. + +## Architectural invariants + +1. Portable EasyTier policy has one owner in `easytier-core`. +2. Core does not perform real OS socket, DNS, TUN, route, filesystem + configuration, process, or service-manager operations. +3. Host-OS policy is runtime input; a WASI compile target is not Host policy. +4. Every real socket and DNS operation crosses a Host capability seam. +5. Core owns socket scheduling, backpressure, protocol state, and cancellation. +6. Dial, accept, and hole-punch paths produce sockets before protocol upgrade. +7. Peer admission consumes upgraded transports and does not create OS + resources. +8. Each instance owns its mutable peer, route, connectivity, gateway, and + runtime configuration state. +9. One Host composition has one canonical UUID-to-instance manager. +10. Process-level runtimes do not capture instance state. +11. `CoreInstance::new` is the sole normalized direct construction entry. +12. The manager owns membership, not lifecycle or presentation. +13. Management consumes the manager; the manager does not return management + projections. +14. Unknown protobuf fields in reflected route information survive forwarding + and credential filtering. +15. Feature selection is localized at cohesive Module/Adapter boundaries. +16. Unsupported configured capabilities fail explicitly rather than changing + wire protocol or silently falling back. + +## Validation + +Changes to these boundaries should run, at minimum: + +```text +cargo fmt --all -- --check +cargo check -p easytier-core -p easytier-proto -p easytier --features full +cargo test -p easytier-core --lib +``` + +Feature work should add focused checks for the changed no-default, isolated, +default, full, and cross-target profiles. Socket, TUN, namespace, protocol +engine, and multi-node changes require the relevant Docker integration tests. +WASI ABI or Adapter changes require a `wasm32-wasip1` build and target-side +tests. These compiler-resolved profiles are the authority for feature and +target boundaries. + +CI path filters include `easytier-core`, `easytier-proto`, native, web, GUI +Tauri, and contrib. The archived Rust test suite contains both `easytier` and +`easytier-core`. + +## Known limitations and debt + +- Some production feature and platform gates still select fields or statements + inside shared implementations. New code should prefer complete Module or + Adapter variants, and existing cases should move only when their owning + Module is changed. +- Connectivity retains localized Adapter implementations that name + `PeerManagerCore`; further decoupling requires an interface extraction, not + a visibility-only move. +- Native Linux namespace guards exist in paths that can cross async suspension. + Because `setns` is thread-local, those operations should eventually be kept + on one non-migrating execution context. +- QUIC session retirement after failed or exhausted accepted sessions remains + separate native-engine correctness work; it must preserve multiple + connections sharing one QUIC endpoint/session. + +These limitations are not reasons to add fallback owners or parallel state. +Fixes should preserve the ownership rules above and address the responsible +Module directly. diff --git a/docs/data-plane-runtime-plan.md b/docs/data-plane-runtime-plan.md new file mode 100644 index 00000000..2cfbbfe8 --- /dev/null +++ b/docs/data-plane-runtime-plan.md @@ -0,0 +1,1133 @@ +# Data Plane Runtime and Event-Driven Host Plan + +## Status + +Accepted. + +This document is the implementation plan for restructuring the EasyTier data +plane and exposing it through native FFI and the standalone +`easytier-go-host` project. It describes a target architecture, not the current +implementation. + +The implementation scope is: + +- `easytier-core`; +- the native `easytier` composition that injects the existing optional KCP + backend; +- `easytier-contrib/easytier-ffi`; +- the WASI guest ABI implemented by `easytier-core`; +- `/data/project/easytier-go-host`; +- TCP, UDP, and smoltcp data-plane paths; +- moving KCP route selection and source-connection ownership below the + `DataPlaneRuntime` Interface without making KCP portable. + +The implementation explicitly excludes integration-specific code. Downstream +consumers must be able to use the resulting standard Go network interfaces +without requiring consumer-specific code in either repository. + +It also excludes: + +- porting `kcp-sys` to `wasm32-wasip1`; +- changing KCP timers or cancellation behavior; +- installing KCP in the WASI composition; +- exposing KCP through the FFI v2 or Go data-plane session; +- compatibility with the existing native data-plane FFI. + +## Motivation + +The current gateway data plane works, but its locality is poor: + +- `gateway/dataplane/mod.rs` owns smoltcp lifecycle, packet demultiplexing, + flow state, SOCKS5 listener tasks, port forwarding, public data-plane + sockets, and route selection; +- public data-plane TCP connections use a connector named and designed for + SOCKS5; +- route and packet-flow types that are shared data-plane concepts live under + the SOCKS5 Module; +- the public data-plane connect path can silently select a Host direct socket + for a non-overlay destination; +- native FFI presents asynchronous work as per-operation + `start/status-or-wait/finish` calls rather than one instance completion + stream; +- the WASI runtime is cooperatively driven, but the data-plane result direction + is not exposed to a Host; +- the current Go artifact is not built with the smoltcp data plane. + +The objective is a deep `DataPlaneRuntime` Module. Callers select an explicit +route policy and receive TCP or UDP behavior without understanding PeerManager +lookups, smoltcp flows, optional KCP backend readiness, Host port reservations, +or cleanup rules. SOCKS5 and port forwarding become Adapters over that Module. +Native FFI and WASI use an instance-scoped operation broker and completion +queue. Go exposes only standard `net` interfaces. + +## Architectural decisions + +The following decisions are part of this plan. + +| Caller | Route policy | +| --- | --- | +| Go data plane | `OverlayOnly`, smoltcp only | +| Native FFI v2 data plane | `OverlayOnly`, smoltcp only | +| SOCKS5 TCP | `OverlayOrDirect`, existing optional KCP backend | +| Port forwarding | Preserve current behavior | + +Additional decisions: + +- Public FFI v2, WASI, and Go TCP use smoltcp. +- Native gateway Adapters may request the existing optional KCP backend. +- KCP selection and source-connection ownership live behind + `DataPlaneRuntime`; SOCKS5 no longer implements that policy. +- This refactor does not change the existing KCP backend implementation, + destination proxy, timers, or cancellation semantics. +- A selected KCP connection keeps the existing no-fallback behavior. +- An existing overlay route that is temporarily unusable never falls back to + a Host direct connection. +- The first version supports IPv4 only and rejects IPv6 explicitly. +- The first Go interface accepts IP literals only. Name resolution must + eventually be owned by core rather than being silently delegated to an + unrelated Host resolver. +- Public Go and FFI data-plane callers cannot opt into direct fallback. +- The cooperative WASI drive contract remains. It is hidden by the Go engine + and triggered only by real commands, completions, packets, runnable work, or + protocol deadlines. +- Go does not implement routing or smoltcp. + +The local virtual address is an overlay destination. An exact local +data-plane listener wins. Otherwise, a connection to the local virtual address +may use a `LocalHost` path to the Host loopback. This is not equivalent to +falling back to an unrelated direct destination. + +## Target architecture + +```text +CoreInstance + | + +-- DataPlaneRuntime + | - route planning + | - flow ownership + | - smoltcp stack generations + | - local endpoints + | - optional native KCP backend selection + | | + | v + | PeerManager / ACL / packet forwarding / Host socket seams + | + +-- SOCKS5 Adapter ----------- calls DataPlaneRuntime + | + +-- PortForward Adapter ------ calls DataPlaneRuntime + +Go Instance: Dial / Listen / ListenPacket + | + v +Go engine: one driver and completion dispatcher + | + v +coreabi Adapter: WASM exports, memory, canonical wire encoding + | + v +WASI DataPlane Adapter + | + v +DataPlaneSession + - ResourceTable + - data-plane operation metadata and quotas + | + +-- foundation::OperationBroker + - operation lifecycle + - CompletionQueue + | + +----------------------- calls DataPlaneRuntime +``` + +`DataPlaneRuntime`, `Socks5GatewayAdapter`, and `PortForwardAdapter` are sibling +components owned by `CoreInstance`. The dependency is one-way: both gateway +Adapters call the runtime's crate-private Interface. The runtime never owns, +starts, stops, imports, or names either Adapter. + +There are two opposite-direction seams and they must remain separate: + +1. `hostabi` and the Go reactor implement guest-to-Host underlay operations. +2. The data-plane ABI implements caller-to-guest overlay operations. + +Guest data-plane resources must not be stored in the Go Host reactor, and +underlay Host resources must not be exposed as public data-plane handles. + +## EasyTier core design + +### Module layout + +The target layout is: + +```text +easytier-core/src/gateway/ + dataplane/ + mod.rs + runtime.rs + route.rs + flow.rs + packet.rs + stack.rs + tcp.rs + udp.rs + session.rs + operation.rs + resource.rs + socks5/ + adapter.rs + codec.rs + server.rs + host.rs + port_forward.rs +``` + +The exact file split may change during implementation, but ownership must +follow these Modules rather than the current physical layout. + +### `DataPlaneRuntime` + +`CoreInstance` owns one `DataPlaneRuntime`. The Module owns: + +- smoltcp stack-generation lifecycle; +- PeerManager route queries used by the data plane; +- optional native KCP backend capability queries and selection; +- TCP flow and listener registration; +- UDP flow registration; +- peer-packet classification and delivery; +- local data-plane endpoint registration; +- stack, pipeline, and task shutdown; +- data-plane route and I/O errors. + +Its public Interface remains narrow: + +- connect an overlay TCP stream; +- bind an overlay TCP listener; +- bind an overlay UDP socket; +- start and stop with the owning core instance. + +Crate-private Adapter entry points additionally accept route policy, transport +preference, socket purpose, source hints, and an absolute deadline. The public +data-plane entry points always use `OverlayOnly` and `SmoltcpOnly`. + +The runtime must not return SOCKS5 errors or depend on SOCKS5 request types. +SOCKS5 error mapping belongs to the SOCKS5 Adapter. + +### TCP route planner + +`Socks5TcpConnectPlan` becomes a data-plane route planner. It is one pure +implementation, not a trait with a single Adapter. + +The planner returns one of: + +- `LocalEndpoint`; +- `LocalHost`; +- `Smoltcp`; +- `KcpBackend`; +- `Direct`. + +The route matrix is: + +1. An exact local data-plane listener selects `LocalEndpoint`. +2. The local virtual address selects `LocalHost` when no data-plane listener + owns the destination port. +3. An overlay route selects the optional KCP backend only when a native + gateway Adapter requests `PreferKcp`, the source engine is ready, and the + peer chain allows KCP; otherwise it selects smoltcp. +4. A destination without an overlay route returns `NoOverlayRoute` under + `OverlayOnly`. +5. A destination without an overlay route selects `Direct` only under + `OverlayOrDirect`. +6. A known overlay destination whose data-plane path is temporarily + unavailable returns `NotReady`; it never becomes `Direct`. + +One connection uses one route snapshot. A route update does not change a +stream's selected path after establishment. + +KCP backend connection failure is returned to the gateway Adapter. Adding a +fallback policy or exposing KCP to public data-plane sessions requires a +separate explicit design. + +### Deadlines and errors + +Readiness, route lookup, port reservation, and connection establishment use +one absolute deadline. The current combination of an outer duration and an +inner timeout rounded to seconds is removed. + +Core defines stable data-plane error kinds: + +- `Cancelled`; +- `DeadlineExceeded`; +- `InstanceStopped`; +- `HandleClosed`; +- `NoOverlayRoute`; +- `PathNotReady`; +- `AddressFamilyUnsupported`; +- `AddressInUse`; +- `ConnectionRefused`; +- `NetworkChanged`; +- `ResourceLimit`; +- `BufferTooSmall`; +- `Io`. + +Adapters may attach detailed messages, but consumers must not infer behavior +from strings. + +### Flow ownership + +The following SOCKS5-named types move into the data plane: + +| Current type | Target type | +| --- | --- | +| `Socks5Entry` | `FlowKey` | +| `Socks5EntryKind` | `FlowKind` | +| `Socks5EntryTable` | `FlowTable` | +| `Socks5EntryGuard` | `FlowLease` | +| `Socks5PeerPacketRoute` | `PeerPacketRoute` | + +Each resource owns explicit leases: + +- an outbound smoltcp stream owns its exact flow lease and Host port + reservation; +- a stream returned by the optional KCP backend owns its logical source-port + lease through the backend connection wrapper; +- a listener owns its wildcard listener lease; +- an accepted stream owns an independent exact flow lease; +- a UDP socket owns the destination flow leases it created; +- a direct or local stream reports its actual address and does not fabricate a + smoltcp address. + +Dropping a listener stops new accepts but does not invalidate already accepted +streams. Cancelling a connect future releases both its flow and port +reservation. + +UDP cleanup must release owned leases directly. It must not scan the complete +flow table using pointer equality. + +### Packet delivery + +Packet classification behavior must remain compatible: + +- malformed or unmatched packets pass to the next pipeline; +- exact TCP flows win over wildcard listeners; +- only supported EasyTier data packet kinds are consumed; +- modified-source packets that are not local loopback traffic pass through; +- fragmented UDP response routing remains supported; +- a packet is consumed only when the owning flow exists. + +`DataPlaneRuntime`, not the SOCKS5 Adapter, registers and owns the peer-packet +pipeline. + +### smoltcp stack generations + +The smoltcp packet bridge moves out of `Socks5ServerNet` into a +`SmoltcpPlane` implementation. + +A stack generation owns: + +- the virtual IPv4 address and prefix; +- the smoltcp `Net`; +- its flow table; +- peer-to-smoltcp and smoltcp-to-peer tasks; +- a cancellation token and closed state. + +When the virtual address or prefix changes, the old generation closes +deterministically. Existing sockets return `NetworkChanged` or `Closed` on +subsequent I/O. Old flows cannot survive into the new generation. + +The runtime starts as a lightweight idle Module with the core instance. A +runtime lease creates or retains the smoltcp generation. Data-plane sockets, +an enabled SOCKS5 Adapter, and active port forwards hold leases. + +Dropping the final lease immediately wakes the runtime supervisor. The current +possibility of waiting for a 120-second timer before releasing an unused +stack is removed. Protocol timers such as UDP expiry remain real deadlines, +not polling substitutes. + +### SOCKS5 Adapter + +The following SOCKS5 implementation remains: + +- codec and target-address parsing; +- handshake and authentication; +- command parsing; +- reply encoding; +- server framing and request lifecycle. + +`Socks5GatewayAdapter` owns: + +- the Host listener; +- accepted SOCKS5 sessions; +- a connector Adapter that calls `DataPlaneRuntime` with + `OverlayOrDirect`; +- mapping `DataPlaneError` to SOCKS5 replies; +- the SOCKS5 task group. + +The current SOCKS5 UDP behavior is preserved during this implementation. +Adding overlay-aware SOCKS5 UDP is a separate behavior change and must not be +hidden inside the mechanical refactor. + +### Port-forward Adapter + +Port-forward ownership moves out of the core data-plane implementation into a +gateway Adapter. Its current TCP and UDP route behavior remains compatible. +It uses crate-private data-plane entry points rather than reading runtime +fields directly. + +### Core instance lifecycle + +The start order is: + +```text +PeerManager +-> WrappedTransport +-> DataPlaneRuntime +-> SOCKS5 and PortForward Adapters +``` + +Stop uses the reverse order. + +The data-plane runtime is available whenever its feature is compiled. It is +not gated by whether configuration enables a SOCKS5 listener or a port-forward +rule. Gateway configuration controls the Adapters, not whether public Go or +FFI data-plane sockets are possible. + +The runtime must: + +- clean up partial start; +- reject new work after stopping begins; +- wake pending connect, bind, accept, read, write, and receive operations; +- avoid awaiting Host, PeerManager, optional KCP backend, or smoltcp work + while holding its main state lock; +- make repeated stop and resource close safe. + +## KCP backend boundary + +The existing concrete `KcpProxyService` remains in the native `easytier` +composition and continues to use `kcp-sys`. It is injected through the +existing wrapped-transport Interface. + +`DataPlaneRuntime` owns the decision to use that backend and wraps the returned +stream with the flow and source-port leases required by the gateway caller. +SOCKS5 and port forwarding do not query KCP readiness or peer policy directly. + +The FFI v2 and WASI session request `SmoltcpOnly`, so their TCP listeners +accept smoltcp streams only. The WASI composition installs no KCP backend and +the data-plane capability set contains no KCP bit. Consequently this plan does +not require: + +- moving the concrete engine into `easytier-core`; +- compiling `kcp-sys` for WASI; +- tracking KCP timers in the externally driven runtime; +- delivering KCP streams to FFI or Go listeners; +- changing current KCP hedge or cancellation behavior. + +These omissions are an explicit scope boundary, not a Go-side KCP fallback. + +## Data-plane session and operation broker + +### Instance scope + +Each core instance owns a `DataPlaneSession`: + +```text +DataPlaneSession + ResourceTable + data-plane operation metadata and quotas + foundation::OperationBroker + CompletionNotifier +``` + +The native C ABI may still require process-global opaque session handles, but +the actual resource and operation namespaces are instance scoped. + +The crate-private foundation broker is a concrete Module rather than an +Adapter trait. It owns ID allocation, terminal-operation arbitration, retained +outcomes, and completion queue transitions without depending on data-plane +types. `DataPlaneSession` composes it under the session's existing lock so +resource creation, quota settlement, and completion publication remain one +atomic transition. + +Resource handles identify: + +- TCP streams; +- TCP listeners; +- UDP sockets. + +Operation handles identify: + +- TCP connect; +- TCP bind; +- TCP accept; +- TCP read; +- TCP write; +- UDP bind; +- UDP receive; +- UDP send. + +Handle zero is invalid. IDs are monotonically allocated within the session and +are not reused while a stale reference may remain. + +### Operation state + +An operation slot progresses through: + +```text +Pending +-> Queued(result or error) +-> Drained(result or error) +-> Consumed + +Pending -> Discarding -> Discarded +Queued -> Discarded +``` + +Completion and cancellation are linearized in one state transition: + +- the first transition from `Pending` wins; +- an operation is queued at most once; +- cancelling a queued or drained operation preserves its result; +- cancelling an absent or consumed operation is idempotent for cleanup paths; +- closing a resource turns its pending operations into `HandleClosed`; +- stopping a session turns its pending operations into `InstanceStopped`. + +Free and the completion queue are linearized under the same broker lock: + +- freeing `Pending` moves it to an invisible `Discarding` tombstone until its + task acknowledges cancellation; +- freeing `Queued` atomically marks its descriptor discarded so a batch drain + skips it; +- freeing `Drained` drops the retained result; +- a discarded successful result that owns a new resource closes that resource; +- discarded operations never produce an unexplained completion. + +The completion queue contains fixed descriptors, not potentially large +payloads. A descriptor contains: + +- operation ID; +- operation kind; +- stable status code. + +The typed result remains in the operation slot. Draining a descriptor does not +consume the result. Result take is exactly once and consumes only after a +successful copy into a sufficiently large destination. + +The implementation enforces per-instance limits for: + +- open resources; +- outstanding operations; +- total retained result bytes; +- a single read allocation. + +Admission reserves one completion credit for every accepted operation. +Therefore an operation that was accepted can always enter `Queued`; completion +cannot fail because the queue became full. Once operation or result-byte +limits are reached, new submissions fail with `ResourceLimit`. + +### Completion notification + +The completion notifier sends one edge when the queue changes from empty to +non-empty. It does not run user code and does not enter a guest. + +Native FFI uses a condition variable initially. WASI uses queue-ready state +observed by the driver after a guest turn. Optional Unix `eventfd` and Windows +event handles may be added later without changing operation semantics. + +## Native FFI + +### New session Interface + +The new native Interface provides: + +- open and close an instance data-plane session; +- submit typed TCP or UDP operations; +- cancel or free an operation; +- close a resource; +- wait until any completion exists; +- batch-drain completion descriptors; +- query result size; +- take a typed result. + +`completion_wait` blocks on the complete session queue. A caller needs one +dispatcher wait rather than one loop per operation. Timeout means only that no +completion arrived during the requested wait. + +No completion callback is added initially. Calling arbitrary C or Go code from +a Tokio runtime thread would create avoidable reentrancy, lock-order, and +lifetime hazards. + +### Replacement policy + +The existing synchronous and +`start/status/wait/finish/free` data-plane functions are removed rather than +adapted. FFI v2 is the only native data-plane Interface after this change. + +This also removes the runtime-readiness loop that currently calls a manager +method in 50-millisecond chunks even though that method does not actually +wait. Compatibility for those unused symbols and their previous direct +fallback behavior is explicitly outside this plan. + +## WASI guest ABI + +### Direction + +Existing `easytier_host` imports remain guest-to-Host underlay operations. +They are not reversed or reused for public data-plane sockets. + +New guest exports provide: + +- typed TCP and UDP operation submission; +- operation cancellation and free; +- resource close; +- batch completion-descriptor take; +- result length; +- typed result take; +- data-plane ABI version and capabilities. + +`CORE_INSTANCE_CONFIG_VERSION` remains the create-envelope schema version. It +is not reused as a general WASI ABI version. + +Capabilities include at least: + +- data plane; +- TCP; +- UDP. + +The Go `coreabi` Adapter checks required exports, version, and capabilities +before creating an instance. + +### Wire and memory ownership + +The data-plane wire format has an explicit byte order and field encoding. It +does not expose Rust `repr(C)` layout or padding. + +ABI v2 uses big-endian integers and the existing 27-byte WASI socket-address +record. Its fixed records are: + +| Record | Bytes | Fields | +| --- | ---: | --- | +| Operation ID | 8 | `u64` | +| Completion | 12 | operation `u64`, kind `u16`, status `u16` | +| TCP connect/accept | 62 | resource `u64`, local address, peer address | +| TCP/UDP bind | 35 | resource `u64`, local address | +| TCP read metadata | 1 | EOF flag | +| UDP receive metadata | 28 | peer address, truncated flag | + +The v2 address decoder accepts IPv4 only even though the shared address record +reserves an IPv6 representation for a future capability revision. + +Every submit export writes its operation ID into a caller-allocated eight-byte +guest buffer and returns a stable status. `UINT64_MAX` is the no-deadline +sentinel; zero is an immediate deadline. Completion drain writes a dense array +of 12-byte descriptors. Payload result size is queried before typed result +take. + +Memory rules: + +- submission copies and validates all request and write bytes before the + export returns; +- the guest never retains a Host memory pointer; +- read payloads and typed results remain guest-owned until result take; +- result take writes only to a currently valid guest buffer; +- insufficient capacity reports `BufferTooSmall` without consuming; +- the Host copies result bytes before freeing the guest buffer; +- no guest-memory slice escapes a serialized `coreabi` call. + +### Event-driven drive contract + +WASM cannot execute while the Host is not inside the guest. The externally +driven current-thread Tokio runtime therefore remains, but the public Go +caller never drives or polls it. + +The Go driver enters the guest after: + +- a caller command; +- a real Host reactor completion; +- a tracked Tokio deadline; +- raw packet ingress; +- immediately runnable guest work, represented by a zero next deadline. + +After each bounded drive turn, the driver batch-drains data-plane completion +descriptors and dispatches them to Go waiters. + +A guest data-plane completion is created only while the guest is executing a +turn. It does not need to call back into Go from a background thread. + +The current unconditional Host-completion notification inside +`WasiInstance::drive` is removed. Only a real reactor completion causes: + +```text +NotifyCompletions +-> Drive +-> drain data-plane completions +-> NextDeadline +``` + +Host import callbacks and reactor workers may signal a Go channel but must +never re-enter the guest. + +### Shutdown + +Graceful stop: + +1. takes the session lock and linearizes `Running -> Stopping`; +2. rejects new submissions and turns every operation still pending at that + point into an `InstanceStopped` completion; +3. detaches and closes underlying resources without reclassifying those + operations as `HandleClosed`; +4. lets the Host drain terminal completions; +5. stops the data-plane runtime. + +If explicit resource close linearizes first, its pending operations complete +as `HandleClosed`. If stop linearizes first, they complete as +`InstanceStopped`. + +Forced guest drop additionally requires the Go Adapter to complete any local +waiters with `net.ErrClosed`, because guest results cannot be queried after +the module is dropped. + +An unclaimed successful completion that owns a newly created resource closes +that resource when cancelled, freed, or dropped. + +## Go host design + +### Public Interface + +The root `Instance` gains: + +```go +Dial(ctx context.Context, network, address string) (net.Conn, error) +Listen(network, address string) (net.Listener, error) +ListenPacket(network, address string) (net.PacketConn, error) +``` + +The first version supports: + +- `tcp`; +- `tcp4`; +- `udp`; +- `udp4`; +- IPv4 literals. + +It explicitly rejects: + +- `tcp6` and `udp6`; +- hostnames; +- non-overlay destinations; +- Host direct fallback. + +Existing raw packet `SendPacket` and `ReceivePacket` calls remain independent +of the socket data plane. + +No public type exposes: + +- wazero; +- guest memory; +- resource handles; +- operation IDs; +- `start`, `take`, `drive`, or completion polling. + +### Package ownership + +Suggested additions: + +```text +dataplane.go +internal/coreabi/dataplane.go +internal/coreabi/dataplane_wire.go +internal/coreabi/errors.go +internal/engine/dataplane.go +internal/engine/conn.go +internal/engine/listener.go +internal/engine/packet_conn.go +internal/engine/deadline.go +``` + +The exact internal files may be consolidated where that improves locality. + +Responsibilities: + +- the root package validates the small public Interface and delegates; +- `internal/coreabi` owns guest export lookup, memory copies, wire encoding, + status conversion, and ABI validation; +- `internal/engine` owns serialized guest execution, pending-operation + dispatch, standard Go network behavior, cancellation, deadlines, and + resource lifecycle; +- `platform`, `internal/hostabi`, and `internal/reactor` continue to own only + underlay Host capabilities. + +The data plane does not add another public bridge or expose raw guest handles. + +### Single driver + +After module creation, every guest export and guest-memory access runs on the +instance's single driver goroutine. + +Public socket calls submit one of three internal command classes: + +- submit a data-plane operation; +- cancel an operation; +- close a resource. + +The driver owns the pending-operation map. Reactor workers, public goroutines, +and timer callbacks communicate only through channels. + +The event loop is: + +```text +command, Host completion, or deadline +-> optional NotifyCompletions +-> bounded Drive turns +-> batch-drain data-plane completions +-> dispatch result channels +-> query NextDeadline +-> wait +``` + +The driver records a pending waiter before a completion can be drained. An +unknown operation completion, wrong kind, or wrong resource is an ABI protocol +error rather than something silently ignored. + +Completion draining has a fairness budget. If a batch limit is reached while +more work remains, the driver schedules an immediate self-turn rather than +waiting for an unrelated event. + +### Cancellation races + +Context or deadline cancellation sends a guest cancel command and waits for +the terminal arbitration. It does not return while leaving an unowned +operation active. + +If completion wins, its result is delivered. If cancellation wins, the +operation completes as cancelled. If a resource-creating operation completes +after its Go caller can no longer receive the result, the driver closes the +orphan resource immediately. + +Pending map entries remain until one terminal completion is processed. This +prevents late completions from becoming unexplained protocol errors. + +### TCP connection Adapter + +The guest session owns the actual stream. Go owns an opaque resource ID, +address snapshots, deadline state, and its instance reference. + +Concurrency: + +- one read and one write may proceed concurrently; +- multiple reads serialize through a read mutex; +- multiple writes serialize through a write mutex; +- close and deadline changes may run concurrently with either direction; +- close does not wait for a read or write mutex. + +Read behavior: + +- zero-length reads return immediately; +- explicit EOF maps to `io.EOF`; +- cancellation cannot leave a hidden operation consuming future bytes; +- closing the connection wakes a blocked read. + +Write behavior: + +- zero-length writes return immediately; +- the guest tracks bytes written before an error; +- Go returns `n, err` for partial progress; +- successful completion returns the full source length. + +`LocalAddr` and `RemoteAddr` use immutable snapshots returned when the stream +is established or accepted. + +### Listener Adapter + +Accept calls are serialized. Close is idempotent and wakes a blocked Accept +with `net.ErrClosed`. + +If an Accept/Close race creates a child stream after the caller can no longer +receive it, the driver closes the child immediately. + +The public session listener accepts smoltcp streams through the core listener +Interface. + +### Packet connection Adapter + +Read and write directions serialize independently. + +- UDP writes are atomic. +- A short datagram write becomes `io.ErrShortWrite`. +- `ReadFrom` returns `*net.UDPAddr`. +- A zero-length datagram is a valid datagram and is not EOF. +- Completion results identify truncation explicitly. +- closing the socket wakes blocked receive operations. + +`Dial` for UDP creates an ephemeral UDP resource with a fixed peer. Its read +side accepts only that peer; `ListenPacket` exposes the full datagram Interface. + +### Deadline Adapter + +Read and write directions each own: + +- an absolute deadline; +- a generation or changed channel; +- the current guest operation. + +Changing a deadline wakes the current waiter so it can recompute: + +- shortening affects an existing operation; +- extending does not cause the old timer to cancel the operation; +- clearing removes the timer; +- close remains independent of direction locks. + +When a deadline expires, Go submits cancellation and waits for the +cancel-versus-completion arbitration before returning +`os.ErrDeadlineExceeded`. + +### Error mapping + +Stable guest error codes map to: + +- `*net.OpError`; +- `context.Canceled`; +- `context.DeadlineExceeded`; +- `os.ErrDeadlineExceeded`; +- `net.ErrClosed`. + +Go must not classify errors by string matching. + +### Instance shutdown + +Instance close order is: + +1. reject new public operations; +2. ask the guest session to close resources and cancel operations; +3. dispatch or locally complete all pending waiters; +4. stop and drop the guest instance; +5. close the Host reactor and wazero module/runtime. + +The existing distinction between graceful `Stop` and resource-reclaiming +`Close` remains. + +## Implementation sequence + +### Phase 0: freeze semantics and define the ABI + +Repository: EasyTier. + +- Add golden tests for current route selection, Host direct calls, flow + cleanup, listener/accepted-stream lifetime, stack lease release, stop, IP + update, local addresses, SOCKS5 UDP, and port-forward behavior. +- Define the data-plane ABI version, capabilities, wire format, errors, and + resource limits. + +### Phase 1: move shared data-plane ownership + +Repository: EasyTier. + +- Move and rename the flow table and peer-packet classifier. +- Move the smoltcp packet bridge out of `Socks5ServerNet`. +- Split TCP and UDP resource implementations. +- Preserve all behavior and existing tests. + +### Phase 2: introduce `DataPlaneRuntime` + +Repository: EasyTier. + +- Add stack generations and runtime leases. +- Make shutdown and final-lease release event driven. +- Add typed errors and one absolute deadline. +- Make `CoreInstance` own and start the runtime independently of configured + gateway Adapters. + +### Phase 3: route policy and gateway Adapters + +Repository: EasyTier. + +- Add `OverlayOnly` and `OverlayOrDirect`. +- Move route selection into the data plane. +- Make public data-plane calls overlay-only. +- Adapt SOCKS5 TCP and port forwarding without changing their behavior. +- Add guard tests proving overlay-only never connects an unrelated Host + destination. + +### Phase 4: optional native KCP backend integration + +Repository: EasyTier. + +- Move KCP selection and source-connection ownership into + `DataPlaneRuntime`. +- Keep the existing concrete backend injected by native `easytier`. +- Make public FFI/WASI sessions explicitly request `SmoltcpOnly`. +- Add selected-path diagnostics and route-selection tests without changing + `kcp-sys` or the wrapped-transport destination. + +### Phase 5: session and operation broker + +Repository: EasyTier. + +- Add the resource table, operation slots, completion queue, and notifier. +- Implement operation limits, cancellation, close, stop, result ownership, and + batch completion drain. +- Add deterministic race and lost-wakeup tests. + +### Phase 6: native FFI + +Repository: EasyTier. + +- Add the new session Interface. +- Remove the existing synchronous and per-operation asynchronous data-plane + Interfaces. +- Remove runtime-readiness polling. +- Update examples and conformance tests. + +### Phase 7: WASI ABI + +Repository: EasyTier. + +- Add guest data-plane exports and canonical wire encoding. +- Add ABI version and capabilities. +- Remove unconditional duplicate Host completion notification. +- Enable smoltcp in the artifact build profile. +- Add WASI target and memory-ownership tests. + +### Phase 8: Go `coreabi` and engine + +Repository: `easytier-go-host`. + +- Add typed data-plane guest calls and wire codecs. +- Extend the single driver with submit, cancel, close, and completion drain. +- Add pending-operation, orphan-resource, shutdown, and ABI-protocol handling. +- Test with a deterministic fake guest before adding public sockets. + +### Phase 9: Go standard network Adapters and artifact + +Repository: `easytier-go-host`. + +- Add `Dial`, `Listen`, and `ListenPacket`. +- Implement TCP, UDP, deadlines, cancellation, close, and error mapping. +- Build the final WASM from the exact final EasyTier commit on + `codex/cfg-refactor-pr`. +- Update embedded artifact provenance. +- Run real two-instance TCP and UDP integration tests. + +Each phase ends in a reviewable commit. Commit messages use a 72-column text +width. The complete task receives one final review across both repositories. +The operation-broker commit may receive an additional high-risk incremental +review because it contains concurrency logic. + +## Validation plan + +### EasyTier core + +Required focused tests: + +- complete route-policy truth table; +- no Host direct call under `OverlayOnly`; +- optional native KCP backend selection remains below `DataPlaneRuntime`; +- public session calls remain smoltcp-only even when KCP is available; +- port and flow release on cancel and drop; +- listener and accepted-stream independent lifetime; +- UDP multi-destination flow cleanup; +- deterministic close on stack-generation change; +- stop while connecting, accepting, reading, writing, or receiving; +- SOCKS5 TCP wire compatibility; +- broker completion/cancel and completion/close races; +- exactly-once completion and result take; +- buffer-too-small without consumption; +- session isolation and resource limits; +- blocking completion wait without lost wakeups. + +Feature and target profiles include: + +```text +cargo check -p easytier-core --no-default-features +cargo check -p easytier-core --no-default-features \ + --features proxy-smoltcp-stack +cargo test -p easytier-core --no-default-features \ + --features proxy-smoltcp-stack,test-utils +cargo test -p easytier-core \ + --features proxy-smoltcp-stack,test-utils +cargo build -p easytier-core --release --target wasm32-wasip1 \ + --features proxy-smoltcp-stack +``` + +Exact commands may change with the final feature names. + +Docker integration in the shared `rust` container covers: + +- two-node TCP and UDP data-plane traffic; +- three-node relay TCP; +- SOCKS5 virtual and direct targets; +- stop during connection and accept; +- runtime virtual IPv4 change. + +### Native FFI + +- one wait receives completions from multiple operation kinds; +- timeout does not spin; +- cross-thread wait and take; +- instance deletion wakes waiters; +- stress with many outstanding operations; +- no leaked result buffers, operations, or resources. + +### Go + +Deterministic engine tests cover: + +- all guest calls remain serialized; +- Host completion ordering is notify before drive; +- synchronous and batch completion delivery; +- fairness self-wake; +- complete/cancel race; +- orphan resource close; +- instance shutdown with pending work; +- unknown or mismatched completion as an ABI error. + +Standard network tests cover: + +- TCP full duplex and EOF; +- partial write with progress; +- one concurrent read and write; +- listener port zero and Accept/Close races; +- UDP address and truncation semantics; +- connected UDP peer filtering; +- deadline shortening, extension, clearing, and expiry; +- Close unblocking Read, Accept, and Receive; +- `errors.Is` mappings; +- race-detector execution. + +Real embedded-WASM tests cover: + +- overlay TCP dial and listen; +- overlay UDP packet connection; +- context cancellation and deadlines; +- repeated close; +- a non-overlay dial failure with a recording Host socket factory proving no + data-plane direct connection occurred; +- idle operation with no drive turns other than genuine protocol deadlines. + +Final Go validation includes: + +```text +go test -count=1 ./... +go test -race ./... +go vet ./... +``` + +## Definition of done + +The implementation is complete when: + +- EasyTier owns all data-plane route, smoltcp, flow, and resource policy; +- optional native KCP selection and source-stream ownership are behind + `DataPlaneRuntime`; +- SOCKS5 and port forwarding use Data Plane through explicit Adapters; +- public Go and new FFI data-plane calls are overlay-only; +- a Go data-plane TCP listener accepts smoltcp connections; +- native FFI can wait for any instance completion without per-operation + polling; +- the Go caller sees only standard `net` interfaces; +- the Go driver is event driven and never re-entered from a Host callback; +- cancellation, deadline, close, and instance stop produce exactly one terminal + operation outcome; +- the WASM artifact includes the smoltcp data plane and does not advertise + KCP; +- the unused legacy data-plane FFI has been removed; +- all focused, Docker, race, ABI, and embedded-WASM tests pass; +- artifact provenance identifies the exact final EasyTier commit; +- no consumer-specific code exists in either repository. diff --git a/easytier-contrib/easytier-android-jni/Cargo.toml b/easytier-contrib/easytier-android-jni/Cargo.toml index 73212296..f810eaac 100644 --- a/easytier-contrib/easytier-android-jni/Cargo.toml +++ b/easytier-contrib/easytier-android-jni/Cargo.toml @@ -14,4 +14,4 @@ android_logger = "0.13" serde = { version = "1.0.220", features = ["derive"] } serde_json = "1.0" easytier = { path = "../../easytier" } -easytier-ffi = { path = "../easytier-ffi", default-features = false, features = ["ffi-dataplane"] } +easytier-ffi = { path = "../easytier-ffi", default-features = false } diff --git a/easytier-contrib/easytier-android-jni/exports.map b/easytier-contrib/easytier-android-jni/exports.map index 9c047617..3f2b15f1 100644 --- a/easytier-contrib/easytier-android-jni/exports.map +++ b/easytier-contrib/easytier-android-jni/exports.map @@ -1,7 +1,6 @@ { global: Java_com_easytier_jni_EasyTierJNI_*; - Java_com_easytier_jni_EasyTierDataPlaneJNI_*; local: *; }; diff --git a/easytier-contrib/easytier-android-jni/kotlin/com/easytier/jni/EasyTierDataPlaneJNI.kt b/easytier-contrib/easytier-android-jni/kotlin/com/easytier/jni/EasyTierDataPlaneJNI.kt deleted file mode 100644 index 8293c6e7..00000000 --- a/easytier-contrib/easytier-android-jni/kotlin/com/easytier/jni/EasyTierDataPlaneJNI.kt +++ /dev/null @@ -1,451 +0,0 @@ -package com.easytier.jni - -import kotlinx.coroutines.CancellationException -import kotlinx.coroutines.Dispatchers -import kotlinx.coroutines.currentCoroutineContext -import kotlinx.coroutines.ensureActive -import kotlinx.coroutines.withContext - -/** - * EasyTier data-plane API for Android. - * - * Dataplane APIs do not create or start an EasyTier instance by themselves. - * Start an instance with [EasyTierJNI.runNetworkInstance] first, then pass the - * same `instanceName` to [EasyTierDataPlane.tcpConnect], - * [EasyTierDataPlane.tcpBind], or [EasyTierDataPlane.udpBind]. If that instance - * is not running, the native start call fails and the coroutine wrapper throws - * the last EasyTier FFI error. - * - * Typical setup: - * ``` - * val instanceName = "android-dataplane-demo" - * val config = """ - * instance_name = "$instanceName" - * ipv4 = "10.144.0.1" - * listeners = ["tcp://0.0.0.0:11010"] - * - * [network_identity] - * network_name = "android-dataplane-demo" - * network_secret = "replace-with-a-real-secret" - * - * [[peer]] - * uri = "tcp://peer.example.com:11010" - * - * [flags] - * no_tun = true - * bind_device = false - * """.trimIndent() - * - * EasyTierJNI.runNetworkInstance(config) - * ``` - * - * After the instance is running, most callers should use [EasyTierDataPlane] - * and the socket/stream classes below. [EasyTierDataPlaneJNI] is the low-level - * native op-handle ABI used by the coroutine wrappers. - * - * TCP client usage: - * ``` - * val stream = EasyTierDataPlane.tcpConnect(instanceName, "10.144.0.2", 8080, 5_000) - * try { - * stream.write("ping".toByteArray(), 5_000) - * val reply = stream.read(4096, 5_000) - * } finally { - * stream.close() - * } - * ``` - * - * TCP server usage: - * ``` - * val listener = EasyTierDataPlane.tcpBind(instanceName, 8080, 5_000) - * try { - * val stream = listener.accept(30_000) - * try { - * stream.write(stream.read(4096, 5_000), 5_000) - * } finally { - * stream.close() - * } - * } finally { - * listener.close() - * } - * ``` - * - * UDP usage: - * ``` - * val socket = EasyTierDataPlane.udpBind(instanceName, 0, 5_000) - * try { - * socket.sendTo("10.144.0.2", 9000, "ping".toByteArray(), 5_000) - * val packet = socket.recvFrom(4096, 5_000) - * } finally { - * socket.close() - * } - * ``` - * - * Operation model: - * - Each suspend function starts one native async op, waits on Dispatchers.IO, - * then consumes the op with the matching finish call. - * - Coroutine cancellation cancels and frees the native op. - * - Returned stream/listener/socket handles must be closed by the caller. - * - Input ByteArray data is copied by the native start call; output data is - * copied into Kotlin ByteArray before the native buffer is freed. - */ - -/** Data-plane IPv4/port pair returned by EasyTier FFI. */ -data class DataPlaneSocketAddress(val ip: String, val port: Int) - -/** Result of a completed TCP connect op. */ -data class DataPlaneTcpConnectResult(val handle: Long, val localAddress: DataPlaneSocketAddress) - -/** Result of a completed TCP bind op. */ -data class DataPlaneTcpBindResult(val handle: Long, val localAddress: DataPlaneSocketAddress) - -/** Result of a completed TCP accept op. */ -data class DataPlaneTcpAcceptResult( - val handle: Long, - val localAddress: DataPlaneSocketAddress, - val peerAddress: DataPlaneSocketAddress -) - -/** Result of a completed TCP read op. */ -data class DataPlaneTcpReadResult(val data: ByteArray) - -/** Result of a completed UDP bind op. */ -data class DataPlaneUdpBindResult(val handle: Long, val localAddress: DataPlaneSocketAddress) - -/** Result of a completed UDP recv_from op. */ -data class DataPlaneUdpRecvResult( - val data: ByteArray, - val peerAddress: DataPlaneSocketAddress -) - -/** TCP data-plane stream handle. Call [close] when the stream is no longer needed. */ -class DataPlaneTcpStream( - val handle: Long, - val localAddress: DataPlaneSocketAddress? = null, - val peerAddress: DataPlaneSocketAddress? = null -) { - /** Read up to [maxLength] bytes, waiting at most [timeoutMs] in native code. */ - suspend fun read(maxLength: Int, timeoutMs: Long): ByteArray = - EasyTierDataPlane.tcpRead(this, maxLength, timeoutMs) - - /** Write [data], waiting at most [timeoutMs] in native code. */ - suspend fun write(data: ByteArray, timeoutMs: Long): Int = - EasyTierDataPlane.tcpWrite(this, data, timeoutMs) - - /** Close the native TCP stream handle. */ - fun close(): Int = EasyTierDataPlaneJNI.dataPlaneTcpClose(handle) -} - -/** TCP data-plane listener handle. Call [close] when the listener is no longer needed. */ -class DataPlaneTcpListener(val handle: Long, val localAddress: DataPlaneSocketAddress) { - /** Accept one TCP data-plane stream. */ - suspend fun accept(timeoutMs: Long): DataPlaneTcpStream = - EasyTierDataPlane.tcpAccept(this, timeoutMs) - - /** Close the native TCP listener handle. */ - fun close(): Int = EasyTierDataPlaneJNI.dataPlaneTcpListenerClose(handle) -} - -/** UDP data-plane socket handle. Call [close] when the socket is no longer needed. */ -class DataPlaneUdpSocket(val handle: Long, val localAddress: DataPlaneSocketAddress) { - /** Send one UDP datagram to [dstIp]:[dstPort]. */ - suspend fun sendTo( - dstIp: String, - dstPort: Int, - data: ByteArray, - timeoutMs: Long - ): Int = EasyTierDataPlane.udpSendTo(this, dstIp, dstPort, data, timeoutMs) - - /** Receive one UDP datagram and its peer address. */ - suspend fun recvFrom(maxLength: Int, timeoutMs: Long): DataPlaneUdpRecvResult = - EasyTierDataPlane.udpRecvFrom(this, maxLength, timeoutMs) - - /** Close the native UDP socket handle. */ - fun close(): Int = EasyTierDataPlaneJNI.dataPlaneUdpClose(handle) -} - -/** - * Low-level native data-plane JNI entry points. - * - * These functions mirror the Rust FFI op-handle ABI directly. They are exposed - * for completeness, but most Android callers should use [EasyTierDataPlane] - * instead so coroutine cancellation and op cleanup are handled consistently. - */ -object EasyTierDataPlaneJNI { - init { - System.loadLibrary("easytier_android_jni") - } - - @JvmStatic external fun dataPlaneAsyncOpStatus(handle: Long): Int - - @JvmStatic external fun dataPlaneAsyncOpWait(handle: Long, timeoutMs: Long): Int - - @JvmStatic external fun dataPlaneAsyncOpCancel(handle: Long): Int - - @JvmStatic external fun dataPlaneAsyncOpFree(handle: Long): Int - - @JvmStatic - external fun dataPlaneTcpConnectStart( - instanceName: String, - dstIp: String, - dstPort: Int, - timeoutMs: Long - ): Long - - @JvmStatic external fun dataPlaneTcpConnectFinish(op: Long): DataPlaneTcpConnectResult? - - @JvmStatic - external fun dataPlaneTcpBindStart( - instanceName: String, - localPort: Int, - timeoutMs: Long - ): Long - - @JvmStatic external fun dataPlaneTcpBindFinish(op: Long): DataPlaneTcpBindResult? - - @JvmStatic external fun dataPlaneTcpAcceptStart(handle: Long, timeoutMs: Long): Long - - @JvmStatic external fun dataPlaneTcpAcceptFinish(op: Long): DataPlaneTcpAcceptResult? - - @JvmStatic external fun dataPlaneTcpReadStart(handle: Long, maxLength: Int, timeoutMs: Long): Long - - @JvmStatic external fun dataPlaneTcpReadFinish(op: Long): DataPlaneTcpReadResult? - - @JvmStatic external fun dataPlaneTcpWriteStart(handle: Long, data: ByteArray, timeoutMs: Long): Long - - @JvmStatic external fun dataPlaneTcpWriteFinish(op: Long): Int - - @JvmStatic - external fun dataPlaneUdpBindStart( - instanceName: String, - localPort: Int, - timeoutMs: Long - ): Long - - @JvmStatic external fun dataPlaneUdpBindFinish(op: Long): DataPlaneUdpBindResult? - - @JvmStatic - external fun dataPlaneUdpSendToStart( - handle: Long, - dstIp: String, - dstPort: Int, - data: ByteArray, - timeoutMs: Long - ): Long - - @JvmStatic external fun dataPlaneUdpSendToFinish(op: Long): Int - - @JvmStatic external fun dataPlaneUdpRecvFromStart(handle: Long, maxLength: Int, timeoutMs: Long): Long - - @JvmStatic external fun dataPlaneUdpRecvFromFinish(op: Long): DataPlaneUdpRecvResult? - - @JvmStatic external fun dataPlaneTcpClose(handle: Long): Int - - @JvmStatic external fun dataPlaneTcpListenerClose(handle: Long): Int - - @JvmStatic external fun dataPlaneUdpClose(handle: Long): Int -} - -/** Coroutine-friendly Android data-plane API. */ -object EasyTierDataPlane { - private const val DATA_PLANE_OP_PENDING = 0 - private const val DATA_PLANE_OP_READY = 1 - private const val DATA_PLANE_OP_FAILED = -1 - private const val DATA_PLANE_OP_INVALID = -2 - private const val DATA_PLANE_WAIT_SLICE_MS = 50L - - /** Connect to a TCP endpoint through the named EasyTier instance. */ - @JvmStatic - suspend fun tcpConnect( - instanceName: String, - dstIp: String, - dstPort: Int, - timeoutMs: Long - ): DataPlaneTcpStream { - val op = - requireOp( - EasyTierDataPlaneJNI.dataPlaneTcpConnectStart( - instanceName, - dstIp, - dstPort, - timeoutMs - ) - ) - val result = awaitOp(op) { - EasyTierDataPlaneJNI.dataPlaneTcpConnectFinish(it) ?: throw lastDataPlaneException() - } - return DataPlaneTcpStream(result.handle, result.localAddress) - } - - /** Bind a TCP data-plane listener on [localPort]. Port 0 asks EasyTier to allocate one. */ - @JvmStatic - suspend fun tcpBind( - instanceName: String, - localPort: Int, - timeoutMs: Long - ): DataPlaneTcpListener { - val op = - requireOp( - EasyTierDataPlaneJNI.dataPlaneTcpBindStart( - instanceName, - localPort, - timeoutMs - ) - ) - val result = awaitOp(op) { - EasyTierDataPlaneJNI.dataPlaneTcpBindFinish(it) ?: throw lastDataPlaneException() - } - return DataPlaneTcpListener(result.handle, result.localAddress) - } - - /** Accept one TCP stream from [listener]. */ - @JvmStatic - suspend fun tcpAccept(listener: DataPlaneTcpListener, timeoutMs: Long): DataPlaneTcpStream { - val op = - requireOp( - EasyTierDataPlaneJNI.dataPlaneTcpAcceptStart(listener.handle, timeoutMs) - ) - val result = awaitOp(op) { - EasyTierDataPlaneJNI.dataPlaneTcpAcceptFinish(it) ?: throw lastDataPlaneException() - } - return DataPlaneTcpStream(result.handle, result.localAddress, result.peerAddress) - } - - /** Read up to [maxLength] bytes from [stream]. */ - @JvmStatic - suspend fun tcpRead( - stream: DataPlaneTcpStream, - maxLength: Int, - timeoutMs: Long - ): ByteArray { - val op = - requireOp( - EasyTierDataPlaneJNI.dataPlaneTcpReadStart( - stream.handle, - maxLength, - timeoutMs - ) - ) - return awaitOp(op) { - EasyTierDataPlaneJNI.dataPlaneTcpReadFinish(it)?.data - ?: throw lastDataPlaneException() - } - } - - /** Write [data] to [stream]. */ - @JvmStatic - suspend fun tcpWrite(stream: DataPlaneTcpStream, data: ByteArray, timeoutMs: Long): Int { - val op = - requireOp( - EasyTierDataPlaneJNI.dataPlaneTcpWriteStart( - stream.handle, - data, - timeoutMs - ) - ) - return awaitOp(op) { EasyTierDataPlaneJNI.dataPlaneTcpWriteFinish(it) } - } - - /** Bind a UDP data-plane socket on [localPort]. Port 0 asks EasyTier to allocate one. */ - @JvmStatic - suspend fun udpBind( - instanceName: String, - localPort: Int, - timeoutMs: Long - ): DataPlaneUdpSocket { - val op = - requireOp( - EasyTierDataPlaneJNI.dataPlaneUdpBindStart( - instanceName, - localPort, - timeoutMs - ) - ) - val result = awaitOp(op) { - EasyTierDataPlaneJNI.dataPlaneUdpBindFinish(it) ?: throw lastDataPlaneException() - } - return DataPlaneUdpSocket(result.handle, result.localAddress) - } - - /** Send one UDP datagram through [socket]. */ - @JvmStatic - suspend fun udpSendTo( - socket: DataPlaneUdpSocket, - dstIp: String, - dstPort: Int, - data: ByteArray, - timeoutMs: Long - ): Int { - val op = - requireOp( - EasyTierDataPlaneJNI.dataPlaneUdpSendToStart( - socket.handle, - dstIp, - dstPort, - data, - timeoutMs - ) - ) - return awaitOp(op) { EasyTierDataPlaneJNI.dataPlaneUdpSendToFinish(it) } - } - - /** Receive one UDP datagram through [socket]. */ - @JvmStatic - suspend fun udpRecvFrom( - socket: DataPlaneUdpSocket, - maxLength: Int, - timeoutMs: Long - ): DataPlaneUdpRecvResult { - val op = - requireOp( - EasyTierDataPlaneJNI.dataPlaneUdpRecvFromStart( - socket.handle, - maxLength, - timeoutMs - ) - ) - return awaitOp(op) { - EasyTierDataPlaneJNI.dataPlaneUdpRecvFromFinish(it) ?: throw lastDataPlaneException() - } - } - - private fun requireOp(op: Long): Long { - if (op == 0L) { - throw lastDataPlaneException() - } - return op - } - - private suspend fun awaitOp(op: Long, finish: (Long) -> T): T = - withContext(Dispatchers.IO) { - var consumed = false - try { - awaitReady(op) - val result = finish(op) - consumed = true - result - } catch (e: CancellationException) { - EasyTierDataPlaneJNI.dataPlaneAsyncOpCancel(op) - throw e - } finally { - if (!consumed) { - EasyTierDataPlaneJNI.dataPlaneAsyncOpFree(op) - } - } - } - - private suspend fun awaitReady(op: Long) { - while (true) { - currentCoroutineContext().ensureActive() - when (EasyTierDataPlaneJNI.dataPlaneAsyncOpWait(op, DATA_PLANE_WAIT_SLICE_MS)) { - DATA_PLANE_OP_READY, DATA_PLANE_OP_FAILED -> return - DATA_PLANE_OP_PENDING -> Unit - DATA_PLANE_OP_INVALID -> throw RuntimeException("Data-plane async operation is invalid") - else -> throw RuntimeException("Unknown data-plane async operation status") - } - } - } - - private fun lastDataPlaneException(): RuntimeException { - return RuntimeException(EasyTierJNI.getLastError() ?: "EasyTier data-plane call failed") - } -} diff --git a/easytier-contrib/easytier-android-jni/src/data_plane_api.rs b/easytier-contrib/easytier-android-jni/src/data_plane_api.rs deleted file mode 100644 index 574054be..00000000 --- a/easytier-contrib/easytier-android-jni/src/data_plane_api.rs +++ /dev/null @@ -1,673 +0,0 @@ -use std::{ - ffi::{CStr, c_char}, - ptr, -}; - -use easytier_ffi::{ - data_plane_async_op_cancel, data_plane_async_op_free, data_plane_async_op_status, - data_plane_async_op_wait, data_plane_free_bytes, data_plane_tcp_accept_finish, - data_plane_tcp_accept_start, data_plane_tcp_bind_finish, data_plane_tcp_bind_start, - data_plane_tcp_close, data_plane_tcp_connect_finish, data_plane_tcp_connect_start, - data_plane_tcp_listener_close, data_plane_tcp_read_finish, data_plane_tcp_read_start, - data_plane_tcp_write_finish, data_plane_tcp_write_start, data_plane_udp_bind_finish, - data_plane_udp_bind_start, data_plane_udp_close, data_plane_udp_recv_from_finish, - data_plane_udp_recv_from_start, data_plane_udp_send_to_finish, data_plane_udp_send_to_start, - free_string, -}; -use jni::{ - JNIEnv, - objects::{JByteArray, JClass, JObject, JString, JValue}, - sys::{jint, jlong, jobject}, -}; - -use crate::{ - error::{get_last_error, throw_exception}, - strings::jstring_to_cstring, -}; - -const SOCKET_ADDR_CLASS: &str = "com/easytier/jni/DataPlaneSocketAddress"; -const TCP_CONNECT_RESULT_CLASS: &str = "com/easytier/jni/DataPlaneTcpConnectResult"; -const TCP_BIND_RESULT_CLASS: &str = "com/easytier/jni/DataPlaneTcpBindResult"; -const TCP_ACCEPT_RESULT_CLASS: &str = "com/easytier/jni/DataPlaneTcpAcceptResult"; -const TCP_READ_RESULT_CLASS: &str = "com/easytier/jni/DataPlaneTcpReadResult"; -const UDP_BIND_RESULT_CLASS: &str = "com/easytier/jni/DataPlaneUdpBindResult"; -const UDP_RECV_RESULT_CLASS: &str = "com/easytier/jni/DataPlaneUdpRecvResult"; - -fn timeout_from_jlong(timeout_ms: jlong) -> u64 { - timeout_ms.max(0) as u64 -} - -fn port_from_jint(env: &mut JNIEnv, value: jint, name: &str) -> Option { - match u16::try_from(value) { - Ok(port) => Some(port), - Err(_) => { - throw_exception(env, &format!("Invalid {}: {}", name, value)); - None - } - } -} - -fn len_from_jint(env: &mut JNIEnv, value: jint, name: &str) -> Option { - match u32::try_from(value) { - Ok(len) => Some(len), - Err(_) => { - throw_exception(env, &format!("Invalid {}: {}", name, value)); - None - } - } -} - -fn throw_last(env: &mut JNIEnv) { - let message = get_last_error().unwrap_or_else(|| "EasyTier data-plane call failed".to_string()); - throw_exception(env, &message); -} - -unsafe fn take_ffi_string(ptr: *const c_char) -> String { - if ptr.is_null() { - return String::new(); - } - let value = unsafe { CStr::from_ptr(ptr) } - .to_string_lossy() - .into_owned(); - free_string(ptr); - value -} - -fn new_socket_addr<'local>( - env: &mut JNIEnv<'local>, - ip: String, - port: u16, -) -> Option> { - let class = match env.find_class(SOCKET_ADDR_CLASS) { - Ok(class) => class, - Err(err) => { - throw_exception( - env, - &format!("Failed to find socket address class: {:?}", err), - ); - return None; - } - }; - let ip = match env.new_string(ip) { - Ok(ip) => ip, - Err(err) => { - throw_exception(env, &format!("Failed to create IP string: {:?}", err)); - return None; - } - }; - match env.new_object( - class, - "(Ljava/lang/String;I)V", - &[JValue::Object(&ip), JValue::Int(port as jint)], - ) { - Ok(addr) => Some(addr), - Err(err) => { - throw_exception(env, &format!("Failed to create socket address: {:?}", err)); - None - } - } -} - -fn new_handle_addr_result( - env: &mut JNIEnv, - class_name: &str, - handle: u64, - ip: String, - port: u16, -) -> jobject { - let Some(addr) = new_socket_addr(env, ip, port) else { - return ptr::null_mut(); - }; - let class = match env.find_class(class_name) { - Ok(class) => class, - Err(err) => { - throw_exception(env, &format!("Failed to find result class: {:?}", err)); - return ptr::null_mut(); - } - }; - let sig = format!("(JL{};)V", SOCKET_ADDR_CLASS); - match env.new_object( - class, - sig.as_str(), - &[JValue::Long(handle as jlong), JValue::Object(&addr)], - ) { - Ok(result) => result.into_raw(), - Err(err) => { - throw_exception(env, &format!("Failed to create result object: {:?}", err)); - ptr::null_mut() - } - } -} - -fn close_tcp_stream_on_null(result: jobject, handle: u64) -> jobject { - if result.is_null() { - let _ = data_plane_tcp_close(handle); - } - result -} - -fn close_tcp_listener_on_null(result: jobject, handle: u64) -> jobject { - if result.is_null() { - let _ = data_plane_tcp_listener_close(handle); - } - result -} - -fn close_udp_socket_on_null(result: jobject, handle: u64) -> jobject { - if result.is_null() { - let _ = data_plane_udp_close(handle); - } - result -} - -fn read_owned_bytes(ptr: *const u8, len: u32) -> Vec { - if ptr.is_null() || len == 0 { - return Vec::new(); - } - let bytes = unsafe { std::slice::from_raw_parts(ptr, len as usize) }.to_vec(); - data_plane_free_bytes(ptr, len); - bytes -} - -pub(crate) fn async_op_status_jni(_env: JNIEnv, _class: JClass, handle: jlong) -> jint { - data_plane_async_op_status(handle as u64) -} - -pub(crate) fn async_op_wait_jni( - _env: JNIEnv, - _class: JClass, - handle: jlong, - timeout_ms: jlong, -) -> jint { - data_plane_async_op_wait(handle as u64, timeout_ms.max(0) as u64) -} - -pub(crate) fn async_op_cancel_jni(_env: JNIEnv, _class: JClass, handle: jlong) -> jint { - data_plane_async_op_cancel(handle as u64) -} - -pub(crate) fn async_op_free_jni(_env: JNIEnv, _class: JClass, handle: jlong) -> jint { - data_plane_async_op_free(handle as u64) -} - -pub(crate) fn tcp_connect_start_jni( - mut env: JNIEnv, - _class: JClass, - inst_name: JString, - dst_ip: JString, - dst_port: jint, - timeout_ms: jlong, -) -> jlong { - let inst_name = match jstring_to_cstring(&mut env, &inst_name) { - Ok(value) => value, - Err(err) => { - throw_exception(&mut env, &format!("Invalid instance name: {}", err)); - return 0; - } - }; - let dst_ip = match jstring_to_cstring(&mut env, &dst_ip) { - Ok(value) => value, - Err(err) => { - throw_exception(&mut env, &format!("Invalid destination IP: {}", err)); - return 0; - } - }; - let Some(dst_port) = port_from_jint(&mut env, dst_port, "destination port") else { - return 0; - }; - let op = unsafe { - data_plane_tcp_connect_start( - inst_name.as_ptr(), - dst_ip.as_ptr(), - dst_port, - timeout_ms.max(0) as u64, - ) - }; - if op == 0 { - throw_last(&mut env); - } - op as jlong -} - -pub(crate) fn tcp_connect_finish_jni(mut env: JNIEnv, _class: JClass, op: jlong) -> jobject { - let mut ip: *const c_char = ptr::null(); - let mut port = 0u16; - let handle = unsafe { data_plane_tcp_connect_finish(op as u64, &mut ip, &mut port) }; - if handle == 0 { - throw_last(&mut env); - return ptr::null_mut(); - } - close_tcp_stream_on_null( - new_handle_addr_result( - &mut env, - TCP_CONNECT_RESULT_CLASS, - handle, - unsafe { take_ffi_string(ip) }, - port, - ), - handle, - ) -} - -pub(crate) fn tcp_bind_start_jni( - mut env: JNIEnv, - _class: JClass, - inst_name: JString, - local_port: jint, - timeout_ms: jlong, -) -> jlong { - let inst_name = match jstring_to_cstring(&mut env, &inst_name) { - Ok(value) => value, - Err(err) => { - throw_exception(&mut env, &format!("Invalid instance name: {}", err)); - return 0; - } - }; - let Some(local_port) = port_from_jint(&mut env, local_port, "local port") else { - return 0; - }; - let op = unsafe { - data_plane_tcp_bind_start( - inst_name.as_ptr(), - local_port, - timeout_from_jlong(timeout_ms), - ) - }; - if op == 0 { - throw_last(&mut env); - } - op as jlong -} - -pub(crate) fn tcp_bind_finish_jni(mut env: JNIEnv, _class: JClass, op: jlong) -> jobject { - let mut ip: *const c_char = ptr::null(); - let mut port = 0u16; - let handle = unsafe { data_plane_tcp_bind_finish(op as u64, &mut ip, &mut port) }; - if handle == 0 { - throw_last(&mut env); - return ptr::null_mut(); - } - close_tcp_listener_on_null( - new_handle_addr_result( - &mut env, - TCP_BIND_RESULT_CLASS, - handle, - unsafe { take_ffi_string(ip) }, - port, - ), - handle, - ) -} - -pub(crate) fn tcp_accept_start_jni( - mut env: JNIEnv, - _class: JClass, - handle: jlong, - timeout_ms: jlong, -) -> jlong { - let op = unsafe { data_plane_tcp_accept_start(handle as u64, timeout_from_jlong(timeout_ms)) }; - if op == 0 { - throw_last(&mut env); - } - op as jlong -} - -pub(crate) fn tcp_accept_finish_jni(mut env: JNIEnv, _class: JClass, op: jlong) -> jobject { - let mut local_ip: *const c_char = ptr::null(); - let mut local_port = 0u16; - let mut peer_ip: *const c_char = ptr::null(); - let mut peer_port = 0u16; - let handle = unsafe { - data_plane_tcp_accept_finish( - op as u64, - &mut local_ip, - &mut local_port, - &mut peer_ip, - &mut peer_port, - ) - }; - if handle == 0 { - throw_last(&mut env); - return ptr::null_mut(); - } - let Some(local_addr) = - new_socket_addr(&mut env, unsafe { take_ffi_string(local_ip) }, local_port) - else { - free_string(peer_ip); - let _ = data_plane_tcp_close(handle); - return ptr::null_mut(); - }; - let Some(peer_addr) = new_socket_addr(&mut env, unsafe { take_ffi_string(peer_ip) }, peer_port) - else { - let _ = data_plane_tcp_close(handle); - return ptr::null_mut(); - }; - let class = match env.find_class(TCP_ACCEPT_RESULT_CLASS) { - Ok(class) => class, - Err(err) => { - throw_exception( - &mut env, - &format!("Failed to find accept result class: {:?}", err), - ); - let _ = data_plane_tcp_close(handle); - return ptr::null_mut(); - } - }; - let sig = format!("(JL{};L{};)V", SOCKET_ADDR_CLASS, SOCKET_ADDR_CLASS); - let result = match env.new_object( - class, - sig.as_str(), - &[ - JValue::Long(handle as jlong), - JValue::Object(&local_addr), - JValue::Object(&peer_addr), - ], - ) { - Ok(result) => result.into_raw(), - Err(err) => { - throw_exception( - &mut env, - &format!("Failed to create accept result: {:?}", err), - ); - ptr::null_mut() - } - }; - close_tcp_stream_on_null(result, handle) -} - -pub(crate) fn tcp_read_start_jni( - mut env: JNIEnv, - _class: JClass, - handle: jlong, - max_len: jint, - timeout_ms: jlong, -) -> jlong { - let Some(max_len) = len_from_jint(&mut env, max_len, "max length") else { - return 0; - }; - let op = unsafe { - data_plane_tcp_read_start(handle as u64, max_len, timeout_from_jlong(timeout_ms)) - }; - if op == 0 { - throw_last(&mut env); - } - op as jlong -} - -pub(crate) fn tcp_read_finish_jni(mut env: JNIEnv, _class: JClass, op: jlong) -> jobject { - let mut ptr: *const u8 = ptr::null(); - let mut len = 0u32; - let ret = unsafe { data_plane_tcp_read_finish(op as u64, &mut ptr, &mut len) }; - if ret < 0 { - throw_last(&mut env); - return ptr::null_mut(); - } - let bytes = read_owned_bytes(ptr, len); - let array = match env.byte_array_from_slice(&bytes) { - Ok(array) => array, - Err(err) => { - throw_exception(&mut env, &format!("Failed to create byte array: {:?}", err)); - return ptr::null_mut(); - } - }; - let class = match env.find_class(TCP_READ_RESULT_CLASS) { - Ok(class) => class, - Err(err) => { - throw_exception( - &mut env, - &format!("Failed to find read result class: {:?}", err), - ); - return ptr::null_mut(); - } - }; - match env.new_object(class, "([B)V", &[JValue::Object(&array)]) { - Ok(result) => result.into_raw(), - Err(err) => { - throw_exception( - &mut env, - &format!("Failed to create read result: {:?}", err), - ); - ptr::null_mut() - } - } -} - -pub(crate) fn tcp_write_start_jni( - mut env: JNIEnv, - _class: JClass, - handle: jlong, - data: JByteArray, - timeout_ms: jlong, -) -> jlong { - let data = match env.convert_byte_array(&data) { - Ok(data) => data, - Err(err) => { - throw_exception(&mut env, &format!("Invalid write buffer: {:?}", err)); - return 0; - } - }; - let ptr = if data.is_empty() { - ptr::null() - } else { - data.as_ptr() - }; - let op = unsafe { - data_plane_tcp_write_start( - handle as u64, - ptr, - data.len() as u32, - timeout_from_jlong(timeout_ms), - ) - }; - if op == 0 { - throw_last(&mut env); - } - op as jlong -} - -pub(crate) fn tcp_write_finish_jni(mut env: JNIEnv, _class: JClass, op: jlong) -> jint { - let ret = data_plane_tcp_write_finish(op as u64); - if ret < 0 { - throw_last(&mut env); - } - ret -} - -pub(crate) fn udp_bind_start_jni( - mut env: JNIEnv, - _class: JClass, - inst_name: JString, - local_port: jint, - timeout_ms: jlong, -) -> jlong { - let inst_name = match jstring_to_cstring(&mut env, &inst_name) { - Ok(value) => value, - Err(err) => { - throw_exception(&mut env, &format!("Invalid instance name: {}", err)); - return 0; - } - }; - let Some(local_port) = port_from_jint(&mut env, local_port, "local port") else { - return 0; - }; - let op = unsafe { - data_plane_udp_bind_start( - inst_name.as_ptr(), - local_port, - timeout_from_jlong(timeout_ms), - ) - }; - if op == 0 { - throw_last(&mut env); - } - op as jlong -} - -pub(crate) fn udp_bind_finish_jni(mut env: JNIEnv, _class: JClass, op: jlong) -> jobject { - let mut ip: *const c_char = ptr::null(); - let mut port = 0u16; - let handle = unsafe { data_plane_udp_bind_finish(op as u64, &mut ip, &mut port) }; - if handle == 0 { - throw_last(&mut env); - return ptr::null_mut(); - } - close_udp_socket_on_null( - new_handle_addr_result( - &mut env, - UDP_BIND_RESULT_CLASS, - handle, - unsafe { take_ffi_string(ip) }, - port, - ), - handle, - ) -} - -pub(crate) fn udp_send_to_start_jni( - mut env: JNIEnv, - _class: JClass, - handle: jlong, - dst_ip: JString, - dst_port: jint, - data: JByteArray, - timeout_ms: jlong, -) -> jlong { - let dst_ip = match jstring_to_cstring(&mut env, &dst_ip) { - Ok(value) => value, - Err(err) => { - throw_exception(&mut env, &format!("Invalid destination IP: {}", err)); - return 0; - } - }; - let Some(dst_port) = port_from_jint(&mut env, dst_port, "destination port") else { - return 0; - }; - let data = match env.convert_byte_array(&data) { - Ok(data) => data, - Err(err) => { - throw_exception(&mut env, &format!("Invalid UDP send buffer: {:?}", err)); - return 0; - } - }; - let ptr = if data.is_empty() { - ptr::null() - } else { - data.as_ptr() - }; - let op = unsafe { - data_plane_udp_send_to_start( - handle as u64, - dst_ip.as_ptr(), - dst_port, - ptr, - data.len() as u32, - timeout_from_jlong(timeout_ms), - ) - }; - if op == 0 { - throw_last(&mut env); - } - op as jlong -} - -pub(crate) fn udp_send_to_finish_jni(mut env: JNIEnv, _class: JClass, op: jlong) -> jint { - let ret = data_plane_udp_send_to_finish(op as u64); - if ret < 0 { - throw_last(&mut env); - } - ret -} - -pub(crate) fn udp_recv_from_start_jni( - mut env: JNIEnv, - _class: JClass, - handle: jlong, - max_len: jint, - timeout_ms: jlong, -) -> jlong { - let Some(max_len) = len_from_jint(&mut env, max_len, "max length") else { - return 0; - }; - let op = unsafe { - data_plane_udp_recv_from_start(handle as u64, max_len, timeout_from_jlong(timeout_ms)) - }; - if op == 0 { - throw_last(&mut env); - } - op as jlong -} - -pub(crate) fn udp_recv_from_finish_jni(mut env: JNIEnv, _class: JClass, op: jlong) -> jobject { - let mut ptr: *const u8 = ptr::null(); - let mut len = 0u32; - let mut ip: *const c_char = ptr::null(); - let mut port = 0u16; - let ret = unsafe { - data_plane_udp_recv_from_finish(op as u64, &mut ptr, &mut len, &mut ip, &mut port) - }; - if ret < 0 { - throw_last(&mut env); - return ptr::null_mut(); - } - let bytes = read_owned_bytes(ptr, len); - let array = match env.byte_array_from_slice(&bytes) { - Ok(array) => array, - Err(err) => { - free_string(ip); - throw_exception(&mut env, &format!("Failed to create byte array: {:?}", err)); - return ptr::null_mut(); - } - }; - let Some(peer_addr) = new_socket_addr(&mut env, unsafe { take_ffi_string(ip) }, port) else { - return ptr::null_mut(); - }; - let class = match env.find_class(UDP_RECV_RESULT_CLASS) { - Ok(class) => class, - Err(err) => { - throw_exception( - &mut env, - &format!("Failed to find UDP recv result class: {:?}", err), - ); - return ptr::null_mut(); - } - }; - let sig = format!("([BL{};)V", SOCKET_ADDR_CLASS); - match env.new_object( - class, - sig.as_str(), - &[JValue::Object(&array), JValue::Object(&peer_addr)], - ) { - Ok(result) => result.into_raw(), - Err(err) => { - throw_exception( - &mut env, - &format!("Failed to create UDP recv result: {:?}", err), - ); - ptr::null_mut() - } - } -} - -pub(crate) fn tcp_close_jni(mut env: JNIEnv, _class: JClass, handle: jlong) -> jint { - let ret = data_plane_tcp_close(handle as u64); - if ret != 0 { - throw_last(&mut env); - } - ret -} - -pub(crate) fn tcp_listener_close_jni(mut env: JNIEnv, _class: JClass, handle: jlong) -> jint { - let ret = data_plane_tcp_listener_close(handle as u64); - if ret != 0 { - throw_last(&mut env); - } - ret -} - -pub(crate) fn udp_close_jni(mut env: JNIEnv, _class: JClass, handle: jlong) -> jint { - let ret = data_plane_udp_close(handle as u64); - if ret != 0 { - throw_last(&mut env); - } - ret -} diff --git a/easytier-contrib/easytier-android-jni/src/lib.rs b/easytier-contrib/easytier-android-jni/src/lib.rs index f673023e..e70180d8 100644 --- a/easytier-contrib/easytier-android-jni/src/lib.rs +++ b/easytier-contrib/easytier-android-jni/src/lib.rs @@ -22,13 +22,8 @@ //! Error API: //! - `getLastError()`: return the latest FFI/JNI error string for the calling thread. //! -//! Data-plane APIs: -//! - `EasyTierDataPlaneJNI.*`: low-level async op-handle data-plane JNI. -//! - `EasyTierJNI.dataPlane*`: compatibility exports for older callers. - mod callback; mod config_server_api; -mod data_plane_api; mod error; mod json_rpc_api; mod logger; @@ -36,8 +31,8 @@ mod network_api; mod strings; use jni::JNIEnv; -use jni::objects::{JByteArray, JClass, JObject, JObjectArray, JString}; -use jni::sys::{jboolean, jint, jlong, jobject, jstring}; +use jni::objects::{JClass, JObject, JObjectArray, JString}; +use jni::sys::{jboolean, jint, jstring}; /// Attach a TUN file descriptor to an EasyTier network instance. /// @@ -256,522 +251,3 @@ pub extern "system" fn Java_com_easytier_jni_EasyTierJNI_isConfigServerClientCon logger::init(); config_server_api::is_config_server_client_connected_jni(env, class) } - -macro_rules! export_data_plane_jni { - ( - $op_status:ident, - $op_wait:ident, - $op_cancel:ident, - $op_free:ident, - $tcp_connect_start:ident, - $tcp_connect_finish:ident, - $tcp_bind_start:ident, - $tcp_bind_finish:ident, - $tcp_accept_start:ident, - $tcp_accept_finish:ident, - $tcp_read_start:ident, - $tcp_read_finish:ident, - $tcp_write_start:ident, - $tcp_write_finish:ident, - $udp_bind_start:ident, - $udp_bind_finish:ident, - $udp_send_to_start:ident, - $udp_send_to_finish:ident, - $udp_recv_from_start:ident, - $udp_recv_from_finish:ident, - $tcp_close:ident, - $tcp_listener_close:ident, - $udp_close:ident - ) => { - #[unsafe(no_mangle)] - pub extern "system" fn $op_status(env: JNIEnv, class: JClass, handle: jlong) -> jint { - logger::init(); - data_plane_api::async_op_status_jni(env, class, handle) - } - - #[unsafe(no_mangle)] - pub extern "system" fn $op_wait( - env: JNIEnv, - class: JClass, - handle: jlong, - timeout_ms: jlong, - ) -> jint { - logger::init(); - data_plane_api::async_op_wait_jni(env, class, handle, timeout_ms) - } - - #[unsafe(no_mangle)] - pub extern "system" fn $op_cancel(env: JNIEnv, class: JClass, handle: jlong) -> jint { - logger::init(); - data_plane_api::async_op_cancel_jni(env, class, handle) - } - - #[unsafe(no_mangle)] - pub extern "system" fn $op_free(env: JNIEnv, class: JClass, handle: jlong) -> jint { - logger::init(); - data_plane_api::async_op_free_jni(env, class, handle) - } - - #[unsafe(no_mangle)] - pub extern "system" fn $tcp_connect_start( - env: JNIEnv, - class: JClass, - inst_name: JString, - dst_ip: JString, - dst_port: jint, - timeout_ms: jlong, - ) -> jlong { - logger::init(); - data_plane_api::tcp_connect_start_jni( - env, class, inst_name, dst_ip, dst_port, timeout_ms, - ) - } - - #[unsafe(no_mangle)] - pub extern "system" fn $tcp_connect_finish( - env: JNIEnv, - class: JClass, - op: jlong, - ) -> jobject { - logger::init(); - data_plane_api::tcp_connect_finish_jni(env, class, op) - } - - #[unsafe(no_mangle)] - pub extern "system" fn $tcp_bind_start( - env: JNIEnv, - class: JClass, - inst_name: JString, - local_port: jint, - timeout_ms: jlong, - ) -> jlong { - logger::init(); - data_plane_api::tcp_bind_start_jni(env, class, inst_name, local_port, timeout_ms) - } - - #[unsafe(no_mangle)] - pub extern "system" fn $tcp_bind_finish(env: JNIEnv, class: JClass, op: jlong) -> jobject { - logger::init(); - data_plane_api::tcp_bind_finish_jni(env, class, op) - } - - #[unsafe(no_mangle)] - pub extern "system" fn $tcp_accept_start( - env: JNIEnv, - class: JClass, - handle: jlong, - timeout_ms: jlong, - ) -> jlong { - logger::init(); - data_plane_api::tcp_accept_start_jni(env, class, handle, timeout_ms) - } - - #[unsafe(no_mangle)] - pub extern "system" fn $tcp_accept_finish( - env: JNIEnv, - class: JClass, - op: jlong, - ) -> jobject { - logger::init(); - data_plane_api::tcp_accept_finish_jni(env, class, op) - } - - #[unsafe(no_mangle)] - pub extern "system" fn $tcp_read_start( - env: JNIEnv, - class: JClass, - handle: jlong, - max_len: jint, - timeout_ms: jlong, - ) -> jlong { - logger::init(); - data_plane_api::tcp_read_start_jni(env, class, handle, max_len, timeout_ms) - } - - #[unsafe(no_mangle)] - pub extern "system" fn $tcp_read_finish(env: JNIEnv, class: JClass, op: jlong) -> jobject { - logger::init(); - data_plane_api::tcp_read_finish_jni(env, class, op) - } - - #[unsafe(no_mangle)] - pub extern "system" fn $tcp_write_start( - env: JNIEnv, - class: JClass, - handle: jlong, - data: JByteArray, - timeout_ms: jlong, - ) -> jlong { - logger::init(); - data_plane_api::tcp_write_start_jni(env, class, handle, data, timeout_ms) - } - - #[unsafe(no_mangle)] - pub extern "system" fn $tcp_write_finish(env: JNIEnv, class: JClass, op: jlong) -> jint { - logger::init(); - data_plane_api::tcp_write_finish_jni(env, class, op) - } - - #[unsafe(no_mangle)] - pub extern "system" fn $udp_bind_start( - env: JNIEnv, - class: JClass, - inst_name: JString, - local_port: jint, - timeout_ms: jlong, - ) -> jlong { - logger::init(); - data_plane_api::udp_bind_start_jni(env, class, inst_name, local_port, timeout_ms) - } - - #[unsafe(no_mangle)] - pub extern "system" fn $udp_bind_finish(env: JNIEnv, class: JClass, op: jlong) -> jobject { - logger::init(); - data_plane_api::udp_bind_finish_jni(env, class, op) - } - - #[unsafe(no_mangle)] - pub extern "system" fn $udp_send_to_start( - env: JNIEnv, - class: JClass, - handle: jlong, - dst_ip: JString, - dst_port: jint, - data: JByteArray, - timeout_ms: jlong, - ) -> jlong { - logger::init(); - data_plane_api::udp_send_to_start_jni( - env, class, handle, dst_ip, dst_port, data, timeout_ms, - ) - } - - #[unsafe(no_mangle)] - pub extern "system" fn $udp_send_to_finish(env: JNIEnv, class: JClass, op: jlong) -> jint { - logger::init(); - data_plane_api::udp_send_to_finish_jni(env, class, op) - } - - #[unsafe(no_mangle)] - pub extern "system" fn $udp_recv_from_start( - env: JNIEnv, - class: JClass, - handle: jlong, - max_len: jint, - timeout_ms: jlong, - ) -> jlong { - logger::init(); - data_plane_api::udp_recv_from_start_jni(env, class, handle, max_len, timeout_ms) - } - - #[unsafe(no_mangle)] - pub extern "system" fn $udp_recv_from_finish( - env: JNIEnv, - class: JClass, - op: jlong, - ) -> jobject { - logger::init(); - data_plane_api::udp_recv_from_finish_jni(env, class, op) - } - - #[unsafe(no_mangle)] - pub extern "system" fn $tcp_close(env: JNIEnv, class: JClass, handle: jlong) -> jint { - logger::init(); - data_plane_api::tcp_close_jni(env, class, handle) - } - - #[unsafe(no_mangle)] - pub extern "system" fn $tcp_listener_close( - env: JNIEnv, - class: JClass, - handle: jlong, - ) -> jint { - logger::init(); - data_plane_api::tcp_listener_close_jni(env, class, handle) - } - - #[unsafe(no_mangle)] - pub extern "system" fn $udp_close(env: JNIEnv, class: JClass, handle: jlong) -> jint { - logger::init(); - data_plane_api::udp_close_jni(env, class, handle) - } - }; -} - -export_data_plane_jni!( - Java_com_easytier_jni_EasyTierDataPlaneJNI_dataPlaneAsyncOpStatus, - Java_com_easytier_jni_EasyTierDataPlaneJNI_dataPlaneAsyncOpWait, - Java_com_easytier_jni_EasyTierDataPlaneJNI_dataPlaneAsyncOpCancel, - Java_com_easytier_jni_EasyTierDataPlaneJNI_dataPlaneAsyncOpFree, - Java_com_easytier_jni_EasyTierDataPlaneJNI_dataPlaneTcpConnectStart, - Java_com_easytier_jni_EasyTierDataPlaneJNI_dataPlaneTcpConnectFinish, - Java_com_easytier_jni_EasyTierDataPlaneJNI_dataPlaneTcpBindStart, - Java_com_easytier_jni_EasyTierDataPlaneJNI_dataPlaneTcpBindFinish, - Java_com_easytier_jni_EasyTierDataPlaneJNI_dataPlaneTcpAcceptStart, - Java_com_easytier_jni_EasyTierDataPlaneJNI_dataPlaneTcpAcceptFinish, - Java_com_easytier_jni_EasyTierDataPlaneJNI_dataPlaneTcpReadStart, - Java_com_easytier_jni_EasyTierDataPlaneJNI_dataPlaneTcpReadFinish, - Java_com_easytier_jni_EasyTierDataPlaneJNI_dataPlaneTcpWriteStart, - Java_com_easytier_jni_EasyTierDataPlaneJNI_dataPlaneTcpWriteFinish, - Java_com_easytier_jni_EasyTierDataPlaneJNI_dataPlaneUdpBindStart, - Java_com_easytier_jni_EasyTierDataPlaneJNI_dataPlaneUdpBindFinish, - Java_com_easytier_jni_EasyTierDataPlaneJNI_dataPlaneUdpSendToStart, - Java_com_easytier_jni_EasyTierDataPlaneJNI_dataPlaneUdpSendToFinish, - Java_com_easytier_jni_EasyTierDataPlaneJNI_dataPlaneUdpRecvFromStart, - Java_com_easytier_jni_EasyTierDataPlaneJNI_dataPlaneUdpRecvFromFinish, - Java_com_easytier_jni_EasyTierDataPlaneJNI_dataPlaneTcpClose, - Java_com_easytier_jni_EasyTierDataPlaneJNI_dataPlaneTcpListenerClose, - Java_com_easytier_jni_EasyTierDataPlaneJNI_dataPlaneUdpClose -); - -// Compatibility exports for older Kotlin/Java callers that used EasyTierJNI -// directly for data-plane operations. - -#[unsafe(no_mangle)] -pub extern "system" fn Java_com_easytier_jni_EasyTierJNI_dataPlaneAsyncOpStatus( - env: JNIEnv, - class: JClass, - handle: jlong, -) -> jint { - logger::init(); - data_plane_api::async_op_status_jni(env, class, handle) -} - -#[unsafe(no_mangle)] -pub extern "system" fn Java_com_easytier_jni_EasyTierJNI_dataPlaneAsyncOpWait( - env: JNIEnv, - class: JClass, - handle: jlong, - timeout_ms: jlong, -) -> jint { - logger::init(); - data_plane_api::async_op_wait_jni(env, class, handle, timeout_ms) -} - -#[unsafe(no_mangle)] -pub extern "system" fn Java_com_easytier_jni_EasyTierJNI_dataPlaneAsyncOpCancel( - env: JNIEnv, - class: JClass, - handle: jlong, -) -> jint { - logger::init(); - data_plane_api::async_op_cancel_jni(env, class, handle) -} - -#[unsafe(no_mangle)] -pub extern "system" fn Java_com_easytier_jni_EasyTierJNI_dataPlaneAsyncOpFree( - env: JNIEnv, - class: JClass, - handle: jlong, -) -> jint { - logger::init(); - data_plane_api::async_op_free_jni(env, class, handle) -} - -#[unsafe(no_mangle)] -pub extern "system" fn Java_com_easytier_jni_EasyTierJNI_dataPlaneTcpConnectStart( - env: JNIEnv, - class: JClass, - inst_name: JString, - dst_ip: JString, - dst_port: jint, - timeout_ms: jlong, -) -> jlong { - logger::init(); - data_plane_api::tcp_connect_start_jni(env, class, inst_name, dst_ip, dst_port, timeout_ms) -} - -#[unsafe(no_mangle)] -pub extern "system" fn Java_com_easytier_jni_EasyTierJNI_dataPlaneTcpConnectFinish( - env: JNIEnv, - class: JClass, - op: jlong, -) -> jobject { - logger::init(); - data_plane_api::tcp_connect_finish_jni(env, class, op) -} - -#[unsafe(no_mangle)] -pub extern "system" fn Java_com_easytier_jni_EasyTierJNI_dataPlaneTcpBindStart( - env: JNIEnv, - class: JClass, - inst_name: JString, - local_port: jint, - timeout_ms: jlong, -) -> jlong { - logger::init(); - data_plane_api::tcp_bind_start_jni(env, class, inst_name, local_port, timeout_ms) -} - -#[unsafe(no_mangle)] -pub extern "system" fn Java_com_easytier_jni_EasyTierJNI_dataPlaneTcpBindFinish( - env: JNIEnv, - class: JClass, - op: jlong, -) -> jobject { - logger::init(); - data_plane_api::tcp_bind_finish_jni(env, class, op) -} - -#[unsafe(no_mangle)] -pub extern "system" fn Java_com_easytier_jni_EasyTierJNI_dataPlaneTcpAcceptStart( - env: JNIEnv, - class: JClass, - handle: jlong, - timeout_ms: jlong, -) -> jlong { - logger::init(); - data_plane_api::tcp_accept_start_jni(env, class, handle, timeout_ms) -} - -#[unsafe(no_mangle)] -pub extern "system" fn Java_com_easytier_jni_EasyTierJNI_dataPlaneTcpAcceptFinish( - env: JNIEnv, - class: JClass, - op: jlong, -) -> jobject { - logger::init(); - data_plane_api::tcp_accept_finish_jni(env, class, op) -} - -#[unsafe(no_mangle)] -pub extern "system" fn Java_com_easytier_jni_EasyTierJNI_dataPlaneTcpReadStart( - env: JNIEnv, - class: JClass, - handle: jlong, - max_len: jint, - timeout_ms: jlong, -) -> jlong { - logger::init(); - data_plane_api::tcp_read_start_jni(env, class, handle, max_len, timeout_ms) -} - -#[unsafe(no_mangle)] -pub extern "system" fn Java_com_easytier_jni_EasyTierJNI_dataPlaneTcpReadFinish( - env: JNIEnv, - class: JClass, - op: jlong, -) -> jobject { - logger::init(); - data_plane_api::tcp_read_finish_jni(env, class, op) -} - -#[unsafe(no_mangle)] -pub extern "system" fn Java_com_easytier_jni_EasyTierJNI_dataPlaneTcpWriteStart( - env: JNIEnv, - class: JClass, - handle: jlong, - data: JByteArray, - timeout_ms: jlong, -) -> jlong { - logger::init(); - data_plane_api::tcp_write_start_jni(env, class, handle, data, timeout_ms) -} - -#[unsafe(no_mangle)] -pub extern "system" fn Java_com_easytier_jni_EasyTierJNI_dataPlaneTcpWriteFinish( - env: JNIEnv, - class: JClass, - op: jlong, -) -> jint { - logger::init(); - data_plane_api::tcp_write_finish_jni(env, class, op) -} - -#[unsafe(no_mangle)] -pub extern "system" fn Java_com_easytier_jni_EasyTierJNI_dataPlaneUdpBindStart( - env: JNIEnv, - class: JClass, - inst_name: JString, - local_port: jint, - timeout_ms: jlong, -) -> jlong { - logger::init(); - data_plane_api::udp_bind_start_jni(env, class, inst_name, local_port, timeout_ms) -} - -#[unsafe(no_mangle)] -pub extern "system" fn Java_com_easytier_jni_EasyTierJNI_dataPlaneUdpBindFinish( - env: JNIEnv, - class: JClass, - op: jlong, -) -> jobject { - logger::init(); - data_plane_api::udp_bind_finish_jni(env, class, op) -} - -#[unsafe(no_mangle)] -pub extern "system" fn Java_com_easytier_jni_EasyTierJNI_dataPlaneUdpSendToStart( - env: JNIEnv, - class: JClass, - handle: jlong, - dst_ip: JString, - dst_port: jint, - data: JByteArray, - timeout_ms: jlong, -) -> jlong { - logger::init(); - data_plane_api::udp_send_to_start_jni(env, class, handle, dst_ip, dst_port, data, timeout_ms) -} - -#[unsafe(no_mangle)] -pub extern "system" fn Java_com_easytier_jni_EasyTierJNI_dataPlaneUdpSendToFinish( - env: JNIEnv, - class: JClass, - op: jlong, -) -> jint { - logger::init(); - data_plane_api::udp_send_to_finish_jni(env, class, op) -} - -#[unsafe(no_mangle)] -pub extern "system" fn Java_com_easytier_jni_EasyTierJNI_dataPlaneUdpRecvFromStart( - env: JNIEnv, - class: JClass, - handle: jlong, - max_len: jint, - timeout_ms: jlong, -) -> jlong { - logger::init(); - data_plane_api::udp_recv_from_start_jni(env, class, handle, max_len, timeout_ms) -} - -#[unsafe(no_mangle)] -pub extern "system" fn Java_com_easytier_jni_EasyTierJNI_dataPlaneUdpRecvFromFinish( - env: JNIEnv, - class: JClass, - op: jlong, -) -> jobject { - logger::init(); - data_plane_api::udp_recv_from_finish_jni(env, class, op) -} - -#[unsafe(no_mangle)] -pub extern "system" fn Java_com_easytier_jni_EasyTierJNI_dataPlaneTcpClose( - env: JNIEnv, - class: JClass, - handle: jlong, -) -> jint { - logger::init(); - data_plane_api::tcp_close_jni(env, class, handle) -} - -#[unsafe(no_mangle)] -pub extern "system" fn Java_com_easytier_jni_EasyTierJNI_dataPlaneTcpListenerClose( - env: JNIEnv, - class: JClass, - handle: jlong, -) -> jint { - logger::init(); - data_plane_api::tcp_listener_close_jni(env, class, handle) -} - -#[unsafe(no_mangle)] -pub extern "system" fn Java_com_easytier_jni_EasyTierJNI_dataPlaneUdpClose( - env: JNIEnv, - class: JClass, - handle: jlong, -) -> jint { - logger::init(); - data_plane_api::udp_close_jni(env, class, handle) -} diff --git a/easytier-contrib/easytier-ffi/Cargo.toml b/easytier-contrib/easytier-ffi/Cargo.toml index 1b1cb41f..d75af201 100644 --- a/easytier-contrib/easytier-ffi/Cargo.toml +++ b/easytier-contrib/easytier-ffi/Cargo.toml @@ -9,20 +9,20 @@ crate-type = ["cdylib", "rlib"] [features] default = ["c-abi", "ffi-dataplane"] c-abi = [] -ffi-dataplane = ["easytier/ffi-dataplane"] +ffi-dataplane = [ + "easytier/ffi-dataplane", + "easytier-core/proxy-smoltcp-stack", +] [dependencies] -easytier = { path = "../../easytier" } +easytier = { path = "../../easytier", features = ["tracing-log"] } +easytier-core = { path = "../../easytier-core" } once_cell = "1.18.0" -dashmap = "6.0" tokio = { version = "1", features = ["rt-multi-thread", "io-util", "time", "sync", "macros"] } async-trait = "0.1" log = "0.4" -percent-encoding = "2.3" -url = "2" serde = { version = "1.0", features = ["derive"] } serde_json = "1" uuid = "1.17.0" -tokio-util = "0.7" diff --git a/easytier-contrib/easytier-ffi/DATA_PLANE_ABI.md b/easytier-contrib/easytier-ffi/DATA_PLANE_ABI.md new file mode 100644 index 00000000..ec7b1a8c --- /dev/null +++ b/easytier-contrib/easytier-ffi/DATA_PLANE_ABI.md @@ -0,0 +1,100 @@ +# Native data-plane ABI v2 + +The native data-plane ABI is a thin adapter over the instance-owned +`DataPlaneSession`. It does not own sockets, operation state, completion +queues, routing policy, or timeouts. + +## Conventions + +- Every immediate call returns `0` on success or a negative + `DataPlaneErrorKind` value on failure. +- `data_plane_completion_wait` returns `1` when a completion is ready, `0` on + timeout or session close, and a negative error value on failure. +- `data_plane_completion_drain` returns a non-negative descriptor count or a + negative error value. +- Handle zero is invalid. +- `timeout_ms == UINT64_MAX` means no deadline. Every other timeout starts when + submission is accepted, including time spent waiting for an I/O direction + lock. +- Request and write bytes are copied before a submit call returns. +- Socket-address fields use native-endian integers. Address bytes are in + network order. ABI v2 accepts IPv4 only. + +`DataPlaneSocketAddr` is: + +```c +typedef struct { + uint16_t family; /* 4 */ + uint16_t port; + uint8_t address[16]; /* IPv4 uses the first four bytes */ +} DataPlaneSocketAddr; +``` + +`DataPlaneCompletion` is: + +```c +typedef struct { + uint64_t operation_id; + uint16_t operation_kind; + uint16_t status; /* 0 or DataPlaneErrorKind */ +} DataPlaneCompletion; +``` + +## Lifecycle + +One native session may be open for an EasyTier instance at a time: + +```text +data_plane_session_open + -> submit operations + -> completion_wait + -> completion_drain + -> typed result_take + -> resource_close / operation_free +data_plane_session_close +``` + +Closing a native session cancels and discards its outstanding operations and +resources and wakes a thread blocked in `data_plane_completion_wait`. + +The resource and operation IDs returned by the ABI belong to that session. +They must always be passed together with the same session handle. + +## Completion and result ownership + +Submission returns an operation ID immediately. Completion descriptors carry +only the operation ID, operation kind, and terminal status. Draining a +descriptor makes its typed result available but does not consume it. + +`data_plane_result_size` reports the TCP-read or UDP-receive payload size. +Typed result-take functions consume the result exactly once. If a supplied +buffer is too small, they return `-BufferTooSmall` and leave the result +available for a later call. + +Call `data_plane_operation_free` when a drained result is intentionally +abandoned. Call `data_plane_resource_close` for TCP streams, listeners, and +UDP sockets. + +## Operation kinds + +| Value | Operation | +| ---: | --- | +| 1 | TCP connect | +| 2 | TCP bind | +| 3 | TCP accept | +| 4 | TCP read | +| 5 | TCP write | +| 6 | UDP bind | +| 7 | UDP receive | +| 8 | UDP send | + +The exported function families are: + +- `data_plane_tcp_*_submit` +- `data_plane_udp_*_submit` +- `data_plane_completion_wait` +- `data_plane_completion_drain` +- `data_plane_*_result_take` +- `data_plane_operation_cancel` +- `data_plane_operation_free` +- `data_plane_resource_close` diff --git a/easytier-contrib/easytier-ffi/examples/example_data_plane_async.c b/easytier-contrib/easytier-ffi/examples/example_data_plane_async.c deleted file mode 100644 index 24f3bd9d..00000000 --- a/easytier-contrib/easytier-ffi/examples/example_data_plane_async.c +++ /dev/null @@ -1,429 +0,0 @@ -#include -#include -#include -#include - -#define DATA_PLANE_OP_PENDING 0 -#define DATA_PLANE_OP_READY 1 -#define DATA_PLANE_OP_FAILED -1 -#define DATA_PLANE_OP_INVALID -2 - -extern int run_network_instance(const char *cfg_str); -extern void get_error_msg(const char **out); -extern void free_string(const char *s); - -extern int data_plane_async_op_status(uint64_t op); -extern int data_plane_async_op_wait(uint64_t op, uint64_t timeout_ms); -extern int data_plane_async_op_cancel(uint64_t op); -extern int data_plane_async_op_free(uint64_t op); -extern void data_plane_free_bytes(const uint8_t *ptr, uint32_t len); - -extern uint64_t data_plane_tcp_connect_start( - const char *inst_name, - const char *dst_ip, - uint16_t dst_port, - uint64_t timeout_ms); -extern uint64_t data_plane_tcp_connect_finish( - uint64_t op, - const char **out_local_ip, - uint16_t *out_local_port); -extern uint64_t data_plane_tcp_bind_start( - const char *inst_name, - uint16_t local_port, - uint64_t timeout_ms); -extern uint64_t data_plane_tcp_bind_finish( - uint64_t op, - const char **out_local_ip, - uint16_t *out_local_port); -extern uint64_t data_plane_tcp_accept_start(uint64_t listener, uint64_t timeout_ms); -extern uint64_t data_plane_tcp_accept_finish( - uint64_t op, - const char **out_local_ip, - uint16_t *out_local_port, - const char **out_peer_ip, - uint16_t *out_peer_port); -extern uint64_t data_plane_tcp_read_start( - uint64_t stream, - uint32_t max_len, - uint64_t timeout_ms); -extern int data_plane_tcp_read_finish( - uint64_t op, - const uint8_t **out_buf, - uint32_t *out_len); -extern uint64_t data_plane_tcp_write_start( - uint64_t stream, - const uint8_t *buf, - uint32_t len, - uint64_t timeout_ms); -extern int data_plane_tcp_write_finish(uint64_t op); -extern int data_plane_tcp_close(uint64_t stream); -extern int data_plane_tcp_listener_close(uint64_t listener); - -extern uint64_t data_plane_udp_bind_start( - const char *inst_name, - uint16_t local_port, - uint64_t timeout_ms); -extern uint64_t data_plane_udp_bind_finish( - uint64_t op, - const char **out_local_ip, - uint16_t *out_local_port); -extern uint64_t data_plane_udp_send_to_start( - uint64_t socket, - const char *dst_ip, - uint16_t dst_port, - const uint8_t *buf, - uint32_t len, - uint64_t timeout_ms); -extern int data_plane_udp_send_to_finish(uint64_t op); -extern uint64_t data_plane_udp_recv_from_start( - uint64_t socket, - uint32_t max_len, - uint64_t timeout_ms); -extern int data_plane_udp_recv_from_finish( - uint64_t op, - const uint8_t **out_buf, - uint32_t *out_len, - const char **out_ip, - uint16_t *out_port); -extern int data_plane_udp_close(uint64_t socket); - -static void print_last_error(const char *prefix) { - const char *err = NULL; - get_error_msg(&err); - if (err) { - fprintf(stderr, "%s: %s\n", prefix, err); - free_string(err); - } else { - fprintf(stderr, "%s\n", prefix); - } -} - -static int parse_ip_port(const char *value, char *ip, size_t ip_len, uint16_t *port) { - const char *colon = strrchr(value, ':'); - if (!colon || colon == value || !colon[1]) { - fprintf(stderr, "expected IPv4 target in IP:PORT form, got %s\n", value); - return -1; - } - size_t host_len = (size_t)(colon - value); - if (host_len >= ip_len) { - fprintf(stderr, "IP address is too long: %s\n", value); - return -1; - } - char *end = NULL; - long parsed_port = strtol(colon + 1, &end, 10); - if (!end || *end != '\0' || parsed_port < 0 || parsed_port > 65535) { - fprintf(stderr, "invalid port in %s\n", value); - return -1; - } - memcpy(ip, value, host_len); - ip[host_len] = '\0'; - *port = (uint16_t)parsed_port; - return 0; -} - -static int wait_op(uint64_t op, uint64_t timeout_ms) { - uint64_t waited = 0; - while (waited < timeout_ms) { - int status = data_plane_async_op_wait(op, 50); - if (status != DATA_PLANE_OP_PENDING) { - return status; - } - waited += 50; - } - return data_plane_async_op_status(op); -} - -static int wait_or_cancel(uint64_t op, uint64_t timeout_ms, const char *what) { - int status = wait_op(op, timeout_ms); - if (status == DATA_PLANE_OP_READY || status == DATA_PLANE_OP_FAILED) { - return status; - } - if (status == DATA_PLANE_OP_PENDING) { - fprintf(stderr, "%s did not finish within %llu ms\n", what, (unsigned long long)timeout_ms); - data_plane_async_op_cancel(op); - data_plane_async_op_free(op); - return DATA_PLANE_OP_INVALID; - } - fprintf(stderr, "%s returned invalid op status %d\n", what, status); - return status; -} - -static int async_tcp_read_once(uint64_t stream, uint64_t timeout_ms) { - uint64_t op = data_plane_tcp_read_start(stream, 512, timeout_ms); - if (!op) { - print_last_error("tcp read start failed"); - return -1; - } - if (wait_or_cancel(op, timeout_ms + 1000, "tcp read") == DATA_PLANE_OP_INVALID) { - return -1; - } - - const uint8_t *buf = NULL; - uint32_t len = 0; - int ret = data_plane_tcp_read_finish(op, &buf, &len); - if (ret < 0) { - print_last_error("tcp read finish failed"); - return -1; - } - printf("tcp read %d bytes: %.*s\n", ret, ret, buf ? (const char *)buf : ""); - data_plane_free_bytes(buf, len); - return 0; -} - -static int async_tcp_write_all(uint64_t stream, const char *data, uint64_t timeout_ms) { - uint64_t op = data_plane_tcp_write_start( - stream, - (const uint8_t *)data, - (uint32_t)strlen(data), - timeout_ms); - if (!op) { - print_last_error("tcp write start failed"); - return -1; - } - if (wait_or_cancel(op, timeout_ms + 1000, "tcp write") == DATA_PLANE_OP_INVALID) { - return -1; - } - int ret = data_plane_tcp_write_finish(op); - if (ret < 0) { - print_last_error("tcp write finish failed"); - return -1; - } - printf("tcp wrote %d bytes\n", ret); - return 0; -} - -static int run_tcp_connect_demo(const char *inst, const char *target) { - char ip[128]; - uint16_t port = 0; - if (parse_ip_port(target, ip, sizeof(ip), &port) != 0) { - return -1; - } - - uint64_t op = data_plane_tcp_connect_start(inst, ip, port, 30000); - if (!op) { - print_last_error("tcp connect start failed"); - return -1; - } - if (wait_or_cancel(op, 31000, "tcp connect") == DATA_PLANE_OP_INVALID) { - return -1; - } - - const char *local_ip = NULL; - uint16_t local_port = 0; - uint64_t stream = data_plane_tcp_connect_finish(op, &local_ip, &local_port); - if (!stream) { - print_last_error("tcp connect finish failed"); - return -1; - } - printf("tcp connected from %s:%u to %s:%u, handle=%llu\n", - local_ip, - local_port, - ip, - port, - (unsigned long long)stream); - free_string(local_ip); - - int ret = async_tcp_read_once(stream, 10000); - data_plane_tcp_close(stream); - return ret; -} - -static int run_tcp_listen_demo(const char *inst, const char *port_text) { - uint16_t port = (uint16_t)strtoul(port_text, NULL, 10); - uint64_t op = data_plane_tcp_bind_start(inst, port, 30000); - if (!op) { - print_last_error("tcp bind start failed"); - return -1; - } - if (wait_or_cancel(op, 31000, "tcp bind") == DATA_PLANE_OP_INVALID) { - return -1; - } - - const char *local_ip = NULL; - uint16_t local_port = 0; - uint64_t listener = data_plane_tcp_bind_finish(op, &local_ip, &local_port); - if (!listener) { - print_last_error("tcp bind finish failed"); - return -1; - } - printf("tcp listening on %s:%u, handle=%llu\n", - local_ip, - local_port, - (unsigned long long)listener); - free_string(local_ip); - - op = data_plane_tcp_accept_start(listener, 60000); - if (!op) { - print_last_error("tcp accept start failed"); - data_plane_tcp_listener_close(listener); - return -1; - } - if (wait_or_cancel(op, 61000, "tcp accept") == DATA_PLANE_OP_INVALID) { - data_plane_tcp_listener_close(listener); - return -1; - } - - const char *peer_ip = NULL; - uint16_t peer_port = 0; - local_ip = NULL; - local_port = 0; - uint64_t stream = data_plane_tcp_accept_finish( - op, - &local_ip, - &local_port, - &peer_ip, - &peer_port); - data_plane_tcp_listener_close(listener); - if (!stream) { - print_last_error("tcp accept finish failed"); - return -1; - } - printf("tcp accepted %s:%u -> %s:%u, stream=%llu\n", - peer_ip, - peer_port, - local_ip, - local_port, - (unsigned long long)stream); - free_string(local_ip); - free_string(peer_ip); - - int ret = async_tcp_read_once(stream, 10000); - if (ret == 0) { - ret = async_tcp_write_all(stream, "pong", 10000); - } - data_plane_tcp_close(stream); - return ret; -} - -static int run_udp_demo(const char *inst, const char *target) { - char ip[128]; - uint16_t port = 0; - if (parse_ip_port(target, ip, sizeof(ip), &port) != 0) { - return -1; - } - - uint64_t op = data_plane_udp_bind_start(inst, 0, 30000); - if (!op) { - print_last_error("udp bind start failed"); - return -1; - } - if (wait_or_cancel(op, 31000, "udp bind") == DATA_PLANE_OP_INVALID) { - return -1; - } - - const char *local_ip = NULL; - uint16_t local_port = 0; - uint64_t socket = data_plane_udp_bind_finish(op, &local_ip, &local_port); - if (!socket) { - print_last_error("udp bind finish failed"); - return -1; - } - printf("udp bound on %s:%u, handle=%llu\n", - local_ip, - local_port, - (unsigned long long)socket); - free_string(local_ip); - - const char payload[] = "ping"; - op = data_plane_udp_send_to_start( - socket, - ip, - port, - (const uint8_t *)payload, - (uint32_t)strlen(payload), - 10000); - if (!op) { - print_last_error("udp send start failed"); - data_plane_udp_close(socket); - return -1; - } - if (wait_or_cancel(op, 11000, "udp send") == DATA_PLANE_OP_INVALID) { - data_plane_udp_close(socket); - return -1; - } - int sent = data_plane_udp_send_to_finish(op); - if (sent < 0) { - print_last_error("udp send finish failed"); - data_plane_udp_close(socket); - return -1; - } - printf("udp sent %d bytes to %s:%u\n", sent, ip, port); - - op = data_plane_udp_recv_from_start(socket, 512, 30000); - if (!op) { - print_last_error("udp recv start failed"); - data_plane_udp_close(socket); - return -1; - } - if (wait_or_cancel(op, 31000, "udp recv") == DATA_PLANE_OP_INVALID) { - data_plane_udp_close(socket); - return -1; - } - - const uint8_t *buf = NULL; - uint32_t len = 0; - const char *peer_ip = NULL; - uint16_t peer_port = 0; - int ret = data_plane_udp_recv_from_finish(op, &buf, &len, &peer_ip, &peer_port); - if (ret < 0) { - print_last_error("udp recv finish failed"); - data_plane_udp_close(socket); - return -1; - } - printf("udp received %d bytes from %s:%u: %.*s\n", - ret, - peer_ip, - peer_port, - ret, - buf ? (const char *)buf : ""); - data_plane_free_bytes(buf, len); - free_string(peer_ip); - data_plane_udp_close(socket); - return 0; -} - -static void print_usage(void) { - printf("Set EASYTIER_FFI_CONFIG and EASYTIER_FFI_INSTANCE to run the async data-plane demo.\n"); - printf("Optional demos:\n"); - printf(" EASYTIER_FFI_TARGET=10.0.0.2:22 async TCP connect/read\n"); - printf(" EASYTIER_FFI_LISTEN_PORT=12345 async TCP bind/accept/read/write\n"); - printf(" EASYTIER_FFI_UDP_TARGET=10.0.0.2:9000 async UDP bind/send_to/recv_from\n"); -} - -int main(void) { - const char *config = getenv("EASYTIER_FFI_CONFIG"); - const char *instance = getenv("EASYTIER_FFI_INSTANCE"); - if (!config || !instance) { - print_usage(); - return 0; - } - - if (run_network_instance(config) != 0) { - print_last_error("run_network_instance failed"); - return 1; - } - printf("network instance started: %s\n", instance); - - int failed = 0; - const char *target = getenv("EASYTIER_FFI_TARGET"); - if (target) { - failed |= run_tcp_connect_demo(instance, target) != 0; - } - - const char *listen_port = getenv("EASYTIER_FFI_LISTEN_PORT"); - if (listen_port) { - failed |= run_tcp_listen_demo(instance, listen_port) != 0; - } - - const char *udp_target = getenv("EASYTIER_FFI_UDP_TARGET"); - if (udp_target) { - failed |= run_udp_demo(instance, udp_target) != 0; - } - - if (!target && !listen_port && !udp_target) { - printf("No dataplane demo env var was set; nothing else to run.\n"); - print_usage(); - } - - return failed ? 1 : 0; -} diff --git a/easytier-contrib/easytier-ffi/examples/go/README.md b/easytier-contrib/easytier-ffi/examples/go/README.md deleted file mode 100644 index bd17f59e..00000000 --- a/easytier-contrib/easytier-ffi/examples/go/README.md +++ /dev/null @@ -1,138 +0,0 @@ -# 1. Go FFI Demo - -This demo wraps EasyTier FFI data-plane TCP as Go `net.Conn` and `net.Listener`. -It can connect to an SSH server through EasyTier and read its banner, or accept a -TCP connection from another EasyTier peer and run a small ping/pong exchange. -The async op-handle wrapper is in `easytier_async.go`; the original synchronous -wrapper stays in `easytier.go`. - -## 1.1. Build the FFI library - -Run from the repository root: - -```sh -cargo build -p easytier-ffi --features ffi-dataplane -``` - -The demo loads the debug library by default: - -```text -target/debug/libeasytier_ffi.so -``` - -To use another library path, export `EASYTIER_FFI_LIB=/path/to/libeasytier_ffi.so`. - -## 1.2. Configure the EasyTier config - -`EASYTIER_FFI_CONFIG` is a string of the EasyTier config in TOML format which is passed to the FFI library. For example: - -```sh -export EASYTIER_FFI_CONFIG='instance_name = "default" -ipv4 = "10.0.0.1" - -[network_identity] -network_name = "testnet" -network_secret = "mysecret" - -[flags] -no_tun = true # disable tun device to avoid permission issues. -bind_device = false # allow loopback peers in local examples. - -[[peer]] -uri = "tcp://123.123.123.123:11010" -' -``` - -You should configure with your own real values. - -Set the local instance name and a SSH server target to connect through EasyTier: - -```sh -export EASYTIER_FFI_INSTANCE=default -export EASYTIER_FFI_TARGET=10.0.0.2:22 -``` - -To run the TCP listen integration test in the same `go test` process as the SSH -test, use a separate instance name and config: - -```sh -export EASYTIER_FFI_LISTEN_CONFIG='instance_name = "listener" -ipv4 = "10.0.0.3" - -[network_identity] -network_name = "testnet" -network_secret = "mysecret" - -[flags] -no_tun = true -bind_device = false - -[[peer]] -uri = "tcp://123.123.123.123:11010" -' -export EASYTIER_FFI_LISTEN_INSTANCE=listener -export EASYTIER_FFI_LISTEN_PORT=12345 -``` - -## 1.3. Run the demo - -`goffi` is built without cgo on Linux, so run the tests with `CGO_ENABLED=0`: - -```sh -cd easytier-contrib/easytier-ffi/examples/go -CGO_ENABLED=0 go test -v ./... -``` - -The synchronous tests use the environment variables above. The async Go tests -are self-contained: they start two local EasyTier instances in the same test -process with `no_tun = true` and `bind_device = false`, then run TCP and UDP -ping/pong over the async data-plane API. - -The synchronous wrapper also exposes `CallJSONRPC(service, method, domain, -payload)` for non-lifecycle EasyTier RPCs. For example, -`CallJSONRPC("api.logger.LoggerRpcService", "get_logger_config", "", "{}")` -returns the logger config as protobuf JSON. Instance lifecycle management RPCs -are intentionally filtered; use the dedicated FFI APIs for starting and -stopping instances. - -To run only the async tests: - -```sh -cd easytier-contrib/easytier-ffi/examples/go -CGO_ENABLED=0 go test -run 'TestAsync' -v ./... -``` - -When the SSH integration environment variables are set, expected synchronous -test output includes an SSH banner similar to: - -```text -attempt 1: got banner "SSH-2.0-..." -PASS -``` - -For `TestTCPListenIntegration`, connect from another EasyTier peer to the local -EasyTier IPv4 address and `EASYTIER_FFI_LISTEN_PORT`, send `ping`, and expect -`pong` in response. - -The async test output should include local TCP bind/connect log lines and finish -with `PASS` without any extra environment variables. - -## 1.4. C async example - -The C async example is kept separate from the basic C example: - -```sh -cargo build -p easytier-ffi --features ffi-dataplane -cc -Wall -Wextra -pedantic \ - ../example_data_plane_async.c \ - -L ../../../../target/debug -leasytier_ffi \ - -Wl,-rpath,../../../../target/debug \ - -o /tmp/easytier_data_plane_async - -/tmp/easytier_data_plane_async -``` - -Without environment variables it prints usage and exits successfully. With -`EASYTIER_FFI_CONFIG`, `EASYTIER_FFI_INSTANCE`, and one of -`EASYTIER_FFI_TARGET`, `EASYTIER_FFI_LISTEN_PORT`, or `EASYTIER_FFI_UDP_TARGET`, -it runs the corresponding async data-plane flow. diff --git a/easytier-contrib/easytier-ffi/examples/go/easytier.go b/easytier-contrib/easytier-ffi/examples/go/easytier.go deleted file mode 100644 index 6ac74b97..00000000 --- a/easytier-contrib/easytier-ffi/examples/go/easytier.go +++ /dev/null @@ -1,593 +0,0 @@ -package easytierffi - -import ( - "context" - "errors" - "fmt" - "io" - "net" - "os" - "runtime" - "strconv" - "strings" - "sync/atomic" - "time" - "unsafe" - - "github.com/go-webgpu/goffi/ffi" - "github.com/go-webgpu/goffi/types" -) - -const defaultTimeout = 30 * time.Second - -type Native struct { - lib unsafe.Pointer - - runNetworkInstance symCall - callJSONRPC symCall - getErrorMsg symCall - freeString symCall - tcpConnect symCall - tcpBind symCall - tcpAccept symCall - tcpRead symCall - tcpWrite symCall - tcpClose symCall - tcpListenerClose symCall -} - -type Conn struct { - native *Native - handle uint64 - local net.Addr - remote net.Addr - closed atomic.Bool - rd atomicDeadline - wd atomicDeadline -} - -type Listener struct { - native *Native - handle uint64 - addr net.Addr - closed atomic.Bool -} - -type symCall struct { - fn unsafe.Pointer - cif types.CallInterface -} - -type atomicDeadline struct{ v atomic.Int64 } - -type timeoutError string - -func Open(path string) (*Native, error) { - lib, err := ffi.LoadLibrary(path) - if err != nil { - return nil, err - } - n := &Native{lib: lib} - if err := n.bind(); err != nil { - ffi.FreeLibrary(lib) - return nil, err - } - return n, nil -} - -func (n *Native) Close() error { - if n.lib == nil { - return nil - } - ffi.FreeLibrary(n.lib) - n.lib = nil - return nil -} - -func (n *Native) RunNetworkInstance(config string) error { - defer pinErrorThread()() - cfg := cString(config) - cfgPtr := unsafe.Pointer(&cfg[0]) - var ret int32 - err := n.runNetworkInstance.call(unsafe.Pointer(&ret), unsafe.Pointer(&cfgPtr)) - runtime.KeepAlive(cfg) - if err != nil { - return err - } - if ret != 0 { - return n.lastError() - } - return nil -} - -func (n *Native) CallJSONRPC(serviceName, methodName, domainName, payloadJSON string) (string, error) { - defer pinErrorThread()() - service := cString(serviceName) - method := cString(methodName) - payload := cString(payloadJSON) - servicePtr := unsafe.Pointer(&service[0]) - methodPtr := unsafe.Pointer(&method[0]) - payloadPtr := unsafe.Pointer(&payload[0]) - var domain []byte - var domainPtr unsafe.Pointer - if domainName != "" { - domain = cString(domainName) - domainPtr = unsafe.Pointer(&domain[0]) - } - var response unsafe.Pointer - responseArg := unsafe.Pointer(&response) - var ret int32 - err := n.callJSONRPC.call( - unsafe.Pointer(&ret), - unsafe.Pointer(&servicePtr), - unsafe.Pointer(&methodPtr), - unsafe.Pointer(&domainPtr), - unsafe.Pointer(&payloadPtr), - unsafe.Pointer(&responseArg), - ) - runtime.KeepAlive(service) - runtime.KeepAlive(method) - runtime.KeepAlive(domain) - runtime.KeepAlive(payload) - if err != nil { - return "", err - } - if ret != 0 { - return "", n.lastError() - } - if response == nil { - return "", errors.New("easytier ffi JSON RPC returned nil response") - } - defer func() { _ = n.freeCString(response) }() - return readCString(response), nil -} - -func (n *Native) DialContext(ctx context.Context, instance, network, address string) (net.Conn, error) { - if network != "tcp" && network != "tcp4" && network != "tcp6" { - return nil, net.UnknownNetworkError(network) - } - ip, port, err := parseIPPort(address) - if err != nil { - return nil, err - } - timeout := defaultTimeout - if deadline, ok := ctx.Deadline(); ok { - timeout = time.Until(deadline) - } - if timeout <= 0 { - return nil, context.DeadlineExceeded - } - if err := ctx.Err(); err != nil { - return nil, err - } - handle, local, err := n.tcpConnectTo(instance, ip.String(), uint16(port), timeout) - if err != nil { - return nil, err - } - return &Conn{native: n, handle: handle, local: local, remote: &net.TCPAddr{IP: ip, Port: port}}, nil -} - -func (n *Native) ListenContext(ctx context.Context, instance, network, address string) (net.Listener, error) { - if network != "tcp" && network != "tcp4" && network != "tcp6" { - return nil, net.UnknownNetworkError(network) - } - port, err := parseListenPort(address) - if err != nil { - return nil, err - } - timeout := defaultTimeout - if deadline, ok := ctx.Deadline(); ok { - timeout = time.Until(deadline) - } - if timeout <= 0 { - return nil, context.DeadlineExceeded - } - if err := ctx.Err(); err != nil { - return nil, err - } - handle, local, err := n.tcpBindTo(instance, uint16(port), timeout) - if err != nil { - return nil, err - } - return &Listener{native: n, handle: handle, addr: local}, nil -} - -func (c *Conn) Read(b []byte) (int, error) { - if c.closed.Load() { - return 0, net.ErrClosed - } - n, err := c.native.tcpReadFrom(c.handle, b, c.rd.timeout(defaultTimeout)) - if err != nil { - return 0, opError("read", c.remote, err) - } - if n == 0 { - return 0, io.EOF - } - return n, nil -} - -func (c *Conn) Write(b []byte) (int, error) { - if c.closed.Load() { - return 0, net.ErrClosed - } - n, err := c.native.tcpWriteTo(c.handle, b, c.wd.timeout(defaultTimeout)) - if err != nil { - return 0, opError("write", c.remote, err) - } - return n, nil -} - -func (c *Conn) Close() error { - if !c.closed.CompareAndSwap(false, true) { - return net.ErrClosed - } - return c.native.tcpCloseHandle(c.handle) -} - -func (c *Conn) LocalAddr() net.Addr { return c.local } -func (c *Conn) RemoteAddr() net.Addr { return c.remote } -func (c *Conn) SetDeadline(t time.Time) error { c.rd.set(t); c.wd.set(t); return nil } -func (c *Conn) SetReadDeadline(t time.Time) error { c.rd.set(t); return nil } -func (c *Conn) SetWriteDeadline(t time.Time) error { c.wd.set(t); return nil } - -func (l *Listener) Accept() (net.Conn, error) { - if l.closed.Load() { - return nil, net.ErrClosed - } - for { - handle, local, peer, err := l.native.tcpAcceptFrom(l.handle, defaultTimeout) - if err == nil { - return &Conn{native: l.native, handle: handle, local: local, remote: peer}, nil - } - if l.closed.Load() { - return nil, net.ErrClosed - } - var netErr net.Error - if errors.As(err, &netErr) && netErr.Timeout() { - continue - } - return nil, opError("accept", l.addr, err) - } -} - -func (l *Listener) Close() error { - if !l.closed.CompareAndSwap(false, true) { - return net.ErrClosed - } - return l.native.tcpListenerCloseHandle(l.handle) -} - -func (l *Listener) Addr() net.Addr { return l.addr } - -func (n *Native) bind() error { - return errors.Join( - n.bindSym(&n.runNetworkInstance, "run_network_instance", types.SInt32TypeDescriptor, types.PointerTypeDescriptor), - n.bindSym(&n.callJSONRPC, "call_json_rpc", types.SInt32TypeDescriptor, types.PointerTypeDescriptor, types.PointerTypeDescriptor, types.PointerTypeDescriptor, types.PointerTypeDescriptor, types.PointerTypeDescriptor), - n.bindSym(&n.getErrorMsg, "get_error_msg", types.VoidTypeDescriptor, types.PointerTypeDescriptor), - n.bindSym(&n.freeString, "free_string", types.VoidTypeDescriptor, types.PointerTypeDescriptor), - n.bindSym(&n.tcpConnect, "data_plane_tcp_connect", types.UInt64TypeDescriptor, types.PointerTypeDescriptor, types.PointerTypeDescriptor, types.UInt16TypeDescriptor, types.UInt64TypeDescriptor, types.PointerTypeDescriptor, types.PointerTypeDescriptor), - n.bindSym(&n.tcpBind, "data_plane_tcp_bind", types.UInt64TypeDescriptor, types.PointerTypeDescriptor, types.UInt16TypeDescriptor, types.UInt64TypeDescriptor, types.PointerTypeDescriptor, types.PointerTypeDescriptor), - n.bindSym(&n.tcpAccept, "data_plane_tcp_accept", types.UInt64TypeDescriptor, types.UInt64TypeDescriptor, types.UInt64TypeDescriptor, types.PointerTypeDescriptor, types.PointerTypeDescriptor, types.PointerTypeDescriptor, types.PointerTypeDescriptor), - n.bindSym(&n.tcpRead, "data_plane_tcp_read", types.SInt32TypeDescriptor, types.UInt64TypeDescriptor, types.PointerTypeDescriptor, types.UInt32TypeDescriptor, types.UInt64TypeDescriptor), - n.bindSym(&n.tcpWrite, "data_plane_tcp_write", types.SInt32TypeDescriptor, types.UInt64TypeDescriptor, types.PointerTypeDescriptor, types.UInt32TypeDescriptor, types.UInt64TypeDescriptor), - n.bindSym(&n.tcpClose, "data_plane_tcp_close", types.SInt32TypeDescriptor, types.UInt64TypeDescriptor), - n.bindSym(&n.tcpListenerClose, "data_plane_tcp_listener_close", types.SInt32TypeDescriptor, types.UInt64TypeDescriptor), - ) -} - -func (n *Native) bindSym(dst *symCall, name string, ret *types.TypeDescriptor, args ...*types.TypeDescriptor) error { - sym, err := ffi.GetSymbol(n.lib, name) - if err != nil { - return err - } - if err := ffi.PrepareCallInterface(&dst.cif, types.DefaultCall, ret, args); err != nil { - return err - } - dst.fn = sym - return nil -} - -func (s *symCall) call(ret unsafe.Pointer, args ...unsafe.Pointer) error { - // `ffi.CallFunction` and libffi `ffi_call` are safe to invoke concurrently - // because `cif` is prepared once during binding and only read afterwards. - return ffi.CallFunction(&s.cif, s.fn, ret, args) -} - -func (n *Native) tcpConnectTo(instance, ip string, port uint16, timeout time.Duration) (uint64, *net.TCPAddr, error) { - defer pinErrorThread()() - inst := cString(instance) - dst := cString(ip) - instPtr := unsafe.Pointer(&inst[0]) - dstPtr := unsafe.Pointer(&dst[0]) - timeoutMS := uint64(timeout / time.Millisecond) - var handle uint64 - var outIP unsafe.Pointer - outIPArg := unsafe.Pointer(&outIP) - var outPort uint16 - outPortArg := unsafe.Pointer(&outPort) - err := n.tcpConnect.call( - unsafe.Pointer(&handle), - unsafe.Pointer(&instPtr), - unsafe.Pointer(&dstPtr), - unsafe.Pointer(&port), - unsafe.Pointer(&timeoutMS), - unsafe.Pointer(&outIPArg), - unsafe.Pointer(&outPortArg), - ) - runtime.KeepAlive(inst) - runtime.KeepAlive(dst) - if err != nil { - return 0, nil, err - } - if handle == 0 { - return 0, nil, n.lastError() - } - return handle, n.takeTCPAddr(outIP, outPort), nil -} - -func (n *Native) tcpBindTo(instance string, port uint16, timeout time.Duration) (uint64, *net.TCPAddr, error) { - defer pinErrorThread()() - inst := cString(instance) - instPtr := unsafe.Pointer(&inst[0]) - timeoutMS := uint64(timeout / time.Millisecond) - var handle uint64 - var outIP unsafe.Pointer - outIPArg := unsafe.Pointer(&outIP) - var outPort uint16 - outPortArg := unsafe.Pointer(&outPort) - err := n.tcpBind.call( - unsafe.Pointer(&handle), - unsafe.Pointer(&instPtr), - unsafe.Pointer(&port), - unsafe.Pointer(&timeoutMS), - unsafe.Pointer(&outIPArg), - unsafe.Pointer(&outPortArg), - ) - runtime.KeepAlive(inst) - if err != nil { - return 0, nil, err - } - if handle == 0 { - return 0, nil, n.lastError() - } - return handle, n.takeTCPAddr(outIP, outPort), nil -} - -func (n *Native) tcpAcceptFrom(handle uint64, timeout time.Duration) (uint64, *net.TCPAddr, *net.TCPAddr, error) { - defer pinErrorThread()() - timeoutMS := uint64(timeout / time.Millisecond) - var stream uint64 - var outLocalIP unsafe.Pointer - outLocalIPArg := unsafe.Pointer(&outLocalIP) - var outLocalPort uint16 - outLocalPortArg := unsafe.Pointer(&outLocalPort) - var outPeerIP unsafe.Pointer - outPeerIPArg := unsafe.Pointer(&outPeerIP) - var outPeerPort uint16 - outPeerPortArg := unsafe.Pointer(&outPeerPort) - err := n.tcpAccept.call( - unsafe.Pointer(&stream), - unsafe.Pointer(&handle), - unsafe.Pointer(&timeoutMS), - unsafe.Pointer(&outLocalIPArg), - unsafe.Pointer(&outLocalPortArg), - unsafe.Pointer(&outPeerIPArg), - unsafe.Pointer(&outPeerPortArg), - ) - if err != nil { - return 0, nil, nil, err - } - if stream == 0 { - return 0, nil, nil, n.lastError() - } - return stream, n.takeTCPAddr(outLocalIP, outLocalPort), n.takeTCPAddr(outPeerIP, outPeerPort), nil -} - -func (n *Native) tcpReadFrom(handle uint64, buf []byte, timeout time.Duration) (int, error) { - if len(buf) == 0 { - return 0, nil - } - defer pinErrorThread()() - var ret int32 - bufPtr := unsafe.Pointer(&buf[0]) - length := uint32(len(buf)) - timeoutMS := uint64(timeout / time.Millisecond) - err := n.tcpRead.call(unsafe.Pointer(&ret), unsafe.Pointer(&handle), unsafe.Pointer(&bufPtr), unsafe.Pointer(&length), unsafe.Pointer(&timeoutMS)) - runtime.KeepAlive(buf) - if err != nil { - return 0, err - } - if ret < 0 { - return 0, n.lastError() - } - return int(ret), nil -} - -func (n *Native) tcpWriteTo(handle uint64, buf []byte, timeout time.Duration) (int, error) { - if len(buf) == 0 { - return 0, nil - } - defer pinErrorThread()() - var ret int32 - bufPtr := unsafe.Pointer(&buf[0]) - length := uint32(len(buf)) - timeoutMS := uint64(timeout / time.Millisecond) - err := n.tcpWrite.call(unsafe.Pointer(&ret), unsafe.Pointer(&handle), unsafe.Pointer(&bufPtr), unsafe.Pointer(&length), unsafe.Pointer(&timeoutMS)) - runtime.KeepAlive(buf) - if err != nil { - return 0, err - } - if ret < 0 { - return 0, n.lastError() - } - return int(ret), nil -} - -func (n *Native) tcpCloseHandle(handle uint64) error { - defer pinErrorThread()() - var ret int32 - if err := n.tcpClose.call(unsafe.Pointer(&ret), unsafe.Pointer(&handle)); err != nil { - return err - } - if ret != 0 { - return n.lastError() - } - return nil -} - -func (n *Native) tcpListenerCloseHandle(handle uint64) error { - defer pinErrorThread()() - var ret int32 - if err := n.tcpListenerClose.call(unsafe.Pointer(&ret), unsafe.Pointer(&handle)); err != nil { - return err - } - if ret != 0 { - return n.lastError() - } - return nil -} - -// pinErrorThread ties an FFI op to the get_error_msg that reads its result: the -// Rust side stores the last error in a thread-local, so the goroutine must not -// migrate to another OS thread between the two calls. Use as `defer pinErrorThread()()` -// at the start of any wrapper that reports failures through lastError. -func pinErrorThread() func() { - runtime.LockOSThread() - return runtime.UnlockOSThread -} - -func (n *Native) lastError() error { - var out unsafe.Pointer - outArg := unsafe.Pointer(&out) - if err := n.getErrorMsg.call(nil, unsafe.Pointer(&outArg)); err != nil { - return err - } - if out == nil { - return errors.New("easytier ffi call failed") - } - msg := readCString(out) - _ = n.freeCString(out) - if strings.Contains(msg, "timed out") { - return timeoutError(msg) - } - return errors.New(msg) -} - -func (n *Native) freeCString(ptr unsafe.Pointer) error { - if ptr == nil { - return nil - } - return n.freeString.call(nil, unsafe.Pointer(&ptr)) -} - -func (n *Native) takeTCPAddr(ipPtr unsafe.Pointer, port uint16) *net.TCPAddr { - if ipPtr == nil { - return nil - } - ip := net.ParseIP(readCString(ipPtr)) - _ = n.freeCString(ipPtr) - return &net.TCPAddr{IP: ip, Port: int(port)} -} - -func (d *atomicDeadline) set(t time.Time) { - if t.IsZero() { - d.v.Store(0) - return - } - d.v.Store(t.UnixNano()) -} - -func (d *atomicDeadline) timeout(fallback time.Duration) time.Duration { - ns := d.v.Load() - if ns == 0 { - return fallback - } - remaining := time.Until(time.Unix(0, ns)) - if remaining <= 0 { - return time.Millisecond - } - return remaining -} - -func (e timeoutError) Error() string { return string(e) } -func (e timeoutError) Timeout() bool { return true } -func (e timeoutError) Temporary() bool { return true } - -func opError(op string, addr net.Addr, err error) error { - return &net.OpError{Op: op, Net: "easytier", Addr: addr, Err: err} -} - -func parseIPPort(address string) (net.IP, int, error) { - host, portStr, err := net.SplitHostPort(address) - if err != nil { - return nil, 0, err - } - ip := net.ParseIP(host) - if ip == nil { - return nil, 0, fmt.Errorf("easytier ffi requires an IP address, got %q", host) - } - port, err := strconv.ParseUint(portStr, 10, 16) - if err != nil { - return nil, 0, err - } - return ip, int(port), nil -} - -func parseListenPort(address string) (int, error) { - host, portStr, err := net.SplitHostPort(address) - if err != nil { - return 0, err - } - if host != "" { - ip := net.ParseIP(host) - if ip == nil { - return 0, fmt.Errorf("easytier ffi requires an IP address, got %q", host) - } - if !ip.IsUnspecified() { - return 0, fmt.Errorf("easytier ffi listen address must be unspecified, got %q", host) - } - } - port, err := strconv.ParseUint(portStr, 10, 16) - if err != nil { - return 0, err - } - return int(port), nil -} - -func cString(s string) []byte { - if strings.ContainsRune(s, 0) { - panic("easytier ffi string contains NUL") - } - return append([]byte(s), 0) -} - -func readCString(ptr unsafe.Pointer) string { - if ptr == nil { - return "" - } - var b []byte - for p := uintptr(ptr); ; p++ { - c := *(*byte)(unsafe.Pointer(p)) - if c == 0 { - return string(b) - } - b = append(b, c) - } -} - -func defaultLibraryPath() string { - if p := os.Getenv("EASYTIER_FFI_LIB"); p != "" { - return p - } - switch runtime.GOOS { - case "darwin": - return "../../../../target/debug/libeasytier_ffi.dylib" - case "windows": - return "..\\..\\..\\..\\target\\debug\\easytier_ffi.dll" - default: - return "../../../../target/debug/libeasytier_ffi.so" - } -} - -var _ net.Conn = (*Conn)(nil) -var _ net.Listener = (*Listener)(nil) diff --git a/easytier-contrib/easytier-ffi/examples/go/easytier_async.go b/easytier-contrib/easytier-ffi/examples/go/easytier_async.go deleted file mode 100644 index 232b9667..00000000 --- a/easytier-contrib/easytier-ffi/examples/go/easytier_async.go +++ /dev/null @@ -1,1018 +0,0 @@ -// Package easytierffi contains a small Go wrapper around the EasyTier FFI -// examples. This file documents the async data-plane surface; the synchronous -// wrapper lives in easytier.go. -// -// Public async entry points: -// -// - OpenAsync(path) loads the EasyTier FFI dynamic library and binds the -// async dataplane symbols. Close releases only the dynamic library handle; -// network instances started through RunNetworkInstance are process-global -// EasyTier state. -// -// - (*AsyncNative).RunNetworkInstance(config) starts one EasyTier instance -// from TOML. The instance name in the config is used by all dataplane calls. -// -// - (*AsyncNative).DialContext(ctx, instance, "tcp", "ip:port") starts an -// async TCP connect and returns an AsyncConn implementing net.Conn. -// -// - (*AsyncNative).ListenContext(ctx, instance, "tcp", "0.0.0.0:port") starts -// an async TCP bind and returns an AsyncListener implementing net.Listener. -// -// - AsyncConn implements net.Conn. Read and Write each start one native async -// read/write op and wait for completion. Deadlines are mapped to operation -// timeouts. Close closes the underlying dataplane stream handle. -// -// - AsyncListener implements net.Listener. Accept starts one native async -// accept op and waits for a stream. Close closes the listener handle. -// -// - (*AsyncNative).UDPBindContext(ctx, instance, port) returns an -// AsyncUDPSocket. AsyncUDPSocket.SendTo and RecvFrom start one native async -// UDP send/receive op and wait for completion. Close closes the socket -// handle. -// -// - TCPConnectContext/TCPBindContext/TCPAcceptContext/TCPReadContext/ -// TCPWriteContext and UDPSendToContext/UDPRecvFromContext are lower-level -// handle helpers used by the examples and tests. External callers should -// prefer DialContext, ListenContext, AsyncConn, AsyncListener, and -// AsyncUDPSocket because raw handle close helpers are intentionally internal -// to this example package. -// -// Async operation semantics: -// -// - Each Context method starts a native async op, polls data_plane_async_op_wait -// in short intervals, then calls the matching finish function. Finish is -// single-consume on the native side. -// -// - If the context is canceled or its deadline expires before completion, the -// wrapper cancels and frees the native op and returns the context error. -// -// - Read and RecvFrom copy Rust-owned output buffers into Go slices and free -// the native allocation before returning. -// -// - Write and SendTo keep the Go input buffer alive for the start call. The -// native async API copies the input buffer during start, so callers do not -// need to keep it alive after the Go method returns. -// -// - FFI calls that read the Rust thread-local error string pin the goroutine -// to one OS thread from the failing call through get_error_msg. -package easytierffi - -import ( - "context" - "errors" - "fmt" - "io" - "net" - "runtime" - "strings" - "sync/atomic" - "time" - "unsafe" - - "github.com/go-webgpu/goffi/ffi" - "github.com/go-webgpu/goffi/types" -) - -const ( - dataPlaneOpPending = int32(0) - dataPlaneOpReady = int32(1) - dataPlaneOpFailed = int32(-1) - dataPlaneOpInvalid = int32(-2) - - asyncPollInterval = 50 * time.Millisecond -) - -type AsyncNative struct { - lib unsafe.Pointer - - runNetworkInstance symCall - deleteNetworkInst symCall - getErrorMsg symCall - freeString symCall - freeBytes symCall - - asyncOpStatus symCall - asyncOpWait symCall - asyncOpCancel symCall - asyncOpFree symCall - - tcpConnectStart symCall - tcpConnectFinish symCall - tcpBindStart symCall - tcpBindFinish symCall - tcpAcceptStart symCall - tcpAcceptFinish symCall - tcpReadStart symCall - tcpReadFinish symCall - tcpWriteStart symCall - tcpWriteFinish symCall - tcpClose symCall - tcpListenerClose symCall - - udpBindStart symCall - udpBindFinish symCall - udpSendToStart symCall - udpSendToFinish symCall - udpRecvFromStart symCall - udpRecvFromFinish symCall - udpClose symCall -} - -type AsyncConn struct { - native *AsyncNative - handle uint64 - local net.Addr - remote net.Addr - closed atomicBool - rd atomicDeadline - wd atomicDeadline -} - -type AsyncListener struct { - native *AsyncNative - handle uint64 - addr net.Addr - closed atomicBool -} - -type AsyncUDPSocket struct { - native *AsyncNative - handle uint64 - addr *net.UDPAddr - closed atomicBool -} - -type atomicBool struct{ v atomic.Bool } - -func OpenAsync(path string) (*AsyncNative, error) { - lib, err := ffi.LoadLibrary(path) - if err != nil { - return nil, err - } - n := &AsyncNative{lib: lib} - if err := n.bind(); err != nil { - ffi.FreeLibrary(lib) - return nil, err - } - return n, nil -} - -func (n *AsyncNative) Close() error { - if n.lib == nil { - return nil - } - ffi.FreeLibrary(n.lib) - n.lib = nil - return nil -} - -func (n *AsyncNative) RunNetworkInstance(config string) error { - defer pinErrorThread()() - cfg := cString(config) - cfgPtr := unsafe.Pointer(&cfg[0]) - var ret int32 - err := n.runNetworkInstance.call(unsafe.Pointer(&ret), unsafe.Pointer(&cfgPtr)) - runtime.KeepAlive(cfg) - if err != nil { - return err - } - if ret != 0 { - return n.lastError() - } - return nil -} - -func (n *AsyncNative) deleteNetworkInstances(names []string) error { - defer pinErrorThread()() - - cNames := make([][]byte, len(names)) - namePtrs := make([]unsafe.Pointer, len(names)) - for i, name := range names { - cNames[i] = cString(name) - namePtrs[i] = unsafe.Pointer(&cNames[i][0]) - } - - var namesPtr unsafe.Pointer - if len(namePtrs) > 0 { - namesPtr = unsafe.Pointer(&namePtrs[0]) - } - length := uint64(len(names)) - var ret int32 - err := n.deleteNetworkInst.call(unsafe.Pointer(&ret), unsafe.Pointer(&namesPtr), unsafe.Pointer(&length)) - runtime.KeepAlive(cNames) - runtime.KeepAlive(namePtrs) - if err != nil { - return err - } - if ret != 0 { - return n.lastError() - } - return nil -} - -func (n *AsyncNative) DialContext(ctx context.Context, instance, network, address string) (net.Conn, error) { - if network != "tcp" && network != "tcp4" && network != "tcp6" { - return nil, net.UnknownNetworkError(network) - } - ip, port, err := parseIPPort(address) - if err != nil { - return nil, err - } - handle, local, err := n.TCPConnectContext(ctx, instance, ip.String(), uint16(port)) - if err != nil { - return nil, err - } - return &AsyncConn{ - native: n, - handle: handle, - local: local, - remote: &net.TCPAddr{IP: ip, Port: port}, - }, nil -} - -func (n *AsyncNative) ListenContext(ctx context.Context, instance, network, address string) (net.Listener, error) { - if network != "tcp" && network != "tcp4" && network != "tcp6" { - return nil, net.UnknownNetworkError(network) - } - port, err := parseListenPort(address) - if err != nil { - return nil, err - } - handle, local, err := n.TCPBindContext(ctx, instance, uint16(port)) - if err != nil { - return nil, err - } - return &AsyncListener{native: n, handle: handle, addr: local}, nil -} - -func (n *AsyncNative) TCPConnectContext(ctx context.Context, instance, ip string, port uint16) (uint64, *net.TCPAddr, error) { - timeout, err := contextTimeout(ctx) - if err != nil { - return 0, nil, err - } - op, err := n.tcpConnectStartCall(instance, ip, port, timeout) - if err != nil { - return 0, nil, err - } - if err := n.waitOp(ctx, op); err != nil { - return 0, nil, err - } - return n.tcpConnectFinishCall(op) -} - -func (n *AsyncNative) TCPBindContext(ctx context.Context, instance string, port uint16) (uint64, *net.TCPAddr, error) { - timeout, err := contextTimeout(ctx) - if err != nil { - return 0, nil, err - } - op, err := n.tcpBindStartCall(instance, port, timeout) - if err != nil { - return 0, nil, err - } - if err := n.waitOp(ctx, op); err != nil { - return 0, nil, err - } - return n.tcpBindFinishCall(op) -} - -func (n *AsyncNative) TCPAcceptContext(ctx context.Context, listener uint64) (uint64, *net.TCPAddr, *net.TCPAddr, error) { - timeout, err := contextTimeout(ctx) - if err != nil { - return 0, nil, nil, err - } - op, err := n.tcpAcceptStartCall(listener, timeout) - if err != nil { - return 0, nil, nil, err - } - if err := n.waitOp(ctx, op); err != nil { - return 0, nil, nil, err - } - return n.tcpAcceptFinishCall(op) -} - -func (n *AsyncNative) TCPReadContext(ctx context.Context, stream uint64, maxLen uint32) ([]byte, error) { - timeout, err := contextTimeout(ctx) - if err != nil { - return nil, err - } - op, err := n.tcpReadStartCall(stream, maxLen, timeout) - if err != nil { - return nil, err - } - if err := n.waitOp(ctx, op); err != nil { - return nil, err - } - return n.tcpReadFinishCall(op) -} - -func (n *AsyncNative) TCPWriteContext(ctx context.Context, stream uint64, data []byte) (int, error) { - timeout, err := contextTimeout(ctx) - if err != nil { - return 0, err - } - op, err := n.tcpWriteStartCall(stream, data, timeout) - if err != nil { - return 0, err - } - if err := n.waitOp(ctx, op); err != nil { - return 0, err - } - return n.tcpWriteFinishCall(op) -} - -func (n *AsyncNative) UDPBindContext(ctx context.Context, instance string, port uint16) (*AsyncUDPSocket, error) { - timeout, err := contextTimeout(ctx) - if err != nil { - return nil, err - } - op, err := n.udpBindStartCall(instance, port, timeout) - if err != nil { - return nil, err - } - if err := n.waitOp(ctx, op); err != nil { - return nil, err - } - handle, local, err := n.udpBindFinishCall(op) - if err != nil { - return nil, err - } - return &AsyncUDPSocket{native: n, handle: handle, addr: local}, nil -} - -func (n *AsyncNative) UDPSendToContext(ctx context.Context, socket uint64, addr *net.UDPAddr, data []byte) (int, error) { - timeout, err := contextTimeout(ctx) - if err != nil { - return 0, err - } - op, err := n.udpSendToStartCall(socket, addr, data, timeout) - if err != nil { - return 0, err - } - if err := n.waitOp(ctx, op); err != nil { - return 0, err - } - return n.udpSendToFinishCall(op) -} - -func (n *AsyncNative) UDPRecvFromContext(ctx context.Context, socket uint64, maxLen uint32) ([]byte, *net.UDPAddr, error) { - timeout, err := contextTimeout(ctx) - if err != nil { - return nil, nil, err - } - op, err := n.udpRecvFromStartCall(socket, maxLen, timeout) - if err != nil { - return nil, nil, err - } - if err := n.waitOp(ctx, op); err != nil { - return nil, nil, err - } - return n.udpRecvFromFinishCall(op) -} - -func (c *AsyncConn) Read(b []byte) (int, error) { - if c.closed.Load() { - return 0, net.ErrClosed - } - if len(b) == 0 { - return 0, nil - } - ctx, cancel := context.WithTimeout(context.Background(), c.rd.timeout(defaultTimeout)) - defer cancel() - data, err := c.native.TCPReadContext(ctx, c.handle, uint32(len(b))) - if err != nil { - return 0, opError("read", c.remote, err) - } - if len(data) == 0 { - return 0, io.EOF - } - return copy(b, data), nil -} - -func (c *AsyncConn) Write(b []byte) (int, error) { - if c.closed.Load() { - return 0, net.ErrClosed - } - ctx, cancel := context.WithTimeout(context.Background(), c.wd.timeout(defaultTimeout)) - defer cancel() - n, err := c.native.TCPWriteContext(ctx, c.handle, b) - if err != nil { - return 0, opError("write", c.remote, err) - } - return n, nil -} - -func (c *AsyncConn) Close() error { - if !c.closed.CompareAndSwap(false, true) { - return net.ErrClosed - } - return c.native.tcpCloseHandle(c.handle) -} - -func (c *AsyncConn) LocalAddr() net.Addr { return c.local } -func (c *AsyncConn) RemoteAddr() net.Addr { return c.remote } -func (c *AsyncConn) SetDeadline(t time.Time) error { c.rd.set(t); c.wd.set(t); return nil } -func (c *AsyncConn) SetReadDeadline(t time.Time) error { c.rd.set(t); return nil } -func (c *AsyncConn) SetWriteDeadline(t time.Time) error { c.wd.set(t); return nil } - -func (l *AsyncListener) Accept() (net.Conn, error) { - if l.closed.Load() { - return nil, net.ErrClosed - } - for { - ctx, cancel := context.WithTimeout(context.Background(), defaultTimeout) - handle, local, peer, err := l.native.TCPAcceptContext(ctx, l.handle) - cancel() - if err == nil { - return &AsyncConn{native: l.native, handle: handle, local: local, remote: peer}, nil - } - if l.closed.Load() { - return nil, net.ErrClosed - } - var netErr net.Error - if errors.As(err, &netErr) && netErr.Timeout() { - continue - } - return nil, opError("accept", l.addr, err) - } -} - -func (l *AsyncListener) Close() error { - if !l.closed.CompareAndSwap(false, true) { - return net.ErrClosed - } - return l.native.tcpListenerCloseHandle(l.handle) -} - -func (l *AsyncListener) Addr() net.Addr { return l.addr } - -func (s *AsyncUDPSocket) SendTo(ctx context.Context, data []byte, addr *net.UDPAddr) (int, error) { - if s.closed.Load() { - return 0, net.ErrClosed - } - return s.native.UDPSendToContext(ctx, s.handle, addr, data) -} - -func (s *AsyncUDPSocket) RecvFrom(ctx context.Context, maxLen uint32) ([]byte, *net.UDPAddr, error) { - if s.closed.Load() { - return nil, nil, net.ErrClosed - } - return s.native.UDPRecvFromContext(ctx, s.handle, maxLen) -} - -func (s *AsyncUDPSocket) Close() error { - if !s.closed.CompareAndSwap(false, true) { - return net.ErrClosed - } - return s.native.udpCloseHandle(s.handle) -} - -func (s *AsyncUDPSocket) LocalAddr() *net.UDPAddr { return s.addr } - -func (n *AsyncNative) bind() error { - return errors.Join( - n.bindSym(&n.runNetworkInstance, "run_network_instance", types.SInt32TypeDescriptor, types.PointerTypeDescriptor), - n.bindSym(&n.deleteNetworkInst, "delete_network_instance", types.SInt32TypeDescriptor, types.PointerTypeDescriptor, types.UInt64TypeDescriptor), - n.bindSym(&n.getErrorMsg, "get_error_msg", types.VoidTypeDescriptor, types.PointerTypeDescriptor), - n.bindSym(&n.freeString, "free_string", types.VoidTypeDescriptor, types.PointerTypeDescriptor), - n.bindSym(&n.freeBytes, "data_plane_free_bytes", types.VoidTypeDescriptor, types.PointerTypeDescriptor, types.UInt32TypeDescriptor), - n.bindSym(&n.asyncOpStatus, "data_plane_async_op_status", types.SInt32TypeDescriptor, types.UInt64TypeDescriptor), - n.bindSym(&n.asyncOpWait, "data_plane_async_op_wait", types.SInt32TypeDescriptor, types.UInt64TypeDescriptor, types.UInt64TypeDescriptor), - n.bindSym(&n.asyncOpCancel, "data_plane_async_op_cancel", types.SInt32TypeDescriptor, types.UInt64TypeDescriptor), - n.bindSym(&n.asyncOpFree, "data_plane_async_op_free", types.SInt32TypeDescriptor, types.UInt64TypeDescriptor), - n.bindSym(&n.tcpConnectStart, "data_plane_tcp_connect_start", types.UInt64TypeDescriptor, types.PointerTypeDescriptor, types.PointerTypeDescriptor, types.UInt16TypeDescriptor, types.UInt64TypeDescriptor), - n.bindSym(&n.tcpConnectFinish, "data_plane_tcp_connect_finish", types.UInt64TypeDescriptor, types.UInt64TypeDescriptor, types.PointerTypeDescriptor, types.PointerTypeDescriptor), - n.bindSym(&n.tcpBindStart, "data_plane_tcp_bind_start", types.UInt64TypeDescriptor, types.PointerTypeDescriptor, types.UInt16TypeDescriptor, types.UInt64TypeDescriptor), - n.bindSym(&n.tcpBindFinish, "data_plane_tcp_bind_finish", types.UInt64TypeDescriptor, types.UInt64TypeDescriptor, types.PointerTypeDescriptor, types.PointerTypeDescriptor), - n.bindSym(&n.tcpAcceptStart, "data_plane_tcp_accept_start", types.UInt64TypeDescriptor, types.UInt64TypeDescriptor, types.UInt64TypeDescriptor), - n.bindSym(&n.tcpAcceptFinish, "data_plane_tcp_accept_finish", types.UInt64TypeDescriptor, types.UInt64TypeDescriptor, types.PointerTypeDescriptor, types.PointerTypeDescriptor, types.PointerTypeDescriptor, types.PointerTypeDescriptor), - n.bindSym(&n.tcpReadStart, "data_plane_tcp_read_start", types.UInt64TypeDescriptor, types.UInt64TypeDescriptor, types.UInt32TypeDescriptor, types.UInt64TypeDescriptor), - n.bindSym(&n.tcpReadFinish, "data_plane_tcp_read_finish", types.SInt32TypeDescriptor, types.UInt64TypeDescriptor, types.PointerTypeDescriptor, types.PointerTypeDescriptor), - n.bindSym(&n.tcpWriteStart, "data_plane_tcp_write_start", types.UInt64TypeDescriptor, types.UInt64TypeDescriptor, types.PointerTypeDescriptor, types.UInt32TypeDescriptor, types.UInt64TypeDescriptor), - n.bindSym(&n.tcpWriteFinish, "data_plane_tcp_write_finish", types.SInt32TypeDescriptor, types.UInt64TypeDescriptor), - n.bindSym(&n.tcpClose, "data_plane_tcp_close", types.SInt32TypeDescriptor, types.UInt64TypeDescriptor), - n.bindSym(&n.tcpListenerClose, "data_plane_tcp_listener_close", types.SInt32TypeDescriptor, types.UInt64TypeDescriptor), - n.bindSym(&n.udpBindStart, "data_plane_udp_bind_start", types.UInt64TypeDescriptor, types.PointerTypeDescriptor, types.UInt16TypeDescriptor, types.UInt64TypeDescriptor), - n.bindSym(&n.udpBindFinish, "data_plane_udp_bind_finish", types.UInt64TypeDescriptor, types.UInt64TypeDescriptor, types.PointerTypeDescriptor, types.PointerTypeDescriptor), - n.bindSym(&n.udpSendToStart, "data_plane_udp_send_to_start", types.UInt64TypeDescriptor, types.UInt64TypeDescriptor, types.PointerTypeDescriptor, types.UInt16TypeDescriptor, types.PointerTypeDescriptor, types.UInt32TypeDescriptor, types.UInt64TypeDescriptor), - n.bindSym(&n.udpSendToFinish, "data_plane_udp_send_to_finish", types.SInt32TypeDescriptor, types.UInt64TypeDescriptor), - n.bindSym(&n.udpRecvFromStart, "data_plane_udp_recv_from_start", types.UInt64TypeDescriptor, types.UInt64TypeDescriptor, types.UInt32TypeDescriptor, types.UInt64TypeDescriptor), - n.bindSym(&n.udpRecvFromFinish, "data_plane_udp_recv_from_finish", types.SInt32TypeDescriptor, types.UInt64TypeDescriptor, types.PointerTypeDescriptor, types.PointerTypeDescriptor, types.PointerTypeDescriptor, types.PointerTypeDescriptor), - n.bindSym(&n.udpClose, "data_plane_udp_close", types.SInt32TypeDescriptor, types.UInt64TypeDescriptor), - ) -} - -func (n *AsyncNative) bindSym(dst *symCall, name string, ret *types.TypeDescriptor, args ...*types.TypeDescriptor) error { - sym, err := ffi.GetSymbol(n.lib, name) - if err != nil { - return err - } - if err := ffi.PrepareCallInterface(&dst.cif, types.DefaultCall, ret, args); err != nil { - return err - } - dst.fn = sym - return nil -} - -func (n *AsyncNative) tcpConnectStartCall(instance, ip string, port uint16, timeout time.Duration) (uint64, error) { - defer pinErrorThread()() - inst := cString(instance) - dst := cString(ip) - instPtr := unsafe.Pointer(&inst[0]) - dstPtr := unsafe.Pointer(&dst[0]) - timeoutMS := durationMillis(timeout) - var op uint64 - err := n.tcpConnectStart.call( - unsafe.Pointer(&op), - unsafe.Pointer(&instPtr), - unsafe.Pointer(&dstPtr), - unsafe.Pointer(&port), - unsafe.Pointer(&timeoutMS), - ) - runtime.KeepAlive(inst) - runtime.KeepAlive(dst) - return n.startResult(op, err) -} - -func (n *AsyncNative) tcpConnectFinishCall(op uint64) (uint64, *net.TCPAddr, error) { - defer pinErrorThread()() - var handle uint64 - var outIP unsafe.Pointer - outIPArg := unsafe.Pointer(&outIP) - var outPort uint16 - outPortArg := unsafe.Pointer(&outPort) - err := n.tcpConnectFinish.call( - unsafe.Pointer(&handle), - unsafe.Pointer(&op), - unsafe.Pointer(&outIPArg), - unsafe.Pointer(&outPortArg), - ) - if err != nil { - return 0, nil, err - } - if handle == 0 { - return 0, nil, n.lastError() - } - return handle, n.takeTCPAddr(outIP, outPort), nil -} - -func (n *AsyncNative) tcpBindStartCall(instance string, port uint16, timeout time.Duration) (uint64, error) { - defer pinErrorThread()() - inst := cString(instance) - instPtr := unsafe.Pointer(&inst[0]) - timeoutMS := durationMillis(timeout) - var op uint64 - err := n.tcpBindStart.call( - unsafe.Pointer(&op), - unsafe.Pointer(&instPtr), - unsafe.Pointer(&port), - unsafe.Pointer(&timeoutMS), - ) - runtime.KeepAlive(inst) - return n.startResult(op, err) -} - -func (n *AsyncNative) tcpBindFinishCall(op uint64) (uint64, *net.TCPAddr, error) { - defer pinErrorThread()() - var handle uint64 - var outIP unsafe.Pointer - outIPArg := unsafe.Pointer(&outIP) - var outPort uint16 - outPortArg := unsafe.Pointer(&outPort) - err := n.tcpBindFinish.call( - unsafe.Pointer(&handle), - unsafe.Pointer(&op), - unsafe.Pointer(&outIPArg), - unsafe.Pointer(&outPortArg), - ) - if err != nil { - return 0, nil, err - } - if handle == 0 { - return 0, nil, n.lastError() - } - return handle, n.takeTCPAddr(outIP, outPort), nil -} - -func (n *AsyncNative) tcpAcceptStartCall(listener uint64, timeout time.Duration) (uint64, error) { - defer pinErrorThread()() - timeoutMS := durationMillis(timeout) - var op uint64 - err := n.tcpAcceptStart.call( - unsafe.Pointer(&op), - unsafe.Pointer(&listener), - unsafe.Pointer(&timeoutMS), - ) - return n.startResult(op, err) -} - -func (n *AsyncNative) tcpAcceptFinishCall(op uint64) (uint64, *net.TCPAddr, *net.TCPAddr, error) { - defer pinErrorThread()() - var handle uint64 - var localIP unsafe.Pointer - localIPArg := unsafe.Pointer(&localIP) - var localPort uint16 - localPortArg := unsafe.Pointer(&localPort) - var peerIP unsafe.Pointer - peerIPArg := unsafe.Pointer(&peerIP) - var peerPort uint16 - peerPortArg := unsafe.Pointer(&peerPort) - err := n.tcpAcceptFinish.call( - unsafe.Pointer(&handle), - unsafe.Pointer(&op), - unsafe.Pointer(&localIPArg), - unsafe.Pointer(&localPortArg), - unsafe.Pointer(&peerIPArg), - unsafe.Pointer(&peerPortArg), - ) - if err != nil { - return 0, nil, nil, err - } - if handle == 0 { - return 0, nil, nil, n.lastError() - } - return handle, n.takeTCPAddr(localIP, localPort), n.takeTCPAddr(peerIP, peerPort), nil -} - -func (n *AsyncNative) tcpReadStartCall(stream uint64, maxLen uint32, timeout time.Duration) (uint64, error) { - defer pinErrorThread()() - timeoutMS := durationMillis(timeout) - var op uint64 - err := n.tcpReadStart.call( - unsafe.Pointer(&op), - unsafe.Pointer(&stream), - unsafe.Pointer(&maxLen), - unsafe.Pointer(&timeoutMS), - ) - return n.startResult(op, err) -} - -func (n *AsyncNative) tcpReadFinishCall(op uint64) ([]byte, error) { - defer pinErrorThread()() - var ret int32 - var ptr unsafe.Pointer - ptrArg := unsafe.Pointer(&ptr) - var len uint32 - lenArg := unsafe.Pointer(&len) - err := n.tcpReadFinish.call( - unsafe.Pointer(&ret), - unsafe.Pointer(&op), - unsafe.Pointer(&ptrArg), - unsafe.Pointer(&lenArg), - ) - if err != nil { - return nil, err - } - if ret < 0 { - return nil, n.lastError() - } - return n.takeBytes(ptr, len), nil -} - -func (n *AsyncNative) tcpWriteStartCall(stream uint64, data []byte, timeout time.Duration) (uint64, error) { - defer pinErrorThread()() - ptr := unsafe.Pointer(nil) - if len(data) > 0 { - ptr = unsafe.Pointer(&data[0]) - } - timeoutMS := durationMillis(timeout) - length := uint32(len(data)) - var op uint64 - err := n.tcpWriteStart.call( - unsafe.Pointer(&op), - unsafe.Pointer(&stream), - unsafe.Pointer(&ptr), - unsafe.Pointer(&length), - unsafe.Pointer(&timeoutMS), - ) - runtime.KeepAlive(data) - return n.startResult(op, err) -} - -func (n *AsyncNative) tcpWriteFinishCall(op uint64) (int, error) { - defer pinErrorThread()() - var ret int32 - err := n.tcpWriteFinish.call(unsafe.Pointer(&ret), unsafe.Pointer(&op)) - if err != nil { - return 0, err - } - if ret < 0 { - return 0, n.lastError() - } - return int(ret), nil -} - -func (n *AsyncNative) udpBindStartCall(instance string, port uint16, timeout time.Duration) (uint64, error) { - defer pinErrorThread()() - inst := cString(instance) - instPtr := unsafe.Pointer(&inst[0]) - timeoutMS := durationMillis(timeout) - var op uint64 - err := n.udpBindStart.call( - unsafe.Pointer(&op), - unsafe.Pointer(&instPtr), - unsafe.Pointer(&port), - unsafe.Pointer(&timeoutMS), - ) - runtime.KeepAlive(inst) - return n.startResult(op, err) -} - -func (n *AsyncNative) udpBindFinishCall(op uint64) (uint64, *net.UDPAddr, error) { - defer pinErrorThread()() - var handle uint64 - var outIP unsafe.Pointer - outIPArg := unsafe.Pointer(&outIP) - var outPort uint16 - outPortArg := unsafe.Pointer(&outPort) - err := n.udpBindFinish.call( - unsafe.Pointer(&handle), - unsafe.Pointer(&op), - unsafe.Pointer(&outIPArg), - unsafe.Pointer(&outPortArg), - ) - if err != nil { - return 0, nil, err - } - if handle == 0 { - return 0, nil, n.lastError() - } - return handle, n.takeUDPAddr(outIP, outPort), nil -} - -func (n *AsyncNative) udpSendToStartCall(socket uint64, addr *net.UDPAddr, data []byte, timeout time.Duration) (uint64, error) { - defer pinErrorThread()() - if addr == nil || addr.IP == nil { - return 0, errors.New("udp destination address is nil") - } - dst := cString(addr.IP.String()) - dstPtr := unsafe.Pointer(&dst[0]) - ptr := unsafe.Pointer(nil) - if len(data) > 0 { - ptr = unsafe.Pointer(&data[0]) - } - port := uint16(addr.Port) - length := uint32(len(data)) - timeoutMS := durationMillis(timeout) - var op uint64 - err := n.udpSendToStart.call( - unsafe.Pointer(&op), - unsafe.Pointer(&socket), - unsafe.Pointer(&dstPtr), - unsafe.Pointer(&port), - unsafe.Pointer(&ptr), - unsafe.Pointer(&length), - unsafe.Pointer(&timeoutMS), - ) - runtime.KeepAlive(dst) - runtime.KeepAlive(data) - return n.startResult(op, err) -} - -func (n *AsyncNative) udpSendToFinishCall(op uint64) (int, error) { - defer pinErrorThread()() - var ret int32 - err := n.udpSendToFinish.call(unsafe.Pointer(&ret), unsafe.Pointer(&op)) - if err != nil { - return 0, err - } - if ret < 0 { - return 0, n.lastError() - } - return int(ret), nil -} - -func (n *AsyncNative) udpRecvFromStartCall(socket uint64, maxLen uint32, timeout time.Duration) (uint64, error) { - defer pinErrorThread()() - timeoutMS := durationMillis(timeout) - var op uint64 - err := n.udpRecvFromStart.call( - unsafe.Pointer(&op), - unsafe.Pointer(&socket), - unsafe.Pointer(&maxLen), - unsafe.Pointer(&timeoutMS), - ) - return n.startResult(op, err) -} - -func (n *AsyncNative) udpRecvFromFinishCall(op uint64) ([]byte, *net.UDPAddr, error) { - defer pinErrorThread()() - var ret int32 - var ptr unsafe.Pointer - ptrArg := unsafe.Pointer(&ptr) - var len uint32 - lenArg := unsafe.Pointer(&len) - var peerIP unsafe.Pointer - peerIPArg := unsafe.Pointer(&peerIP) - var peerPort uint16 - peerPortArg := unsafe.Pointer(&peerPort) - err := n.udpRecvFromFinish.call( - unsafe.Pointer(&ret), - unsafe.Pointer(&op), - unsafe.Pointer(&ptrArg), - unsafe.Pointer(&lenArg), - unsafe.Pointer(&peerIPArg), - unsafe.Pointer(&peerPortArg), - ) - if err != nil { - return nil, nil, err - } - if ret < 0 { - return nil, nil, n.lastError() - } - return n.takeBytes(ptr, len), n.takeUDPAddr(peerIP, peerPort), nil -} - -func (n *AsyncNative) waitOp(ctx context.Context, op uint64) error { - for { - if err := ctx.Err(); err != nil { - n.cancelAndFreeOp(op) - return err - } - - wait := asyncPollInterval - if deadline, ok := ctx.Deadline(); ok { - remaining := time.Until(deadline) - if remaining <= 0 { - n.cancelAndFreeOp(op) - return context.DeadlineExceeded - } - if remaining < wait { - wait = remaining - } - } - - status, err := n.opWaitStatus(op, wait) - if err != nil { - n.cancelAndFreeOp(op) - return err - } - switch status { - case dataPlaneOpPending: - continue - case dataPlaneOpReady, dataPlaneOpFailed: - return nil - case dataPlaneOpInvalid: - return errors.New("data plane async op is invalid") - default: - return fmt.Errorf("unexpected data plane async op status %d", status) - } - } -} - -func (n *AsyncNative) opWaitStatus(op uint64, timeout time.Duration) (int32, error) { - timeoutMS := durationMillis(timeout) - var status int32 - err := n.asyncOpWait.call( - unsafe.Pointer(&status), - unsafe.Pointer(&op), - unsafe.Pointer(&timeoutMS), - ) - return status, err -} - -func (n *AsyncNative) cancelAndFreeOp(op uint64) { - var ret int32 - _ = n.asyncOpCancel.call(unsafe.Pointer(&ret), unsafe.Pointer(&op)) - _ = n.asyncOpFree.call(unsafe.Pointer(&ret), unsafe.Pointer(&op)) -} - -func (n *AsyncNative) startResult(op uint64, err error) (uint64, error) { - if err != nil { - return 0, err - } - if op == 0 { - return 0, n.lastError() - } - return op, nil -} - -func (n *AsyncNative) lastError() error { - var out unsafe.Pointer - outArg := unsafe.Pointer(&out) - if err := n.getErrorMsg.call(nil, unsafe.Pointer(&outArg)); err != nil { - return err - } - if out == nil { - return errors.New("easytier ffi call failed") - } - msg := readCString(out) - _ = n.freeCString(out) - if msg == "" { - return errors.New("easytier ffi call failed") - } - if containsTimeout(msg) { - return timeoutError(msg) - } - return errors.New(msg) -} - -func (n *AsyncNative) freeCString(ptr unsafe.Pointer) error { - if ptr == nil { - return nil - } - return n.freeString.call(nil, unsafe.Pointer(&ptr)) -} - -func (n *AsyncNative) takeTCPAddr(ipPtr unsafe.Pointer, port uint16) *net.TCPAddr { - if ipPtr == nil { - return nil - } - ip := net.ParseIP(readCString(ipPtr)) - _ = n.freeCString(ipPtr) - return &net.TCPAddr{IP: ip, Port: int(port)} -} - -func (n *AsyncNative) takeUDPAddr(ipPtr unsafe.Pointer, port uint16) *net.UDPAddr { - if ipPtr == nil { - return nil - } - ip := net.ParseIP(readCString(ipPtr)) - _ = n.freeCString(ipPtr) - return &net.UDPAddr{IP: ip, Port: int(port)} -} - -func (n *AsyncNative) takeBytes(ptr unsafe.Pointer, len uint32) []byte { - if ptr == nil || len == 0 { - return nil - } - bytes := make([]byte, int(len)) - copy(bytes, unsafe.Slice((*byte)(ptr), int(len))) - _ = n.freeBytes.call(nil, unsafe.Pointer(&ptr), unsafe.Pointer(&len)) - return bytes -} - -func (n *AsyncNative) tcpCloseHandle(handle uint64) error { - defer pinErrorThread()() - var ret int32 - if err := n.tcpClose.call(unsafe.Pointer(&ret), unsafe.Pointer(&handle)); err != nil { - return err - } - if ret != 0 { - return n.lastError() - } - return nil -} - -func (n *AsyncNative) tcpListenerCloseHandle(handle uint64) error { - defer pinErrorThread()() - var ret int32 - if err := n.tcpListenerClose.call(unsafe.Pointer(&ret), unsafe.Pointer(&handle)); err != nil { - return err - } - if ret != 0 { - return n.lastError() - } - return nil -} - -func (n *AsyncNative) udpCloseHandle(handle uint64) error { - defer pinErrorThread()() - var ret int32 - if err := n.udpClose.call(unsafe.Pointer(&ret), unsafe.Pointer(&handle)); err != nil { - return err - } - if ret != 0 { - return n.lastError() - } - return nil -} - -func contextTimeout(ctx context.Context) (time.Duration, error) { - if err := ctx.Err(); err != nil { - return 0, err - } - timeout := defaultTimeout - if deadline, ok := ctx.Deadline(); ok { - timeout = time.Until(deadline) - } - if timeout <= 0 { - return 0, context.DeadlineExceeded - } - return timeout, nil -} - -func durationMillis(d time.Duration) uint64 { - if d <= 0 { - return 0 - } - ms := d / time.Millisecond - if ms <= 0 { - return 1 - } - return uint64(ms) -} - -func containsTimeout(msg string) bool { - return strings.Contains(msg, "timed out") || strings.Contains(msg, "timeout") -} - -func (b *atomicBool) Load() bool { - return b.v.Load() -} - -func (b *atomicBool) CompareAndSwap(old, new bool) bool { - return b.v.CompareAndSwap(old, new) -} - -var _ net.Conn = (*AsyncConn)(nil) -var _ net.Listener = (*AsyncListener)(nil) diff --git a/easytier-contrib/easytier-ffi/examples/go/easytier_async_test.go b/easytier-contrib/easytier-ffi/examples/go/easytier_async_test.go deleted file mode 100644 index 195b1dc0..00000000 --- a/easytier-contrib/easytier-ffi/examples/go/easytier_async_test.go +++ /dev/null @@ -1,360 +0,0 @@ -package easytierffi - -import ( - "context" - "fmt" - "io" - "net" - "os" - "strconv" - "testing" - "time" -) - -const asyncLocalTestTimeout = 120 * time.Second - -func TestAsyncSymbolBinding(t *testing.T) { - n := openAsyncForTest(t) - - status, err := n.opWaitStatus(0, 0) - if err != nil { - t.Fatal(err) - } - if status != dataPlaneOpInvalid { - t.Fatalf("expected invalid status for op 0, got %d", status) - } -} - -func TestAsyncLocalTwoNodeTCPAndUDP(t *testing.T) { - n := openAsyncForTest(t) - topology := startLocalAsyncTopology(t, n) - - ctx, cancel := context.WithTimeout(context.Background(), asyncLocalTestTimeout) - defer cancel() - - runAsyncTCPPingPong(t, ctx, n, topology) - runAsyncUDPPingPong(t, ctx, n, topology) -} - -type localAsyncTopology struct { - dialerInstance string - listenerInstance string - listenerIP string -} - -func openAsyncForTest(t *testing.T) *AsyncNative { - t.Helper() - - libraryPath := defaultLibraryPath() - if _, err := os.Stat(libraryPath); err != nil { - if os.IsNotExist(err) { - t.Skipf("build easytier-ffi with ffi-dataplane before running async tests: %v", err) - } - t.Fatalf("stat async ffi library: %v", err) - } - - n, err := OpenAsync(libraryPath) - if err != nil { - t.Fatalf("open async ffi library: %v", err) - } - t.Cleanup(func() { - if err := n.Close(); err != nil { - t.Errorf("close async native: %v", err) - } - }) - return n -} - -func startLocalAsyncTopology(t *testing.T, n *AsyncNative) localAsyncTopology { - t.Helper() - - suffix := strconv.FormatInt(time.Now().UnixNano(), 10) - networkName := "ffi-async-" + suffix - networkSecret := "ffi-async-secret-" + suffix - listenerInstance := "ffi-async-listener-" + suffix - dialerInstance := "ffi-async-dialer-" + suffix - listenerIP := "10.251.1.2" - dialerIP := "10.251.1.1" - listenerPort := freeLocalTCPPort(t) - listenerEndpoint := fmt.Sprintf("tcp://127.0.0.1:%d", listenerPort) - t.Cleanup(func() { - if err := n.deleteNetworkInstances([]string{dialerInstance, listenerInstance}); err != nil { - t.Errorf("cleanup async test EasyTier instances: %v", err) - } - }) - - listenerConfig := localAsyncConfig( - listenerInstance, - listenerIP, - networkName, - networkSecret, - []string{listenerEndpoint}, - nil, - ) - dialerConfig := localAsyncConfig( - dialerInstance, - dialerIP, - networkName, - networkSecret, - nil, - []string{listenerEndpoint}, - ) - - if err := n.RunNetworkInstance(listenerConfig); err != nil { - t.Fatalf("start listener instance: %v", err) - } - if err := n.RunNetworkInstance(dialerConfig); err != nil { - t.Fatalf("start dialer instance: %v", err) - } - - return localAsyncTopology{ - dialerInstance: dialerInstance, - listenerInstance: listenerInstance, - listenerIP: listenerIP, - } -} - -func localAsyncConfig(instance, ipv4, networkName, networkSecret string, listeners, peers []string) string { - config := fmt.Sprintf(`instance_name = %s -ipv4 = %s -listeners = %s - -[network_identity] -network_name = %s -network_secret = %s - -[flags] -no_tun = true -bind_device = false -`, - strconv.Quote(instance), - strconv.Quote(ipv4), - tomlStringList(listeners), - strconv.Quote(networkName), - strconv.Quote(networkSecret), - ) - for _, peer := range peers { - config += fmt.Sprintf("\n[[peer]]\nuri = %s\n", strconv.Quote(peer)) - } - return config -} - -func tomlStringList(values []string) string { - if len(values) == 0 { - return "[]" - } - - out := "[" - for i, value := range values { - if i > 0 { - out += ", " - } - out += strconv.Quote(value) - } - return out + "]" -} - -func freeLocalTCPPort(t *testing.T) int { - t.Helper() - - listener, err := net.Listen("tcp", "127.0.0.1:0") - if err != nil { - t.Fatalf("allocate local tcp port: %v", err) - } - defer listener.Close() - return listener.Addr().(*net.TCPAddr).Port -} - -func runAsyncTCPPingPong(t *testing.T, ctx context.Context, n *AsyncNative, topology localAsyncTopology) { - t.Helper() - - listener, listenerAddr := eventuallyTCPListen(t, ctx, n, topology.listenerInstance) - - tcpCtx, cancel := context.WithCancel(ctx) - accepted := make(chan error, 1) - defer waitForAsyncHelper(t, accepted, "tcp accept helper") - defer cancel() - defer listener.Close() - go func() { - conn, err := listener.Accept() - if err != nil { - accepted <- fmt.Errorf("accept tcp stream: %w", err) - return - } - defer conn.Close() - _ = conn.SetDeadline(time.Now().Add(30 * time.Second)) - - payload := make([]byte, len("ping")) - if _, err := io.ReadFull(conn, payload); err != nil { - accepted <- fmt.Errorf("read tcp ping: %w", err) - return - } - if string(payload) != "ping" { - accepted <- fmt.Errorf("expected tcp ping, got %q", string(payload)) - return - } - if _, err := conn.Write([]byte("pong")); err != nil { - accepted <- fmt.Errorf("write tcp pong: %w", err) - return - } - accepted <- nil - }() - - conn, err := eventuallyTCPDial(t, tcpCtx, n, topology.dialerInstance, topology.listenerIP, listenerAddr.Port) - if err != nil { - t.Fatal(err) - } - defer conn.Close() - _ = conn.SetDeadline(time.Now().Add(30 * time.Second)) - - if _, err := conn.Write([]byte("ping")); err != nil { - t.Fatalf("write tcp ping: %v", err) - } - payload := make([]byte, len("pong")) - if _, err := io.ReadFull(conn, payload); err != nil { - t.Fatalf("read tcp pong: %v", err) - } - if string(payload) != "pong" { - t.Fatalf("expected tcp pong, got %q", string(payload)) - } -} - -func eventuallyTCPListen(t *testing.T, ctx context.Context, n *AsyncNative, instance string) (net.Listener, *net.TCPAddr) { - t.Helper() - - var lastErr error - for attempt := 1; ctx.Err() == nil; attempt++ { - attemptCtx, cancel := context.WithTimeout(ctx, 5*time.Second) - listener, err := n.ListenContext(attemptCtx, instance, "tcp", "0.0.0.0:0") - cancel() - if err == nil { - addr := listener.Addr().(*net.TCPAddr) - t.Logf("async tcp bind succeeded on attempt %d at %s", attempt, addr) - return listener, addr - } - - lastErr = err - t.Logf("attempt %d: async tcp bind failed: %v", attempt, err) - waitForRetry(ctx, 500*time.Millisecond) - } - t.Fatalf("async tcp bind never succeeded: %v", lastErr) - panic("unreachable") -} - -func eventuallyTCPDial(t *testing.T, ctx context.Context, n *AsyncNative, instance, ip string, port int) (net.Conn, error) { - t.Helper() - - address := net.JoinHostPort(ip, strconv.Itoa(port)) - var lastErr error - for attempt := 1; ctx.Err() == nil; attempt++ { - attemptCtx, cancel := context.WithTimeout(ctx, 5*time.Second) - conn, err := n.DialContext(attemptCtx, instance, "tcp", address) - cancel() - if err == nil { - t.Logf("async tcp connect succeeded on attempt %d to %s", attempt, address) - return conn, nil - } - - lastErr = err - t.Logf("attempt %d: async tcp connect failed: %v", attempt, err) - waitForRetry(ctx, 500*time.Millisecond) - } - return nil, fmt.Errorf("async tcp connect never succeeded: %w", lastErr) -} - -func runAsyncUDPPingPong(t *testing.T, ctx context.Context, n *AsyncNative, topology localAsyncTopology) { - t.Helper() - - dialerSocket, err := n.UDPBindContext(ctx, topology.dialerInstance, 0) - if err != nil { - t.Fatalf("bind dialer udp socket: %v", err) - } - - listenerSocket, err := n.UDPBindContext(ctx, topology.listenerInstance, 0) - if err != nil { - t.Fatalf("bind listener udp socket: %v", err) - } - - udpCtx, cancel := context.WithCancel(ctx) - warmupDone := make(chan error, 1) - received := make(chan error, 1) - defer waitForAsyncHelper(t, received, "udp receive helper") - defer cancel() - defer listenerSocket.Close() - defer dialerSocket.Close() - - go func() { - if _, err := listenerSocket.SendTo(udpCtx, []byte("warmup"), dialerSocket.LocalAddr()); err != nil { - err = fmt.Errorf("send udp warmup: %w", err) - warmupDone <- err - received <- err - return - } - warmupDone <- nil - - payload, from, err := listenerSocket.RecvFrom(udpCtx, 512) - if err != nil { - received <- fmt.Errorf("recv udp ping: %w", err) - return - } - if string(payload) != "ping" { - received <- fmt.Errorf("expected udp ping, got %q", string(payload)) - return - } - if _, err := listenerSocket.SendTo(udpCtx, []byte("pong"), from); err != nil { - received <- fmt.Errorf("send udp pong: %w", err) - return - } - received <- nil - }() - - select { - case err := <-warmupDone: - if err != nil { - t.Fatal(err) - } - case <-udpCtx.Done(): - t.Fatal(udpCtx.Err()) - } - - target := &net.UDPAddr{IP: net.ParseIP(topology.listenerIP), Port: listenerSocket.LocalAddr().Port} - if _, err := dialerSocket.SendTo(udpCtx, []byte("ping"), target); err != nil { - t.Fatalf("send udp ping: %v", err) - } - for { - payload, from, err := dialerSocket.RecvFrom(udpCtx, 512) - if err != nil { - t.Fatalf("recv udp pong: %v", err) - } - if string(payload) == "pong" { - if !from.IP.Equal(target.IP) || from.Port != target.Port { - t.Fatalf("expected udp pong from %s, got %s", target, from) - } - break - } - t.Logf("skipping udp datagram from %s: %q", from, string(payload)) - } -} - -func waitForAsyncHelper(t *testing.T, done <-chan error, name string) { - t.Helper() - - select { - case err := <-done: - if err != nil { - t.Errorf("%s: %v", name, err) - } - case <-time.After(10 * time.Second): - t.Errorf("%s did not stop", name) - } -} - -func waitForRetry(ctx context.Context, delay time.Duration) { - timer := time.NewTimer(delay) - defer timer.Stop() - - select { - case <-timer.C: - case <-ctx.Done(): - } -} diff --git a/easytier-contrib/easytier-ffi/examples/go/easytier_test.go b/easytier-contrib/easytier-ffi/examples/go/easytier_test.go deleted file mode 100644 index 56943864..00000000 --- a/easytier-contrib/easytier-ffi/examples/go/easytier_test.go +++ /dev/null @@ -1,140 +0,0 @@ -package easytierffi - -import ( - "context" - "fmt" - "io" - "net" - "os" - "strconv" - "strings" - "testing" - "time" -) - -func TestSSHIntegration(t *testing.T) { - config := os.Getenv("EASYTIER_FFI_CONFIG") - instance := os.Getenv("EASYTIER_FFI_INSTANCE") - target := os.Getenv("EASYTIER_FFI_TARGET") - if config == "" || instance == "" || target == "" { - t.Skip("set EASYTIER_FFI_CONFIG, EASYTIER_FFI_INSTANCE and EASYTIER_FFI_TARGET to run integration test") - } - - n, err := Open(defaultLibraryPath()) - if err != nil { - t.Fatal(err) - } - defer n.Close() - - if err := n.RunNetworkInstance(config); err != nil { - t.Fatal(err) - } - - ctx, cancel := context.WithTimeout(context.Background(), 120*time.Second) - defer cancel() - - var lastErr error - for attempt := 1; ctx.Err() == nil; attempt++ { - conn, err := n.DialContext(ctx, instance, "tcp", target) - if err != nil { - lastErr = err - t.Logf("attempt %d: dial failed: %v", attempt, err) - time.Sleep(3 * time.Second) - continue - } - - _ = conn.SetReadDeadline(time.Now().Add(10 * time.Second)) - buf := make([]byte, 128) - nn, err := conn.Read(buf) - _ = conn.Close() - if err != nil { - lastErr = err - t.Logf("attempt %d: read failed: %v", attempt, err) - time.Sleep(3 * time.Second) - continue - } - banner := string(buf[:nn]) - if !strings.HasPrefix(banner, "SSH-") { - t.Fatalf("attempt %d: expected SSH banner, got %q", attempt, banner) - } - t.Logf("attempt %d: got banner %q", attempt, strings.TrimRight(banner, "\r\n")) - return - } - t.Fatalf("never got SSH banner, last err: %v", lastErr) -} - -func TestTCPListenIntegration(t *testing.T) { - config := os.Getenv("EASYTIER_FFI_LISTEN_CONFIG") - instance := os.Getenv("EASYTIER_FFI_LISTEN_INSTANCE") - listenPort := os.Getenv("EASYTIER_FFI_LISTEN_PORT") - if config == "" || instance == "" || listenPort == "" { - t.Skip("set EASYTIER_FFI_LISTEN_CONFIG, EASYTIER_FFI_LISTEN_INSTANCE and EASYTIER_FFI_LISTEN_PORT to run integration test") - } - port, err := strconv.ParseUint(listenPort, 10, 16) - if err != nil { - t.Fatal(err) - } - - n, err := Open(defaultLibraryPath()) - if err != nil { - t.Fatal(err) - } - defer n.Close() - - if err := n.RunNetworkInstance(config); err != nil { - t.Fatal(err) - } - - ctx, cancel := context.WithTimeout(context.Background(), 120*time.Second) - defer cancel() - - // Data-plane readiness is asynchronous: the instance must finish starting - // before the data plane accepts binds. Retry until ready or ctx expires. - var listener net.Listener - for attempt := 1; ; attempt++ { - listener, err = n.ListenContext(ctx, instance, "tcp", net.JoinHostPort("0.0.0.0", strconv.Itoa(int(port)))) - if err == nil { - break - } - if ctx.Err() != nil { - t.Fatalf("bind never succeeded, last err: %v", err) - } - t.Logf("attempt %d: bind failed: %v", attempt, err) - time.Sleep(3 * time.Second) - } - t.Logf("listening on %s; connect from another EasyTier peer and send ping", listener.Addr()) - - accepted := make(chan error, 1) - go func() { - conn, err := listener.Accept() - if err != nil { - accepted <- err - return - } - defer conn.Close() - - _ = conn.SetDeadline(time.Now().Add(10 * time.Second)) - buf := make([]byte, 4) - if _, err := io.ReadFull(conn, buf); err != nil { - accepted <- err - return - } - if string(buf) != "ping" { - accepted <- fmt.Errorf("expected %q, got %q", "ping", string(buf)) - return - } - _, err = conn.Write([]byte("pong")) - accepted <- err - }() - - select { - case err := <-accepted: - _ = listener.Close() - if err != nil { - t.Fatal(err) - } - case <-ctx.Done(): - _ = listener.Close() - t.Fatal(ctx.Err()) - } -} diff --git a/easytier-contrib/easytier-ffi/examples/go/go.mod b/easytier-contrib/easytier-ffi/examples/go/go.mod deleted file mode 100644 index cecc1de2..00000000 --- a/easytier-contrib/easytier-ffi/examples/go/go.mod +++ /dev/null @@ -1,5 +0,0 @@ -module easytierffi-example - -go 1.25 - -require github.com/go-webgpu/goffi v0.4.1 diff --git a/easytier-contrib/easytier-ffi/src/config_server.rs b/easytier-contrib/easytier-ffi/src/config_server.rs index 102a120f..ece1e388 100644 --- a/easytier-contrib/easytier-ffi/src/config_server.rs +++ b/easytier-contrib/easytier-ffi/src/config_server.rs @@ -13,18 +13,14 @@ use easytier::{ MachineIdOptions, config::{ConfigLoader as _, TomlConfigLoader}, }, - tunnel::TunnelScheme, - web_client::{WebClient, WebClientHooks, run_web_client}, + web_client::{WebClient, WebClientHooks, parse_config_server_endpoint, run_web_client}, }; use uuid::Uuid; use crate::{ - data_plane::remove_data_plane_handles_by_instance_ids, + data_plane::remove_data_plane_sessions_by_instance_ids, error::set_error_msg, - state::{ - ASYNC_RUNTIME, INSTANCE_MANAGER, INSTANCE_MUTATION_LOCK, INSTANCE_NAME_ID_MAP, - lock_remote_instance_mutation, remove_instance_name_ids, - }, + state::{ffi_context, resolve_instance_id_by_name}, strings::{c_str_to_string, optional_c_str_to_string}, types::ConfigServerEventCallback, }; @@ -76,37 +72,9 @@ pub fn validate_config_server_client_options( return Err("machine_id is empty".to_string()); } - let config_server_url = match url::Url::parse(config_server_url_s) { - Ok(url) => url, - Err(_) => format!( - "udp://config-server.easytier.cn:22020/{}", - config_server_url_s - ) - .parse() - .map_err(|err| format!("failed to parse config server URL: {}", err))?, - }; - - TunnelScheme::try_from(&config_server_url).map_err(|_| { - format!( - "unsupported config server scheme: {}", - config_server_url.scheme() - ) - })?; - - let token = config_server_url - .path_segments() - .and_then(|mut segments| segments.next_back()) - .map(|segment| percent_encoding::percent_decode_str(segment).decode_utf8()) - .transpose() - .map_err(|err| format!("failed to decode config server token: {}", err))? - .map(|token| token.to_string()) - .unwrap_or_default(); - - if token.is_empty() { - return Err("empty token".to_string()); - } - - Ok(()) + parse_config_server_endpoint(config_server_url_s) + .map(|_| ()) + .map_err(|error| error.to_string()) } struct ManagedConfigServerClient { @@ -150,7 +118,8 @@ impl ManagedConfigServerClientHooks { } fn validate_instance_name(&self, inst_name: &str, inst_id: Uuid) -> Result<(), String> { - if let Some(existing_id) = INSTANCE_NAME_ID_MAP.get(inst_name).map(|id| *id) + if let Some(existing_id) = + resolve_instance_id_by_name(inst_name).map_err(|error| error.to_string())? && existing_id != inst_id { return Err(format!("instance name {} already exists", inst_name)); @@ -159,13 +128,6 @@ impl ManagedConfigServerClientHooks { Ok(()) } - fn commit_instance_name(&self, inst_name: String, inst_id: Uuid) -> Result<(), String> { - INSTANCE_NAME_ID_MAP.retain(|_, existing_id| *existing_id != inst_id); - self.validate_instance_name(&inst_name, inst_id)?; - INSTANCE_NAME_ID_MAP.insert(inst_name, inst_id); - Ok(()) - } - pub(crate) fn start_stopping(&self) -> Vec { let _delivery_guard = if in_config_server_callback() { None @@ -199,11 +161,15 @@ impl ManagedConfigServerClientHooks { let Some(callback) = self.callback else { return Ok(()); }; - let instance_name = INSTANCE_MANAGER - .get_instance_name(&instance_id) + let instance_name = ffi_context() + .manager + .instance(instance_id) + .map(|instance| instance.instance_name().to_owned()) .unwrap_or_default(); - let network_name = INSTANCE_MANAGER - .get_network_name(&instance_id) + let network_name = ffi_context() + .manager + .config(instance_id) + .map(|config| config.get_network_identity().network_name) .unwrap_or_default(); let event_json = serde_json::json!({ "event": event, @@ -263,75 +229,27 @@ impl WebClientHooks for ManagedConfigServerClientHooks { .callback_delivery .lock() .map_err(|err| err.to_string())?; - let Some(inst_name) = INSTANCE_MANAGER.get_instance_name(id) else { - if !self.stopping.load(Ordering::Acquire) { - return Err(format!("instance {} not found after start", id)); - } - return Ok(()); + if self.stopping.load(Ordering::Acquire) { + return Err("config server client is stopping".to_string()); + } + let Some(inst_name) = ffi_context() + .manager + .instance(*id) + .map(|instance| instance.instance_name().to_owned()) + else { + return Err(format!("instance {} not found after start", id)); }; - { - let _mutation_guard = INSTANCE_MUTATION_LOCK - .lock() - .map_err(|err| err.to_string())?; - if INSTANCE_MANAGER.get_instance_name(id).is_none() { - if !self.stopping.load(Ordering::Acquire) { - return Err(format!("instance {} not found after start", id)); - } - return Ok(()); - } - - let should_delete = { - let mut guard = self.instance_ids.lock().map_err(|err| err.to_string())?; - if self.stopping.load(Ordering::Acquire) { - true - } else { - guard.insert(*id); - false - } - }; - - if should_delete { - if let Err(err) = INSTANCE_MANAGER.delete_network_instance(vec![*id]) { - return Err(err.to_string()); - } - remove_instance_name_ids(&[*id]); - return Ok(()); - } - - if self.stopping.load(Ordering::Acquire) { - self.remove_tracked_instance_ids(&[*id])?; - remove_instance_name_ids(&[*id]); - return Ok(()); - } - - if let Err(err) = self.commit_instance_name(inst_name.clone(), *id) { - self.remove_tracked_instance_ids(&[*id])?; - if let Err(delete_err) = INSTANCE_MANAGER.delete_network_instance(vec![*id]) { - return Err(format!( - "{}; failed to delete duplicate instance: {}", - err, delete_err - )); - } - return Err(err); - } - - if self.stopping.load(Ordering::Acquire) { - self.remove_tracked_instance_ids(&[*id])?; - remove_instance_name_ids(&[*id]); - return Ok(()); - } - if INSTANCE_MANAGER.get_instance_name(id).is_none() { - self.remove_tracked_instance_ids(&[*id])?; - remove_instance_name_ids(&[*id]); - return Err(format!( - "instance {} was removed before post-run completed", - id - )); - } + self.instance_ids + .lock() + .map_err(|err| err.to_string())? + .insert(*id); + if let Err(error) = self.validate_instance_name(&inst_name, *id) { + self.remove_tracked_instance_ids(&[*id])?; + return Err(error); } - remove_data_plane_handles_by_instance_ids(&[*id]); + remove_data_plane_sessions_by_instance_ids(&[*id]); if let Err(err) = self.emit_event_with_delivery_locked("run_network_instance", *id) { self.note_callback_error(err); @@ -340,15 +258,8 @@ impl WebClientHooks for ManagedConfigServerClientHooks { } async fn post_remove_network_instances(&self, ids: &[Uuid]) -> Result<(), String> { - let removed_ids = { - let _mutation_guard = INSTANCE_MUTATION_LOCK - .lock() - .map_err(|err| err.to_string())?; - let removed_ids = self.remove_tracked_instance_ids(ids)?; - remove_instance_name_ids(ids); - remove_data_plane_handles_by_instance_ids(&removed_ids); - removed_ids - }; + let removed_ids = self.remove_tracked_instance_ids(ids)?; + remove_data_plane_sessions_by_instance_ids(&removed_ids); for id in removed_ids { if let Err(err) = self.emit_event("delete_network_instance", id) { @@ -485,12 +396,12 @@ pub(crate) unsafe fn start_config_server_client( drop(data_plane_usage_guard); let hooks = Arc::new(ManagedConfigServerClientHooks::new(callback, user_data)); - let client = match ASYNC_RUNTIME.block_on(run_web_client( + let client = match ffi_context().runtime.block_on(run_web_client( &config_server_url, config_server_machine_id_options(machine_id), hostname, secure_mode, - INSTANCE_MANAGER.clone(), + ffi_context().manager.clone(), Some(hooks.clone()), )) { Ok(client) => client, @@ -511,7 +422,7 @@ pub(crate) fn stop_config_server_client() -> c_int { return -1; } - let mut guard = match CONFIG_SERVER_CLIENT.lock() { + let guard = match CONFIG_SERVER_CLIENT.lock() { Ok(guard) => guard, Err(err) => { set_error_msg(&format!("failed to lock config server client: {}", err)); @@ -528,29 +439,25 @@ pub(crate) fn stop_config_server_client() -> c_int { return -1; } let hooks = managed.hooks.clone(); - let managed = guard.take().expect("config server client exists"); + // Keep the client discoverable until the canonical transaction drains its + // tracking. Earlier removals must still retire IDs from these same hooks. drop(guard); - let _remote_mutation_guard = lock_remote_instance_mutation(); - let tracked_ids = hooks.start_stopping(); - drop(managed); - - let _mutation_guard = match INSTANCE_MUTATION_LOCK.lock() { - Ok(guard) => guard, + let delete_result = ffi_context().runtime.block_on( + ffi_context() + .process_management + .delete_owned_network_instances_selected_by(|| hooks.start_stopping()), + ); + let managed = match CONFIG_SERVER_CLIENT.lock() { + Ok(mut guard) => guard.take(), Err(err) => { hooks.wait_for_callback_delivery(); CONFIG_SERVER_CLIENT_ACTIVE.store(false, Ordering::Release); - CONFIG_SERVER_CLIENT_STOPPING.store(false, Ordering::Release); - set_error_msg(&format!("failed to lock instance mutation: {}", err)); + set_error_msg(&format!("failed to lock config server client: {err}")); return -1; } }; - let delete_result = INSTANCE_MANAGER.delete_network_instance(tracked_ids.clone()); - if delete_result.is_ok() { - remove_instance_name_ids(&tracked_ids); - remove_data_plane_handles_by_instance_ids(&tracked_ids); - } - drop(_mutation_guard); + drop(managed); hooks.wait_for_callback_delivery(); CONFIG_SERVER_CLIENT_ACTIVE.store(false, Ordering::Release); CONFIG_SERVER_CLIENT_STOPPING.store(false, Ordering::Release); diff --git a/easytier-contrib/easytier-ffi/src/data_plane.rs b/easytier-contrib/easytier-ffi/src/data_plane.rs deleted file mode 100644 index a565c85c..00000000 --- a/easytier-contrib/easytier-ffi/src/data_plane.rs +++ /dev/null @@ -1,928 +0,0 @@ -#[cfg(feature = "ffi-dataplane")] -use std::{ - future::Future, - net::{IpAddr, SocketAddr}, - sync::{ - Arc, RwLock, - atomic::{AtomicU64, Ordering}, - }, - time::Duration, -}; - -#[cfg(feature = "ffi-dataplane")] -use dashmap::DashMap; -#[cfg(feature = "ffi-dataplane")] -use easytier::launcher::{DataPlaneTcpListener, DataPlaneTcpStream, DataPlaneUdpSocket}; -#[cfg(feature = "ffi-dataplane")] -use tokio::io::{AsyncReadExt, AsyncWriteExt, ReadHalf, WriteHalf}; -#[cfg(feature = "ffi-dataplane")] -use tokio_util::sync::CancellationToken; -#[cfg(feature = "ffi-dataplane")] -use uuid::Uuid; - -#[cfg(feature = "ffi-dataplane")] -use crate::{ - config_server::{in_config_server_callback, is_config_server_active_or_stopping}, - error::{free_string, set_error_msg}, - state::{INSTANCE_MANAGER, INSTANCE_NAME_ID_MAP}, -}; - -#[cfg(feature = "ffi-dataplane")] -static NEXT_DATA_PLANE_HANDLE: AtomicU64 = AtomicU64::new(1); -#[cfg(feature = "ffi-dataplane")] -static DATA_PLANE_HANDLES: once_cell::sync::Lazy> = - once_cell::sync::Lazy::new(DashMap::new); -#[cfg(feature = "ffi-dataplane")] -static DATA_PLANE_USAGE_LOCK: once_cell::sync::Lazy> = - once_cell::sync::Lazy::new(|| RwLock::new(())); - -#[cfg(feature = "ffi-dataplane")] -pub(crate) struct DataPlaneHandle { - pub(crate) instance_id: uuid::Uuid, - pub(crate) runtime: tokio::runtime::Handle, - // Cancelled by close() to wake any in-flight op on this handle. - pub(crate) close_token: CancellationToken, - pub(crate) resource: DataPlaneResource, -} - -#[cfg(feature = "ffi-dataplane")] -pub(crate) struct TcpHalves { - pub(crate) read: tokio::sync::Mutex>, - pub(crate) write: tokio::sync::Mutex>, -} - -#[cfg(feature = "ffi-dataplane")] -pub(crate) enum DataPlaneResource { - Tcp(Arc), - TcpListener(Arc>), - Udp(Arc), -} - -// Several helper functions for FFI data plane operations to facilitate logic reuse. - -#[cfg(feature = "ffi-dataplane")] -pub(crate) fn next_handle() -> u64 { - NEXT_DATA_PLANE_HANDLE.fetch_add(1, Ordering::Relaxed) -} - -#[cfg(feature = "ffi-dataplane")] -pub(crate) fn timeout_duration(timeout_ms: u64) -> Duration { - Duration::from_millis(timeout_ms) -} - -#[cfg(feature = "ffi-dataplane")] -pub(crate) unsafe fn cstr_to_string(ptr: *const std::ffi::c_char, name: &str) -> Option { - if ptr.is_null() { - set_error_msg(&format!("{} is null", name)); - return None; - } - Some( - unsafe { std::ffi::CStr::from_ptr(ptr) } - .to_string_lossy() - .into_owned(), - ) -} - -#[cfg(feature = "ffi-dataplane")] -pub(crate) fn get_instance_id(inst_name: &str) -> Option { - INSTANCE_NAME_ID_MAP.get(inst_name).map(|id| *id.value()) -} - -#[cfg(feature = "ffi-dataplane")] -pub(crate) fn parse_socket_addr(host: &str, port: u16) -> Option { - let ip = match host.parse::() { - Ok(ip) => ip, - Err(e) => { - set_error_msg(&format!("failed to parse ip address: {}", e)); - return None; - } - }; - Some(SocketAddr::new(ip, port)) -} - -/// Encode an IP address for FFI return. Returns `*mut c_char` to match -/// `CString::into_raw`; caller releases it via `free_string`. -#[cfg(feature = "ffi-dataplane")] -pub(crate) fn into_ffi_ip_cstring(ip: IpAddr) -> Option<*mut std::ffi::c_char> { - match std::ffi::CString::new(ip.to_string()) { - Ok(s) => Some(s.into_raw()), - Err(e) => { - set_error_msg(&format!("failed to encode ip: {}", e)); - None - } - } -} - -#[cfg(feature = "ffi-dataplane")] -pub(crate) fn get_runtime_handle( - inst_id: &uuid::Uuid, - deadline: std::time::Instant, -) -> Option { - let remaining = deadline.saturating_duration_since(std::time::Instant::now()); - let Some(rt) = INSTANCE_MANAGER.data_plane_wait_runtime_handle(inst_id, remaining) else { - set_error_msg("instance runtime is not ready"); - return None; - }; - Some(rt) -} - -#[cfg(feature = "ffi-dataplane")] -pub(crate) fn insert_tcp_stream_handle( - instance_id: uuid::Uuid, - runtime: tokio::runtime::Handle, - stream: DataPlaneTcpStream, -) -> u64 { - let (rd, wr) = tokio::io::split(stream); - let handle = next_handle(); - DATA_PLANE_HANDLES.insert( - handle, - DataPlaneHandle { - instance_id, - runtime, - close_token: CancellationToken::new(), - resource: DataPlaneResource::Tcp(Arc::new(TcpHalves { - read: tokio::sync::Mutex::new(rd), - write: tokio::sync::Mutex::new(wr), - })), - }, - ); - handle -} - -#[cfg(feature = "ffi-dataplane")] -pub(crate) fn insert_tcp_listener_handle( - instance_id: uuid::Uuid, - runtime: tokio::runtime::Handle, - listener: DataPlaneTcpListener, -) -> u64 { - let handle = next_handle(); - DATA_PLANE_HANDLES.insert( - handle, - DataPlaneHandle { - instance_id, - runtime, - close_token: CancellationToken::new(), - resource: DataPlaneResource::TcpListener(Arc::new(tokio::sync::Mutex::new(listener))), - }, - ); - handle -} - -#[cfg(feature = "ffi-dataplane")] -pub(crate) fn insert_udp_socket_handle( - instance_id: uuid::Uuid, - runtime: tokio::runtime::Handle, - socket: DataPlaneUdpSocket, -) -> u64 { - let handle = next_handle(); - DATA_PLANE_HANDLES.insert( - handle, - DataPlaneHandle { - instance_id, - runtime, - close_token: CancellationToken::new(), - resource: DataPlaneResource::Udp(Arc::new(socket)), - }, - ); - handle -} - -#[cfg(feature = "ffi-dataplane")] -pub(crate) fn get_tcp_stream( - handle: u64, -) -> Option<(Arc, tokio::runtime::Handle, CancellationToken)> { - get_tcp_stream_with_instance(handle) - .map(|(halves, runtime, close_token, _)| (halves, runtime, close_token)) -} - -#[cfg(feature = "ffi-dataplane")] -pub(crate) fn get_tcp_stream_with_instance( - handle: u64, -) -> Option<( - Arc, - tokio::runtime::Handle, - CancellationToken, - uuid::Uuid, -)> { - let Some(h) = DATA_PLANE_HANDLES.get(&handle) else { - set_error_msg("tcp stream handle not found"); - return None; - }; - match &h.resource { - DataPlaneResource::Tcp(halves) => Some(( - halves.clone(), - h.runtime.clone(), - h.close_token.clone(), - h.instance_id, - )), - DataPlaneResource::TcpListener(_) | DataPlaneResource::Udp(_) => { - set_error_msg("handle is not a tcp stream"); - None - } - } -} - -#[cfg(feature = "ffi-dataplane")] -pub(crate) fn get_tcp_listener( - handle: u64, -) -> Option<( - Arc>, - tokio::runtime::Handle, - CancellationToken, - uuid::Uuid, -)> { - let Some(h) = DATA_PLANE_HANDLES.get(&handle) else { - set_error_msg("tcp listener handle not found"); - return None; - }; - match &h.resource { - DataPlaneResource::TcpListener(listener) => Some(( - listener.clone(), - h.runtime.clone(), - h.close_token.clone(), - h.instance_id, - )), - DataPlaneResource::Tcp(_) | DataPlaneResource::Udp(_) => { - set_error_msg("handle is not a tcp listener"); - None - } - } -} - -#[cfg(feature = "ffi-dataplane")] -pub(crate) fn get_udp_socket( - handle: u64, -) -> Option<( - Arc, - tokio::runtime::Handle, - CancellationToken, -)> { - get_udp_socket_with_instance(handle) - .map(|(socket, runtime, close_token, _)| (socket, runtime, close_token)) -} - -#[cfg(feature = "ffi-dataplane")] -pub(crate) fn get_udp_socket_with_instance( - handle: u64, -) -> Option<( - Arc, - tokio::runtime::Handle, - CancellationToken, - uuid::Uuid, -)> { - let Some(h) = DATA_PLANE_HANDLES.get(&handle) else { - set_error_msg("udp socket handle not found"); - return None; - }; - match &h.resource { - DataPlaneResource::Udp(socket) => Some(( - socket.clone(), - h.runtime.clone(), - h.close_token.clone(), - h.instance_id, - )), - DataPlaneResource::Tcp(_) | DataPlaneResource::TcpListener(_) => { - set_error_msg("handle is not a udp socket"); - None - } - } -} - -#[cfg(feature = "ffi-dataplane")] -pub(crate) fn remove_data_plane_handles_by_instance_ids(ids: &[Uuid]) { - if ids.is_empty() { - return; - } - - let _data_plane_usage_guard = DATA_PLANE_USAGE_LOCK - .write() - .unwrap_or_else(|err| err.into_inner()); - - DATA_PLANE_HANDLES.retain(|_, handle| { - if ids.contains(&handle.instance_id) { - handle.close_token.cancel(); - false - } else { - true - } - }); - crate::data_plane_async::remove_ops_by_instance_ids(ids); -} - -#[cfg(not(feature = "ffi-dataplane"))] -pub(crate) fn remove_data_plane_handles_by_instance_ids(_ids: &[uuid::Uuid]) {} - -#[cfg(feature = "ffi-dataplane")] -pub(crate) fn data_plane_rejected() -> bool { - if in_config_server_callback() { - set_error_msg("cannot use data plane from config server callback"); - true - } else if is_config_server_active_or_stopping() { - set_error_msg("cannot use data plane while config server client is active"); - true - } else { - false - } -} - -#[cfg(feature = "ffi-dataplane")] -pub(crate) fn enter_data_plane_operation() -> Option> { - if data_plane_rejected() { - return None; - } - - let guard = match DATA_PLANE_USAGE_LOCK.read() { - Ok(guard) => guard, - Err(err) => { - set_error_msg(&format!("failed to lock data plane usage: {}", err)); - return None; - } - }; - if data_plane_rejected() { - return None; - } - Some(guard) -} - -/// Run an IO op on the resource's owning runtime, supporting -/// timeout and cancellation. -#[cfg(feature = "ffi-dataplane")] -async fn run_with_cancel( - close_token: &CancellationToken, - timeout_ms: u64, - error_prefix: &str, - op: F, -) -> Option> -where - F: Future>, -{ - tokio::select! { - biased; - _ = close_token.cancelled() => { - set_error_msg(&format!("{}: handle closed", error_prefix)); - None - } - res = tokio::time::timeout(timeout_duration(timeout_ms), op) => match res { - Ok(r) => Some(r), - Err(_) => { - set_error_msg(&format!("{} timed out", error_prefix)); - None - } - } - } -} - -#[cfg(feature = "ffi-dataplane")] -pub(crate) fn lock_for_config_server_start() --> Result, String> { - let guard = DATA_PLANE_USAGE_LOCK - .write() - .map_err(|err| format!("failed to lock data plane usage: {}", err))?; - if !DATA_PLANE_HANDLES.is_empty() || crate::data_plane_async::has_live_ops() { - return Err("cannot start config server client while data plane is in use".to_string()); - } - Ok(guard) -} -/// # Safety -/// Open a TCP stream through an EasyTier instance data plane. Returns 0 on -/// failure. On success, writes the local socket address chosen for this -/// connection into `out_local_ip` (a heap-allocated C string the caller must -/// release via `free_string`) and `out_local_port`. Both out pointers must be -/// non-null. -#[cfg(feature = "ffi-dataplane")] -pub(crate) unsafe fn data_plane_tcp_connect( - inst_name: *const std::ffi::c_char, - dst_ip: *const std::ffi::c_char, - dst_port: std::ffi::c_ushort, - timeout_ms: u64, - out_local_ip: *mut *const std::ffi::c_char, - out_local_port: *mut std::ffi::c_ushort, -) -> u64 { - let _data_plane_usage_guard = match enter_data_plane_operation() { - Some(guard) => guard, - None => return 0, - }; - if out_local_ip.is_null() || out_local_port.is_null() { - set_error_msg("output pointer is null"); - return 0; - } - let Some(inst_name) = (unsafe { cstr_to_string(inst_name, "inst_name") }) else { - return 0; - }; - let Some(dst_ip) = (unsafe { cstr_to_string(dst_ip, "dst_ip") }) else { - return 0; - }; - let Some(inst_id) = get_instance_id(&inst_name) else { - set_error_msg("instance not found"); - return 0; - }; - let Some(dst_addr) = parse_socket_addr(&dst_ip, dst_port) else { - return 0; - }; - let deadline = std::time::Instant::now() + timeout_duration(timeout_ms); - let Some(runtime) = get_runtime_handle(&inst_id, deadline) else { - return 0; - }; - - let remaining = deadline.saturating_duration_since(std::time::Instant::now()); - let result = - runtime.block_on(INSTANCE_MANAGER.data_plane_tcp_connect(&inst_id, dst_addr, remaining)); - match result { - Ok(stream) => { - let local_addr = stream.local_addr(); - let Some(local_ip) = into_ffi_ip_cstring(local_addr.ip()) else { - return 0; - }; - let handle = insert_tcp_stream_handle(inst_id, runtime, stream); - unsafe { - *out_local_ip = local_ip as *const std::ffi::c_char; - *out_local_port = local_addr.port(); - } - handle - } - Err(e) => { - set_error_msg(&format!("failed to connect tcp data plane: {}", e)); - 0 - } - } -} - -/// # Safety -/// Bind a TCP listener through an EasyTier instance data plane. Returns 0 on -/// failure. The local address actually bound is written into `out_local_ip` / -/// `out_local_port`; the caller must release `*out_local_ip` via `free_string`. -#[cfg(feature = "ffi-dataplane")] -pub(crate) unsafe fn data_plane_tcp_bind( - inst_name: *const std::ffi::c_char, - local_port: std::ffi::c_ushort, - timeout_ms: u64, - out_local_ip: *mut *const std::ffi::c_char, - out_local_port: *mut std::ffi::c_ushort, -) -> u64 { - let _data_plane_usage_guard = match enter_data_plane_operation() { - Some(guard) => guard, - None => return 0, - }; - if out_local_ip.is_null() || out_local_port.is_null() { - set_error_msg("output pointer is null"); - return 0; - } - let Some(inst_name) = (unsafe { cstr_to_string(inst_name, "inst_name") }) else { - return 0; - }; - let Some(inst_id) = get_instance_id(&inst_name) else { - set_error_msg("instance not found"); - return 0; - }; - let deadline = std::time::Instant::now() + timeout_duration(timeout_ms); - let Some(runtime) = get_runtime_handle(&inst_id, deadline) else { - return 0; - }; - - let remaining = deadline.saturating_duration_since(std::time::Instant::now()); - let result = - runtime.block_on(INSTANCE_MANAGER.data_plane_tcp_bind(&inst_id, local_port, remaining)); - match result { - Ok(listener) => { - let local_addr = listener.local_addr(); - let Some(local_ip) = into_ffi_ip_cstring(local_addr.ip()) else { - return 0; - }; - let handle = insert_tcp_listener_handle(inst_id, runtime, listener); - unsafe { - *out_local_ip = local_ip as *const std::ffi::c_char; - *out_local_port = local_addr.port(); - } - handle - } - Err(e) => { - set_error_msg(&format!("failed to bind tcp data plane: {}", e)); - 0 - } - } -} - -/// # Safety -/// Accept one connection from a TCP data-plane listener. Returns a TCP stream -/// handle, or 0 on failure. Local and peer addresses are written into out -/// parameters; returned IP strings must be released via `free_string`. -#[cfg(feature = "ffi-dataplane")] -pub(crate) unsafe fn data_plane_tcp_accept( - handle: u64, - timeout_ms: u64, - out_local_ip: *mut *const std::ffi::c_char, - out_local_port: *mut std::ffi::c_ushort, - out_peer_ip: *mut *const std::ffi::c_char, - out_peer_port: *mut std::ffi::c_ushort, -) -> u64 { - let _data_plane_usage_guard = match enter_data_plane_operation() { - Some(guard) => guard, - None => return 0, - }; - if out_local_ip.is_null() - || out_local_port.is_null() - || out_peer_ip.is_null() - || out_peer_port.is_null() - { - set_error_msg("output pointer is null"); - return 0; - } - let Some((listener, runtime, close_token, instance_id)) = get_tcp_listener(handle) else { - return 0; - }; - - let ret = runtime.block_on(async move { - let mut listener = listener.lock().await; - run_with_cancel( - &close_token, - timeout_ms, - "tcp data plane accept", - listener.accept(), - ) - .await - }); - - match ret { - Some(Ok((stream, peer_addr))) => { - let local_addr = stream.local_addr(); - let Some(local_ip) = into_ffi_ip_cstring(local_addr.ip()) else { - return 0; - }; - let Some(peer_ip) = into_ffi_ip_cstring(peer_addr.ip()) else { - free_string(local_ip); - return 0; - }; - let stream_handle = insert_tcp_stream_handle(instance_id, runtime, stream); - unsafe { - *out_local_ip = local_ip as *const std::ffi::c_char; - *out_local_port = local_addr.port(); - *out_peer_ip = peer_ip as *const std::ffi::c_char; - *out_peer_port = peer_addr.port(); - } - stream_handle - } - Some(Err(e)) => { - set_error_msg(&format!("failed to accept tcp data plane: {}", e)); - 0 - } - None => 0, - } -} - -/// # Safety -/// Read from a TCP data-plane stream. -#[cfg(feature = "ffi-dataplane")] -pub(crate) unsafe fn data_plane_tcp_read( - handle: u64, - buf: *mut std::ffi::c_uchar, - len: u32, - timeout_ms: u64, -) -> std::ffi::c_int { - let _data_plane_usage_guard = match enter_data_plane_operation() { - Some(guard) => guard, - None => return -1, - }; - if buf.is_null() { - set_error_msg("buf is null"); - return -1; - } - let Some((halves, runtime, close_token)) = get_tcp_stream(handle) else { - return -1; - }; - // Safety: caller-owned buffer outlives this blocking call. - let buf = unsafe { std::slice::from_raw_parts_mut(buf, len as usize) }; - runtime.block_on(async move { - let mut rd = halves.read.lock().await; - match run_with_cancel( - &close_token, - timeout_ms, - "failed to read tcp data plane", - rd.read(buf), - ) - .await - { - Some(Ok(n)) => n as std::ffi::c_int, - Some(Err(e)) => { - set_error_msg(&format!("failed to read tcp data plane: {}", e)); - -1 - } - None => -1, - } - }) -} - -/// # Safety -/// Write to a TCP data-plane stream. -#[cfg(feature = "ffi-dataplane")] -pub(crate) unsafe fn data_plane_tcp_write( - handle: u64, - buf: *const std::ffi::c_uchar, - len: u32, - timeout_ms: u64, -) -> std::ffi::c_int { - let _data_plane_usage_guard = match enter_data_plane_operation() { - Some(guard) => guard, - None => return -1, - }; - if buf.is_null() { - set_error_msg("buf is null"); - return -1; - } - let Some((halves, runtime, close_token)) = get_tcp_stream(handle) else { - return -1; - }; - let total = len as usize; - // Safety: caller-owned buffer outlives this blocking call. - let buf = unsafe { std::slice::from_raw_parts(buf, total) }; - runtime.block_on(async move { - let mut wr = halves.write.lock().await; - // Use `write_all` to honor `net.Conn::Write` semantics on the Go side - // (must write everything or return an error); single `write()` can - // silently short-write and corrupt streams that the caller assumes are - // fully written. - match run_with_cancel( - &close_token, - timeout_ms, - "failed to write tcp data plane", - wr.write_all(buf), - ) - .await - { - Some(Ok(())) => total as std::ffi::c_int, - Some(Err(e)) => { - set_error_msg(&format!("failed to write tcp data plane: {}", e)); - -1 - } - None => -1, - } - }) -} - -#[cfg(feature = "ffi-dataplane")] -pub(crate) fn data_plane_tcp_close(handle: u64) -> std::ffi::c_int { - let _data_plane_usage_guard = match enter_data_plane_operation() { - Some(guard) => guard, - None => return -1, - }; - crate::data_plane_async::cancel_ops_for_handle(handle); - let Some((_, h)) = DATA_PLANE_HANDLES.remove_if(&handle, |_, e| { - matches!(e.resource, DataPlaneResource::Tcp(_)) - }) else { - set_error_msg(if DATA_PLANE_HANDLES.contains_key(&handle) { - "handle is not a tcp stream" - } else { - "tcp stream handle not found" - }); - return -1; - }; - h.close_token.cancel(); - if let DataPlaneResource::Tcp(halves) = h.resource { - // Best-effort half-close; if write half is in use, the in-flight call - // observes the cancel token and releases the lock shortly after. - h.runtime.spawn(async move { - if let Ok(mut wr) = halves.write.try_lock() { - let _ = wr.shutdown().await; - } - }); - } - 0 -} - -#[cfg(feature = "ffi-dataplane")] -pub(crate) fn data_plane_tcp_listener_close(handle: u64) -> std::ffi::c_int { - let _data_plane_usage_guard = match enter_data_plane_operation() { - Some(guard) => guard, - None => return -1, - }; - crate::data_plane_async::cancel_ops_for_handle(handle); - let Some((_, h)) = DATA_PLANE_HANDLES.remove_if(&handle, |_, e| { - matches!(e.resource, DataPlaneResource::TcpListener(_)) - }) else { - set_error_msg(if DATA_PLANE_HANDLES.contains_key(&handle) { - "handle is not a tcp listener" - } else { - "tcp listener handle not found" - }); - return -1; - }; - h.close_token.cancel(); - 0 -} - -/// # Safety -/// Bind a UDP socket through an EasyTier instance data plane. Returns 0 on -/// failure. The local address actually bound (which may differ from the -/// requested port when `local_port == 0`) is written into `out_local_ip` / -/// `out_local_port`; the caller must release `*out_local_ip` via `free_string`. -#[cfg(feature = "ffi-dataplane")] -pub(crate) unsafe fn data_plane_udp_bind( - inst_name: *const std::ffi::c_char, - local_port: std::ffi::c_ushort, - timeout_ms: u64, - out_local_ip: *mut *const std::ffi::c_char, - out_local_port: *mut std::ffi::c_ushort, -) -> u64 { - let _data_plane_usage_guard = match enter_data_plane_operation() { - Some(guard) => guard, - None => return 0, - }; - if out_local_ip.is_null() || out_local_port.is_null() { - set_error_msg("output pointer is null"); - return 0; - } - let Some(inst_name) = (unsafe { cstr_to_string(inst_name, "inst_name") }) else { - return 0; - }; - let Some(inst_id) = get_instance_id(&inst_name) else { - set_error_msg("instance not found"); - return 0; - }; - let deadline = std::time::Instant::now() + timeout_duration(timeout_ms); - let Some(runtime) = get_runtime_handle(&inst_id, deadline) else { - return 0; - }; - - let remaining = deadline.saturating_duration_since(std::time::Instant::now()); - let result = - runtime.block_on(INSTANCE_MANAGER.data_plane_udp_bind(&inst_id, local_port, remaining)); - match result { - Ok(socket) => { - let local_addr = socket.local_addr(); - let Some(local_ip) = into_ffi_ip_cstring(local_addr.ip()) else { - return 0; - }; - let handle = insert_udp_socket_handle(inst_id, runtime, socket); - unsafe { - *out_local_ip = local_ip as *const std::ffi::c_char; - *out_local_port = local_addr.port(); - } - handle - } - Err(e) => { - set_error_msg(&format!("failed to bind udp data plane: {}", e)); - 0 - } - } -} - -/// # Safety -/// Send a datagram through a UDP data-plane socket. -#[cfg(feature = "ffi-dataplane")] -pub(crate) unsafe fn data_plane_udp_send_to( - handle: u64, - dst_ip: *const std::ffi::c_char, - dst_port: std::ffi::c_ushort, - buf: *const std::ffi::c_uchar, - len: u32, - timeout_ms: u64, -) -> std::ffi::c_int { - let _data_plane_usage_guard = match enter_data_plane_operation() { - Some(guard) => guard, - None => return -1, - }; - if buf.is_null() { - set_error_msg("buf is null"); - return -1; - } - let Some(dst_ip) = (unsafe { cstr_to_string(dst_ip, "dst_ip") }) else { - return -1; - }; - let Some(dst_addr) = parse_socket_addr(&dst_ip, dst_port) else { - return -1; - }; - let Some((socket, runtime, close_token)) = get_udp_socket(handle) else { - return -1; - }; - let total = len as usize; - // Safety: caller-owned buffer outlives this blocking call. - let buf = unsafe { std::slice::from_raw_parts(buf, total) }; - runtime.block_on(async move { - match run_with_cancel( - &close_token, - timeout_ms, - "failed to send udp data plane", - socket.send_to(buf, dst_addr), - ) - .await - { - Some(Ok(n)) => n as std::ffi::c_int, - Some(Err(e)) => { - set_error_msg(&format!("failed to send udp data plane: {}", e)); - -1 - } - None => -1, - } - }) -} - -/// # Safety -/// Receive a datagram from a UDP data-plane socket. -#[cfg(feature = "ffi-dataplane")] -pub(crate) unsafe fn data_plane_udp_recv_from( - handle: u64, - buf: *mut std::ffi::c_uchar, - len: u32, - out_ip: *mut *const std::ffi::c_char, - out_port: *mut std::ffi::c_ushort, - timeout_ms: u64, -) -> std::ffi::c_int { - let _data_plane_usage_guard = match enter_data_plane_operation() { - Some(guard) => guard, - None => return -1, - }; - if buf.is_null() || out_ip.is_null() || out_port.is_null() { - set_error_msg("output pointer is null"); - return -1; - } - let Some((socket, runtime, close_token)) = get_udp_socket(handle) else { - return -1; - }; - let total = len as usize; - // Safety: caller-owned buffer outlives this blocking call. - let buf = unsafe { std::slice::from_raw_parts_mut(buf, total) }; - let ret = runtime.block_on(run_with_cancel( - &close_token, - timeout_ms, - "udp data plane receive", - socket.recv_from(buf), - )); - - match ret { - Some(Ok((n, addr))) => { - // The returned ip pointer must be released by the caller via - // `free_string` (which calls `CString::from_raw`, matching - // `CString::into_raw` here). - let Some(ip_cstr) = into_ffi_ip_cstring(addr.ip()) else { - return -1; - }; - unsafe { - *out_ip = ip_cstr as *const std::ffi::c_char; - *out_port = addr.port() as std::ffi::c_ushort; - } - n as std::ffi::c_int - } - Some(Err(e)) => { - set_error_msg(&format!("failed to receive udp data plane: {}", e)); - -1 - } - None => -1, - } -} - -#[cfg(feature = "ffi-dataplane")] -pub(crate) fn data_plane_udp_close(handle: u64) -> std::ffi::c_int { - let _data_plane_usage_guard = match enter_data_plane_operation() { - Some(guard) => guard, - None => return -1, - }; - crate::data_plane_async::cancel_ops_for_handle(handle); - let Some((_, h)) = DATA_PLANE_HANDLES.remove_if(&handle, |_, e| { - matches!(e.resource, DataPlaneResource::Udp(_)) - }) else { - set_error_msg(if DATA_PLANE_HANDLES.contains_key(&handle) { - "handle is not a udp socket" - } else { - "udp socket handle not found" - }); - return -1; - }; - h.close_token.cancel(); - 0 -} - -#[cfg(all(test, feature = "ffi-dataplane"))] -mod tests { - use super::*; - use std::{sync::mpsc, time::Duration}; - - #[test] - fn config_server_start_waits_for_data_plane_operation() { - let read_guard = DATA_PLANE_USAGE_LOCK.read().unwrap(); - let (done_tx, done_rx) = mpsc::channel(); - let waiter = std::thread::spawn(move || { - let _write_guard = lock_for_config_server_start().unwrap(); - done_tx.send(()).unwrap(); - }); - - assert!(done_rx.recv_timeout(Duration::from_millis(100)).is_err()); - drop(read_guard); - done_rx.recv_timeout(Duration::from_secs(5)).unwrap(); - waiter.join().unwrap(); - } - - #[test] - fn instance_cleanup_waits_for_data_plane_operation() { - let read_guard = DATA_PLANE_USAGE_LOCK.read().unwrap(); - let instance_id = Uuid::new_v4(); - let (done_tx, done_rx) = mpsc::channel(); - let cleaner = std::thread::spawn(move || { - remove_data_plane_handles_by_instance_ids(&[instance_id]); - done_tx.send(()).unwrap(); - }); - - assert!(done_rx.recv_timeout(Duration::from_millis(100)).is_err()); - drop(read_guard); - done_rx.recv_timeout(Duration::from_secs(5)).unwrap(); - cleaner.join().unwrap(); - } -} diff --git a/easytier-contrib/easytier-ffi/src/data_plane/abi.rs b/easytier-contrib/easytier-ffi/src/data_plane/abi.rs new file mode 100644 index 00000000..71fc1f16 --- /dev/null +++ b/easytier-contrib/easytier-ffi/src/data_plane/abi.rs @@ -0,0 +1,662 @@ +use std::{ + ffi::{c_char, c_int, c_uchar}, + net::{IpAddr, Ipv4Addr, SocketAddr}, + ptr, +}; + +use easytier_core::gateway::DataPlaneErrorKind; + +use super::session::{self, NativeDataPlaneError, NativeDataPlaneResult}; +use crate::{ + error::set_error_msg, + strings::c_str_to_string, + types::{DataPlaneCompletion, DataPlaneSocketAddr}, +}; + +fn failure(error: NativeDataPlaneError) -> c_int { + set_error_msg(&error.message); + -(error.kind as c_int) +} + +fn status(result: NativeDataPlaneResult<()>) -> c_int { + match result { + Ok(()) => 0, + Err(error) => failure(error), + } +} + +fn invalid(message: impl Into) -> NativeDataPlaneError { + NativeDataPlaneError { + kind: DataPlaneErrorKind::Io, + message: message.into(), + } +} + +fn socket_addr(address: DataPlaneSocketAddr) -> NativeDataPlaneResult { + let ip = match address.family { + 4 => IpAddr::V4(Ipv4Addr::new( + address.address[0], + address.address[1], + address.address[2], + address.address[3], + )), + 6 => { + return Err(NativeDataPlaneError { + kind: DataPlaneErrorKind::AddressFamilyUnsupported, + message: "IPv6 is not supported by data-plane ABI v2".to_string(), + }); + } + family => { + return Err(NativeDataPlaneError { + kind: DataPlaneErrorKind::AddressFamilyUnsupported, + message: format!("unsupported address family {family}"), + }); + } + }; + Ok(SocketAddr::new(ip, address.port)) +} + +fn ffi_socket_addr(address: SocketAddr) -> DataPlaneSocketAddr { + match address.ip() { + IpAddr::V4(ip) => { + let mut bytes = [0; 16]; + bytes[..4].copy_from_slice(&ip.octets()); + DataPlaneSocketAddr { + family: 4, + port: address.port(), + address: bytes, + } + } + IpAddr::V6(ip) => DataPlaneSocketAddr { + family: 6, + port: address.port(), + address: ip.octets(), + }, + } +} + +unsafe fn copy_input(ptr: *const c_uchar, len: u32) -> NativeDataPlaneResult> { + if len == 0 { + return Ok(Vec::new()); + } + if ptr.is_null() { + return Err(invalid("input buffer is null")); + } + Ok(unsafe { std::slice::from_raw_parts(ptr, len as usize) }.to_vec()) +} + +unsafe fn output_slice<'a>(ptr: *mut c_uchar, len: u32) -> NativeDataPlaneResult<&'a mut [u8]> { + if len == 0 { + return Ok(&mut []); + } + if ptr.is_null() { + return Err(invalid("output buffer is null")); + } + Ok(unsafe { std::slice::from_raw_parts_mut(ptr, len as usize) }) +} + +fn write_operation( + out_operation: *mut u64, + submit: impl FnOnce() -> NativeDataPlaneResult, +) -> c_int { + if out_operation.is_null() { + return failure(invalid("out_operation is null")); + } + match submit() { + Ok(operation) => { + unsafe { + *out_operation = operation; + } + 0 + } + Err(error) => failure(error), + } +} + +/// # Safety +/// +/// If non-null, `inst_name` must point to a valid NUL-terminated string. +/// `out_session` must be null or point to writable, properly aligned storage +/// for one `u64`. +#[cfg_attr(feature = "c-abi", unsafe(no_mangle))] +pub unsafe extern "C" fn data_plane_session_open( + inst_name: *const c_char, + out_session: *mut u64, +) -> c_int { + if out_session.is_null() { + return failure(invalid("out_session is null")); + } + unsafe { + *out_session = 0; + } + let inst_name = match unsafe { c_str_to_string(inst_name, "inst_name") } { + Ok(inst_name) => inst_name, + Err(error) => return failure(invalid(error)), + }; + match session::open(&inst_name) { + Ok(handle) => { + unsafe { + *out_session = handle; + } + 0 + } + Err(error) => failure(error), + } +} + +#[cfg_attr(feature = "c-abi", unsafe(no_mangle))] +pub extern "C" fn data_plane_session_close(session: u64) -> c_int { + status(super::session::close(session)) +} + +/// # Safety +/// +/// `out_operation` must be null or point to writable, properly aligned +/// storage for one `u64`. +#[cfg_attr(feature = "c-abi", unsafe(no_mangle))] +pub unsafe extern "C" fn data_plane_tcp_connect_submit( + session: u64, + peer_addr: DataPlaneSocketAddr, + timeout_ms: u64, + out_operation: *mut u64, +) -> c_int { + let peer_addr = match socket_addr(peer_addr) { + Ok(address) => address, + Err(error) => return failure(error), + }; + write_operation(out_operation, || { + super::session::submit_tcp_connect(session, peer_addr, timeout_ms) + }) +} + +/// # Safety +/// +/// `out_operation` must be null or point to writable, properly aligned +/// storage for one `u64`. +#[cfg_attr(feature = "c-abi", unsafe(no_mangle))] +pub unsafe extern "C" fn data_plane_tcp_bind_submit( + session: u64, + local_port: u16, + timeout_ms: u64, + out_operation: *mut u64, +) -> c_int { + write_operation(out_operation, || { + super::session::submit_tcp_bind(session, local_port, timeout_ms) + }) +} + +/// # Safety +/// +/// `out_operation` must be null or point to writable, properly aligned +/// storage for one `u64`. +#[cfg_attr(feature = "c-abi", unsafe(no_mangle))] +pub unsafe extern "C" fn data_plane_tcp_accept_submit( + session: u64, + listener: u64, + timeout_ms: u64, + out_operation: *mut u64, +) -> c_int { + write_operation(out_operation, || { + super::session::submit_tcp_accept(session, listener, timeout_ms) + }) +} + +/// # Safety +/// +/// `out_operation` must be null or point to writable, properly aligned +/// storage for one `u64`. +#[cfg_attr(feature = "c-abi", unsafe(no_mangle))] +pub unsafe extern "C" fn data_plane_tcp_read_submit( + session: u64, + stream: u64, + max_len: u32, + timeout_ms: u64, + out_operation: *mut u64, +) -> c_int { + write_operation(out_operation, || { + super::session::submit_tcp_read(session, stream, max_len, timeout_ms) + }) +} + +/// # Safety +/// +/// When `len` is nonzero, `data` must point to `len` readable bytes. +/// `out_operation` must be null or point to writable, properly aligned +/// storage for one `u64`. +#[cfg_attr(feature = "c-abi", unsafe(no_mangle))] +pub unsafe extern "C" fn data_plane_tcp_write_submit( + session: u64, + stream: u64, + data: *const c_uchar, + len: u32, + timeout_ms: u64, + out_operation: *mut u64, +) -> c_int { + let data = match unsafe { copy_input(data, len) } { + Ok(data) => data, + Err(error) => return failure(error), + }; + write_operation(out_operation, || { + super::session::submit_tcp_write(session, stream, data, timeout_ms) + }) +} + +/// # Safety +/// +/// `out_operation` must be null or point to writable, properly aligned +/// storage for one `u64`. +#[cfg_attr(feature = "c-abi", unsafe(no_mangle))] +pub unsafe extern "C" fn data_plane_udp_bind_submit( + session: u64, + local_port: u16, + timeout_ms: u64, + out_operation: *mut u64, +) -> c_int { + write_operation(out_operation, || { + super::session::submit_udp_bind(session, local_port, timeout_ms) + }) +} + +/// # Safety +/// +/// `out_operation` must be null or point to writable, properly aligned +/// storage for one `u64`. +#[cfg_attr(feature = "c-abi", unsafe(no_mangle))] +pub unsafe extern "C" fn data_plane_udp_receive_submit( + session: u64, + socket: u64, + max_len: u32, + timeout_ms: u64, + out_operation: *mut u64, +) -> c_int { + write_operation(out_operation, || { + super::session::submit_udp_receive(session, socket, max_len, timeout_ms) + }) +} + +/// # Safety +/// +/// When `len` is nonzero, `data` must point to `len` readable bytes. +/// `out_operation` must be null or point to writable, properly aligned +/// storage for one `u64`. +#[cfg_attr(feature = "c-abi", unsafe(no_mangle))] +pub unsafe extern "C" fn data_plane_udp_send_submit( + session: u64, + socket: u64, + peer_addr: DataPlaneSocketAddr, + data: *const c_uchar, + len: u32, + timeout_ms: u64, + out_operation: *mut u64, +) -> c_int { + let peer_addr = match socket_addr(peer_addr) { + Ok(address) => address, + Err(error) => return failure(error), + }; + let data = match unsafe { copy_input(data, len) } { + Ok(data) => data, + Err(error) => return failure(error), + }; + write_operation(out_operation, || { + super::session::submit_udp_send(session, socket, peer_addr, data, timeout_ms) + }) +} + +#[cfg_attr(feature = "c-abi", unsafe(no_mangle))] +pub extern "C" fn data_plane_operation_cancel(session: u64, operation: u64) -> c_int { + status(super::session::cancel_operation(session, operation)) +} + +#[cfg_attr(feature = "c-abi", unsafe(no_mangle))] +pub extern "C" fn data_plane_operation_free(session: u64, operation: u64) -> c_int { + status(super::session::free_operation(session, operation)) +} + +#[cfg_attr(feature = "c-abi", unsafe(no_mangle))] +pub extern "C" fn data_plane_resource_close(session: u64, resource: u64) -> c_int { + status(super::session::close_resource(session, resource)) +} + +#[cfg_attr(feature = "c-abi", unsafe(no_mangle))] +pub extern "C" fn data_plane_completion_wait(session: u64, timeout_ms: u64) -> c_int { + match super::session::completion_wait(session, timeout_ms) { + Ok(true) => 1, + Ok(false) => 0, + Err(error) => failure(error), + } +} + +/// # Safety +/// +/// When `capacity` is nonzero, `completions` must point to writable, properly +/// aligned storage for `capacity` consecutive [`DataPlaneCompletion`] values. +#[cfg_attr(feature = "c-abi", unsafe(no_mangle))] +pub unsafe extern "C" fn data_plane_completion_drain( + session: u64, + completions: *mut DataPlaneCompletion, + capacity: u32, +) -> c_int { + if capacity != 0 && completions.is_null() { + return failure(invalid("completions is null")); + } + let drained = match super::session::drain_completions(session, capacity as usize) { + Ok(drained) => drained, + Err(error) => return failure(error), + }; + for (index, completion) in drained.iter().enumerate() { + unsafe { + ptr::write( + completions.add(index), + DataPlaneCompletion { + operation_id: completion.operation_id.get(), + operation_kind: completion.kind as u16, + status: completion.status.code(), + }, + ); + } + } + drained.len() as c_int +} + +/// # Safety +/// +/// `out_size` must be null or point to writable, properly aligned storage for +/// one `u32`. +#[cfg_attr(feature = "c-abi", unsafe(no_mangle))] +pub unsafe extern "C" fn data_plane_result_size( + session: u64, + operation: u64, + out_size: *mut u32, +) -> c_int { + if out_size.is_null() { + return failure(invalid("out_size is null")); + } + match super::session::result_size(session, operation) { + Ok(size) => match u32::try_from(size) { + Ok(size) => { + unsafe { + *out_size = size; + } + 0 + } + Err(_) => failure(invalid("data-plane result size exceeds u32")), + }, + Err(error) => failure(error), + } +} + +/// # Safety +/// +/// Each output pointer must be null or point to writable, properly aligned +/// storage for its pointee type. Non-null output locations must not overlap. +#[cfg_attr(feature = "c-abi", unsafe(no_mangle))] +pub unsafe extern "C" fn data_plane_tcp_connect_result_take( + session: u64, + operation: u64, + out_stream: *mut u64, + out_local_addr: *mut DataPlaneSocketAddr, + out_peer_addr: *mut DataPlaneSocketAddr, +) -> c_int { + if out_stream.is_null() || out_local_addr.is_null() || out_peer_addr.is_null() { + return failure(invalid("TCP connect result output pointer is null")); + } + match super::session::take_tcp_connect(session, operation) { + Ok(result) => { + unsafe { + *out_stream = result.stream; + *out_local_addr = ffi_socket_addr(result.local_addr); + *out_peer_addr = ffi_socket_addr(result.peer_addr); + } + 0 + } + Err(error) => failure(error), + } +} + +/// # Safety +/// +/// Each output pointer must be null or point to writable, properly aligned +/// storage for its pointee type. Non-null output locations must not overlap. +#[cfg_attr(feature = "c-abi", unsafe(no_mangle))] +pub unsafe extern "C" fn data_plane_tcp_bind_result_take( + session: u64, + operation: u64, + out_listener: *mut u64, + out_local_addr: *mut DataPlaneSocketAddr, +) -> c_int { + if out_listener.is_null() || out_local_addr.is_null() { + return failure(invalid("TCP bind result output pointer is null")); + } + match super::session::take_tcp_bind(session, operation) { + Ok(result) => { + unsafe { + *out_listener = result.listener; + *out_local_addr = ffi_socket_addr(result.local_addr); + } + 0 + } + Err(error) => failure(error), + } +} + +/// # Safety +/// +/// Each output pointer must be null or point to writable, properly aligned +/// storage for its pointee type. Non-null output locations must not overlap. +#[cfg_attr(feature = "c-abi", unsafe(no_mangle))] +pub unsafe extern "C" fn data_plane_tcp_accept_result_take( + session: u64, + operation: u64, + out_stream: *mut u64, + out_local_addr: *mut DataPlaneSocketAddr, + out_peer_addr: *mut DataPlaneSocketAddr, +) -> c_int { + if out_stream.is_null() || out_local_addr.is_null() || out_peer_addr.is_null() { + return failure(invalid("TCP accept result output pointer is null")); + } + match super::session::take_tcp_accept(session, operation) { + Ok(result) => { + unsafe { + *out_stream = result.stream; + *out_local_addr = ffi_socket_addr(result.local_addr); + *out_peer_addr = ffi_socket_addr(result.peer_addr); + } + 0 + } + Err(error) => failure(error), + } +} + +/// # Safety +/// +/// When `capacity` is nonzero, `data` must point to `capacity` writable bytes. +/// Each scalar output pointer must be null or point to writable, properly +/// aligned storage for its pointee type. Non-null output ranges must not +/// overlap. +#[cfg_attr(feature = "c-abi", unsafe(no_mangle))] +pub unsafe extern "C" fn data_plane_tcp_read_result_take( + session: u64, + operation: u64, + data: *mut c_uchar, + capacity: u32, + out_len: *mut u32, + out_eof: *mut bool, +) -> c_int { + if out_len.is_null() || out_eof.is_null() { + return failure(invalid("TCP read result output pointer is null")); + } + let data = match unsafe { output_slice(data, capacity) } { + Ok(data) => data, + Err(error) => return failure(error), + }; + match super::session::take_tcp_read(session, operation, data) { + Ok(result) => { + unsafe { + *out_len = result.len as u32; + *out_eof = result.eof; + } + 0 + } + Err(error) => failure(error), + } +} + +/// # Safety +/// +/// `out_len` must be null or point to writable, properly aligned storage for +/// one `u32`. +#[cfg_attr(feature = "c-abi", unsafe(no_mangle))] +pub unsafe extern "C" fn data_plane_tcp_write_result_take( + session: u64, + operation: u64, + out_len: *mut u32, +) -> c_int { + if out_len.is_null() { + return failure(invalid("out_len is null")); + } + match super::session::take_tcp_write(session, operation) { + Ok(len) => match u32::try_from(len) { + Ok(len) => { + unsafe { + *out_len = len; + } + 0 + } + Err(_) => failure(invalid("TCP write result exceeds u32")), + }, + Err(error) => failure(error), + } +} + +/// # Safety +/// +/// Each output pointer must be null or point to writable, properly aligned +/// storage for its pointee type. Non-null output locations must not overlap. +#[cfg_attr(feature = "c-abi", unsafe(no_mangle))] +pub unsafe extern "C" fn data_plane_udp_bind_result_take( + session: u64, + operation: u64, + out_socket: *mut u64, + out_local_addr: *mut DataPlaneSocketAddr, +) -> c_int { + if out_socket.is_null() || out_local_addr.is_null() { + return failure(invalid("UDP bind result output pointer is null")); + } + match super::session::take_udp_bind(session, operation) { + Ok(result) => { + unsafe { + *out_socket = result.socket; + *out_local_addr = ffi_socket_addr(result.local_addr); + } + 0 + } + Err(error) => failure(error), + } +} + +/// # Safety +/// +/// When `capacity` is nonzero, `data` must point to `capacity` writable bytes. +/// Each scalar output pointer must be null or point to writable, properly +/// aligned storage for its pointee type. Non-null output ranges must not +/// overlap. +#[cfg_attr(feature = "c-abi", unsafe(no_mangle))] +pub unsafe extern "C" fn data_plane_udp_receive_result_take( + session: u64, + operation: u64, + data: *mut c_uchar, + capacity: u32, + out_len: *mut u32, + out_peer_addr: *mut DataPlaneSocketAddr, + out_truncated: *mut bool, +) -> c_int { + if out_len.is_null() || out_peer_addr.is_null() || out_truncated.is_null() { + return failure(invalid("UDP receive result output pointer is null")); + } + let data = match unsafe { output_slice(data, capacity) } { + Ok(data) => data, + Err(error) => return failure(error), + }; + match super::session::take_udp_receive(session, operation, data) { + Ok(result) => { + unsafe { + *out_len = result.len as u32; + *out_peer_addr = ffi_socket_addr(result.peer_addr); + *out_truncated = result.truncated; + } + 0 + } + Err(error) => failure(error), + } +} + +/// # Safety +/// +/// `out_len` must be null or point to writable, properly aligned storage for +/// one `u32`. +#[cfg_attr(feature = "c-abi", unsafe(no_mangle))] +pub unsafe extern "C" fn data_plane_udp_send_result_take( + session: u64, + operation: u64, + out_len: *mut u32, +) -> c_int { + if out_len.is_null() { + return failure(invalid("out_len is null")); + } + match super::session::take_udp_send(session, operation) { + Ok(len) => match u32::try_from(len) { + Ok(len) => { + unsafe { + *out_len = len; + } + 0 + } + Err(_) => failure(invalid("UDP send result exceeds u32")), + }, + Err(error) => failure(error), + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn socket_address_round_trip() { + let address = "127.0.0.1:1234".parse::().unwrap(); + assert_eq!(socket_addr(ffi_socket_addr(address)).unwrap(), address); + } + + #[test] + fn ipv6_is_rejected_by_v2() { + let error = socket_addr(ffi_socket_addr( + "[2001:db8::1]:4321".parse::().unwrap(), + )) + .unwrap_err(); + assert_eq!(error.kind, DataPlaneErrorKind::AddressFamilyUnsupported); + } + + #[test] + fn invalid_address_family_is_stable() { + let error = socket_addr(DataPlaneSocketAddr { + family: 9, + ..Default::default() + }) + .unwrap_err(); + assert_eq!(error.kind, DataPlaneErrorKind::AddressFamilyUnsupported); + } + + #[test] + fn null_operation_output_does_not_submit() { + let submitted = std::cell::Cell::new(false); + + assert_eq!( + write_operation(std::ptr::null_mut(), || { + submitted.set(true); + Ok(1) + }), + -(DataPlaneErrorKind::Io as c_int) + ); + assert!(!submitted.get()); + } +} diff --git a/easytier-contrib/easytier-ffi/src/data_plane/mod.rs b/easytier-contrib/easytier-ffi/src/data_plane/mod.rs new file mode 100644 index 00000000..c02df927 --- /dev/null +++ b/easytier-contrib/easytier-ffi/src/data_plane/mod.rs @@ -0,0 +1,16 @@ +//! Native C ABI adapter for the instance-scoped data-plane operation broker. + +#[cfg(feature = "ffi-dataplane")] +mod abi; +#[cfg(feature = "ffi-dataplane")] +mod session; + +#[cfg(feature = "ffi-dataplane")] +pub use abi::*; +#[cfg(feature = "ffi-dataplane")] +pub(crate) use session::{ + lock_for_config_server_start, remove_data_plane_sessions_by_instance_ids, +}; + +#[cfg(not(feature = "ffi-dataplane"))] +pub(crate) fn remove_data_plane_sessions_by_instance_ids(_ids: &[uuid::Uuid]) {} diff --git a/easytier-contrib/easytier-ffi/src/data_plane/session.rs b/easytier-contrib/easytier-ffi/src/data_plane/session.rs new file mode 100644 index 00000000..eafb4fd7 --- /dev/null +++ b/easytier-contrib/easytier-ffi/src/data_plane/session.rs @@ -0,0 +1,636 @@ +use std::{ + collections::HashMap, + net::SocketAddr, + sync::{ + Arc, Mutex, RwLock, + atomic::{AtomicBool, AtomicU64, Ordering}, + }, + time::Duration, +}; + +use easytier::instance::host::NativeInstanceHost; +use easytier_core::gateway::{ + DataPlaneCompletionDescriptor, DataPlaneError, DataPlaneErrorKind, DataPlaneOperationId, + DataPlaneOperationKind, DataPlaneOperationResult, DataPlaneResourceId, DataPlaneSession, +}; +use uuid::Uuid; + +use crate::{ + config_server::{in_config_server_callback, is_config_server_active_or_stopping}, + state::{ffi_context, resolve_instance_id_by_name}, +}; + +type CoreDataPlaneSession = DataPlaneSession; + +static NEXT_SESSION_HANDLE: AtomicU64 = AtomicU64::new(1); +static SESSIONS: once_cell::sync::Lazy>>> = + once_cell::sync::Lazy::new(|| Mutex::new(HashMap::new())); +static DATA_PLANE_USAGE_LOCK: once_cell::sync::Lazy> = + once_cell::sync::Lazy::new(|| RwLock::new(())); + +#[derive(Debug)] +pub(super) struct NativeDataPlaneError { + pub(super) kind: DataPlaneErrorKind, + pub(super) message: String, +} + +impl NativeDataPlaneError { + fn new(kind: DataPlaneErrorKind, message: impl Into) -> Self { + Self { + kind, + message: message.into(), + } + } + + fn invalid(message: impl Into) -> Self { + Self::new(DataPlaneErrorKind::Io, message) + } + + fn closed(message: impl Into) -> Self { + Self::new(DataPlaneErrorKind::HandleClosed, message) + } +} + +impl From for NativeDataPlaneError { + fn from(error: DataPlaneError) -> Self { + Self::new(error.kind(), error.message()) + } +} + +pub(super) type NativeDataPlaneResult = Result; + +pub(super) struct TcpConnectResult { + pub(super) stream: u64, + pub(super) local_addr: SocketAddr, + pub(super) peer_addr: SocketAddr, +} + +pub(super) struct TcpBindResult { + pub(super) listener: u64, + pub(super) local_addr: SocketAddr, +} + +pub(super) struct TcpAcceptResult { + pub(super) stream: u64, + pub(super) local_addr: SocketAddr, + pub(super) peer_addr: SocketAddr, +} + +pub(super) struct TcpReadResult { + pub(super) len: usize, + pub(super) eof: bool, +} + +pub(super) struct UdpBindResult { + pub(super) socket: u64, + pub(super) local_addr: SocketAddr, +} + +pub(super) struct UdpReceiveResult { + pub(super) len: usize, + pub(super) peer_addr: SocketAddr, + pub(super) truncated: bool, +} + +struct NativeDataPlaneSession { + instance_id: Uuid, + runtime: tokio::runtime::Handle, + core: Arc, + submit_gate: Mutex<()>, + closed: AtomicBool, +} + +impl NativeDataPlaneSession { + fn close(&self) { + let _gate = self + .submit_gate + .lock() + .unwrap_or_else(|error| error.into_inner()); + if self.closed.swap(true, Ordering::AcqRel) { + return; + } + self.core.discard_all(); + } + + fn submit( + &self, + submit: impl FnOnce(&Arc) -> Result, + ) -> NativeDataPlaneResult { + let _gate = self + .submit_gate + .lock() + .map_err(|error| NativeDataPlaneError::invalid(error.to_string()))?; + if self.closed.load(Ordering::Acquire) { + return Err(NativeDataPlaneError::closed( + "native data-plane session is closed", + )); + } + let _runtime = self.runtime.enter(); + submit(&self.core) + .map(DataPlaneOperationId::get) + .map_err(Into::into) + } +} + +fn sessions() +-> NativeDataPlaneResult>>> +{ + SESSIONS + .lock() + .map_err(|error| NativeDataPlaneError::invalid(error.to_string())) +} + +fn get_session(handle: u64) -> NativeDataPlaneResult> { + if handle == 0 { + return Err(NativeDataPlaneError::closed( + "native data-plane session handle is invalid", + )); + } + let session = sessions()? + .get(&handle) + .cloned() + .ok_or_else(|| NativeDataPlaneError::closed("native data-plane session is closed"))?; + if session.closed.load(Ordering::Acquire) { + return Err(NativeDataPlaneError::closed( + "native data-plane session is closed", + )); + } + Ok(session) +} + +fn next_session_handle( + sessions: &HashMap>, +) -> NativeDataPlaneResult { + for _ in 0..sessions.len().saturating_add(2) { + let handle = NEXT_SESSION_HANDLE.fetch_add(1, Ordering::Relaxed); + if handle != 0 && !sessions.contains_key(&handle) { + return Ok(handle); + } + } + Err(NativeDataPlaneError::new( + DataPlaneErrorKind::ResourceLimit, + "native data-plane session handle space is exhausted", + )) +} + +fn reject_data_plane_use() -> NativeDataPlaneResult<()> { + if in_config_server_callback() { + Err(NativeDataPlaneError::invalid( + "cannot use data plane from config server callback", + )) + } else if is_config_server_active_or_stopping() { + Err(NativeDataPlaneError::invalid( + "cannot use data plane while config server client is active", + )) + } else { + Ok(()) + } +} + +pub(super) fn open(inst_name: &str) -> NativeDataPlaneResult { + reject_data_plane_use()?; + let _usage = DATA_PLANE_USAGE_LOCK + .read() + .map_err(|error| NativeDataPlaneError::invalid(error.to_string()))?; + reject_data_plane_use()?; + + let instance_id = resolve_instance_id_by_name(inst_name) + .map_err(NativeDataPlaneError::invalid)? + .ok_or_else(|| NativeDataPlaneError::closed("instance not found"))?; + let manager = &ffi_context().manager; + let core = manager.data_plane_session(&instance_id).ok_or_else(|| { + NativeDataPlaneError::closed("instance data-plane session is unavailable") + })?; + let runtime = manager + .data_plane_runtime_handle(&instance_id) + .ok_or_else(|| NativeDataPlaneError::closed("instance runtime is unavailable"))?; + + let mut sessions = sessions()?; + if sessions + .values() + .any(|session| session.instance_id == instance_id) + { + return Err(NativeDataPlaneError::new( + DataPlaneErrorKind::ResourceLimit, + "instance already has an open native data-plane session", + )); + } + let handle = next_session_handle(&sessions)?; + sessions.insert( + handle, + Arc::new(NativeDataPlaneSession { + instance_id, + runtime, + core, + submit_gate: Mutex::new(()), + closed: AtomicBool::new(false), + }), + ); + Ok(handle) +} + +pub(super) fn close(handle: u64) -> NativeDataPlaneResult<()> { + let _usage = DATA_PLANE_USAGE_LOCK + .read() + .map_err(|error| NativeDataPlaneError::invalid(error.to_string()))?; + let mut sessions = sessions()?; + let session = sessions + .remove(&handle) + .ok_or_else(|| NativeDataPlaneError::closed("native data-plane session is closed"))?; + // Keep the registry locked until the shared core namespace is empty. An + // open for the same instance must not publish a replacement session before + // this old wrapper finishes discarding its operations and resources. + session.close(); + Ok(()) +} + +fn timeout(timeout_ms: u64) -> Option { + (timeout_ms != u64::MAX).then(|| Duration::from_millis(timeout_ms)) +} + +fn operation_id(raw: u64) -> NativeDataPlaneResult { + DataPlaneOperationId::from_raw(raw) + .ok_or_else(|| NativeDataPlaneError::closed("data-plane operation handle is invalid")) +} + +fn resource_id(raw: u64) -> NativeDataPlaneResult { + DataPlaneResourceId::from_raw(raw) + .ok_or_else(|| NativeDataPlaneError::closed("data-plane resource handle is invalid")) +} + +pub(super) fn submit_tcp_connect( + session: u64, + peer_addr: SocketAddr, + timeout_ms: u64, +) -> NativeDataPlaneResult { + get_session(session)?.submit(|core| core.submit_tcp_connect(peer_addr, timeout(timeout_ms))) +} + +pub(super) fn submit_tcp_bind( + session: u64, + local_port: u16, + timeout_ms: u64, +) -> NativeDataPlaneResult { + get_session(session)?.submit(|core| core.submit_tcp_bind(local_port, timeout(timeout_ms))) +} + +pub(super) fn submit_tcp_accept( + session: u64, + listener: u64, + timeout_ms: u64, +) -> NativeDataPlaneResult { + let listener = resource_id(listener)?; + get_session(session)?.submit(|core| core.submit_tcp_accept(listener, timeout(timeout_ms))) +} + +pub(super) fn submit_tcp_read( + session: u64, + stream: u64, + max_len: u32, + timeout_ms: u64, +) -> NativeDataPlaneResult { + let stream = resource_id(stream)?; + get_session(session)? + .submit(|core| core.submit_tcp_read(stream, max_len as usize, timeout(timeout_ms))) +} + +pub(super) fn submit_tcp_write( + session: u64, + stream: u64, + data: Vec, + timeout_ms: u64, +) -> NativeDataPlaneResult { + let stream = resource_id(stream)?; + get_session(session)?.submit(|core| core.submit_tcp_write(stream, data, timeout(timeout_ms))) +} + +pub(super) fn submit_udp_bind( + session: u64, + local_port: u16, + timeout_ms: u64, +) -> NativeDataPlaneResult { + get_session(session)?.submit(|core| core.submit_udp_bind(local_port, timeout(timeout_ms))) +} + +pub(super) fn submit_udp_receive( + session: u64, + socket: u64, + max_len: u32, + timeout_ms: u64, +) -> NativeDataPlaneResult { + let socket = resource_id(socket)?; + get_session(session)? + .submit(|core| core.submit_udp_receive(socket, max_len as usize, timeout(timeout_ms))) +} + +pub(super) fn submit_udp_send( + session: u64, + socket: u64, + peer_addr: SocketAddr, + data: Vec, + timeout_ms: u64, +) -> NativeDataPlaneResult { + let socket = resource_id(socket)?; + get_session(session)? + .submit(|core| core.submit_udp_send(socket, peer_addr, data, timeout(timeout_ms))) +} + +pub(super) fn cancel_operation(session: u64, operation: u64) -> NativeDataPlaneResult<()> { + let operation = operation_id(operation)?; + get_session(session)?.core.cancel_operation(operation); + Ok(()) +} + +pub(super) fn free_operation(session: u64, operation: u64) -> NativeDataPlaneResult<()> { + let operation = operation_id(operation)?; + get_session(session)?.core.free_operation(operation); + Ok(()) +} + +pub(super) fn close_resource(session: u64, resource: u64) -> NativeDataPlaneResult<()> { + let resource = resource_id(resource)?; + get_session(session)?.core.close_resource(resource); + Ok(()) +} + +pub(super) fn completion_wait(session: u64, timeout_ms: u64) -> NativeDataPlaneResult { + let session = get_session(session)?; + let ready = session.core.completion_wait(timeout(timeout_ms)); + Ok(ready && !session.closed.load(Ordering::Acquire)) +} + +pub(super) fn drain_completions( + session: u64, + max_count: usize, +) -> NativeDataPlaneResult> { + Ok(get_session(session)?.core.drain_completions(max_count)) +} + +pub(super) fn result_size(session: u64, operation: u64) -> NativeDataPlaneResult { + let operation = operation_id(operation)?; + get_session(session)? + .core + .result_payload_bytes(operation) + .map_err(Into::into) +} + +fn take_result( + session: u64, + operation: u64, + expected: DataPlaneOperationKind, + take: impl FnOnce(&DataPlaneOperationResult) -> Option, +) -> NativeDataPlaneResult { + let operation = operation_id(operation)?; + let session = get_session(session)?; + let actual = session.core.operation_kind(operation)?; + if actual != expected { + return Err(NativeDataPlaneError::invalid(format!( + "operation kind mismatch: expected {expected:?}, got {actual:?}" + ))); + } + let result = session.core.take_result_with(operation, |outcome| { + Some(match outcome { + Ok(result) => take(result).ok_or_else(|| { + NativeDataPlaneError::invalid("data-plane result variant does not match operation") + }), + Err(kind) => Err(NativeDataPlaneError::new( + *kind, + format!("data-plane operation failed with {kind:?}"), + )), + }) + })?; + result + .ok_or_else(|| NativeDataPlaneError::invalid("data-plane result could not be consumed"))? +} + +pub(super) fn take_tcp_connect( + session: u64, + operation: u64, +) -> NativeDataPlaneResult { + take_result( + session, + operation, + DataPlaneOperationKind::TcpConnect, + |result| match result { + DataPlaneOperationResult::TcpConnected { + stream, + local_addr, + peer_addr, + } => Some(TcpConnectResult { + stream: stream.get(), + local_addr: *local_addr, + peer_addr: *peer_addr, + }), + _ => None, + }, + ) +} + +pub(super) fn take_tcp_bind(session: u64, operation: u64) -> NativeDataPlaneResult { + take_result( + session, + operation, + DataPlaneOperationKind::TcpBind, + |result| match result { + DataPlaneOperationResult::TcpBound { + listener, + local_addr, + } => Some(TcpBindResult { + listener: listener.get(), + local_addr: *local_addr, + }), + _ => None, + }, + ) +} + +pub(super) fn take_tcp_accept( + session: u64, + operation: u64, +) -> NativeDataPlaneResult { + take_result( + session, + operation, + DataPlaneOperationKind::TcpAccept, + |result| match result { + DataPlaneOperationResult::TcpAccepted { + stream, + local_addr, + peer_addr, + } => Some(TcpAcceptResult { + stream: stream.get(), + local_addr: *local_addr, + peer_addr: *peer_addr, + }), + _ => None, + }, + ) +} + +pub(super) fn take_tcp_read( + session: u64, + operation: u64, + output: &mut [u8], +) -> NativeDataPlaneResult { + let required = result_size(session, operation)?; + if output.len() < required { + return Err(NativeDataPlaneError::new( + DataPlaneErrorKind::BufferTooSmall, + format!( + "TCP read result requires {required} bytes, buffer has {}", + output.len() + ), + )); + } + take_result( + session, + operation, + DataPlaneOperationKind::TcpRead, + |result| match result { + DataPlaneOperationResult::TcpRead { data, eof } => { + output[..data.len()].copy_from_slice(data); + Some(TcpReadResult { + len: data.len(), + eof: *eof, + }) + } + _ => None, + }, + ) +} + +pub(super) fn take_tcp_write(session: u64, operation: u64) -> NativeDataPlaneResult { + take_result( + session, + operation, + DataPlaneOperationKind::TcpWrite, + |result| match result { + DataPlaneOperationResult::TcpWritten { len } => Some(*len), + _ => None, + }, + ) +} + +pub(super) fn take_udp_bind(session: u64, operation: u64) -> NativeDataPlaneResult { + take_result( + session, + operation, + DataPlaneOperationKind::UdpBind, + |result| match result { + DataPlaneOperationResult::UdpBound { socket, local_addr } => Some(UdpBindResult { + socket: socket.get(), + local_addr: *local_addr, + }), + _ => None, + }, + ) +} + +pub(super) fn take_udp_receive( + session: u64, + operation: u64, + output: &mut [u8], +) -> NativeDataPlaneResult { + let required = result_size(session, operation)?; + if output.len() < required { + return Err(NativeDataPlaneError::new( + DataPlaneErrorKind::BufferTooSmall, + format!( + "UDP receive result requires {required} bytes, buffer has {}", + output.len() + ), + )); + } + take_result( + session, + operation, + DataPlaneOperationKind::UdpReceive, + |result| match result { + DataPlaneOperationResult::UdpReceived { + data, + peer_addr, + truncated, + } => { + output[..data.len()].copy_from_slice(data); + Some(UdpReceiveResult { + len: data.len(), + peer_addr: *peer_addr, + truncated: *truncated, + }) + } + _ => None, + }, + ) +} + +pub(super) fn take_udp_send(session: u64, operation: u64) -> NativeDataPlaneResult { + take_result( + session, + operation, + DataPlaneOperationKind::UdpSend, + |result| match result { + DataPlaneOperationResult::UdpSent { len } => Some(*len), + _ => None, + }, + ) +} + +pub(crate) fn remove_data_plane_sessions_by_instance_ids(ids: &[Uuid]) { + if ids.is_empty() { + return; + } + let _usage = DATA_PLANE_USAGE_LOCK + .write() + .unwrap_or_else(|error| error.into_inner()); + let removed = { + let mut sessions = SESSIONS.lock().unwrap_or_else(|error| error.into_inner()); + let handles = sessions + .iter() + .filter_map(|(handle, session)| ids.contains(&session.instance_id).then_some(*handle)) + .collect::>(); + handles + .into_iter() + .filter_map(|handle| sessions.remove(&handle)) + .collect::>() + }; + for session in removed { + session.close(); + } +} + +pub(crate) fn lock_for_config_server_start() +-> Result, String> { + let guard = DATA_PLANE_USAGE_LOCK + .write() + .map_err(|error| format!("failed to lock data plane usage: {error}"))?; + if !SESSIONS + .lock() + .map_err(|error| format!("failed to lock data-plane sessions: {error}"))? + .is_empty() + { + return Err("cannot start config server client while data plane is in use".to_string()); + } + Ok(guard) +} + +#[cfg(test)] +mod tests { + use std::{sync::mpsc, time::Duration}; + + use super::*; + + #[test] + fn config_server_start_waits_for_session_open_or_close() { + let read_guard = DATA_PLANE_USAGE_LOCK.read().unwrap(); + let (done_tx, done_rx) = mpsc::channel(); + let waiter = std::thread::spawn(move || { + let _write_guard = lock_for_config_server_start().unwrap(); + done_tx.send(()).unwrap(); + }); + + assert!(done_rx.recv_timeout(Duration::from_millis(100)).is_err()); + drop(read_guard); + done_rx.recv_timeout(Duration::from_secs(5)).unwrap(); + waiter.join().unwrap(); + } +} diff --git a/easytier-contrib/easytier-ffi/src/data_plane_async.rs b/easytier-contrib/easytier-ffi/src/data_plane_async.rs deleted file mode 100644 index 41241b7e..00000000 --- a/easytier-contrib/easytier-ffi/src/data_plane_async.rs +++ /dev/null @@ -1,1162 +0,0 @@ -#[cfg(feature = "ffi-dataplane")] -use std::{ - future::Future, - net::SocketAddr, - sync::{ - Arc, Condvar, Mutex, - atomic::{AtomicU64, Ordering}, - }, - time::{Duration, Instant}, -}; - -#[cfg(feature = "ffi-dataplane")] -use dashmap::DashMap; -#[cfg(feature = "ffi-dataplane")] -use easytier::launcher::{DataPlaneTcpListener, DataPlaneTcpStream, DataPlaneUdpSocket}; -#[cfg(feature = "ffi-dataplane")] -use tokio::io::{AsyncReadExt, AsyncWriteExt}; -#[cfg(feature = "ffi-dataplane")] -use tokio_util::sync::CancellationToken; -#[cfg(feature = "ffi-dataplane")] -use uuid::Uuid; - -#[cfg(feature = "ffi-dataplane")] -use crate::{ - data_plane::{ - TcpHalves, cstr_to_string, enter_data_plane_operation, get_instance_id, get_tcp_listener, - get_tcp_stream_with_instance, get_udp_socket_with_instance, insert_tcp_listener_handle, - insert_tcp_stream_handle, insert_udp_socket_handle, into_ffi_ip_cstring, parse_socket_addr, - timeout_duration, - }, - error::{free_string, set_error_msg}, - state::{ASYNC_RUNTIME, INSTANCE_MANAGER}, -}; - -#[cfg(feature = "ffi-dataplane")] -pub(crate) const DATA_PLANE_OP_PENDING: std::ffi::c_int = 0; -#[cfg(feature = "ffi-dataplane")] -pub(crate) const DATA_PLANE_OP_READY: std::ffi::c_int = 1; -#[cfg(feature = "ffi-dataplane")] -pub(crate) const DATA_PLANE_OP_FAILED: std::ffi::c_int = -1; -#[cfg(feature = "ffi-dataplane")] -pub(crate) const DATA_PLANE_OP_INVALID: std::ffi::c_int = -2; - -#[cfg(feature = "ffi-dataplane")] -static NEXT_DATA_PLANE_OP: AtomicU64 = AtomicU64::new(1); -#[cfg(feature = "ffi-dataplane")] -static DATA_PLANE_OPS: once_cell::sync::Lazy>> = - once_cell::sync::Lazy::new(DashMap::new); - -#[cfg(feature = "ffi-dataplane")] -const MAX_ASYNC_READ_LEN: u32 = 16 * 1024 * 1024; -#[cfg(feature = "ffi-dataplane")] -const MAX_ASYNC_WRITE_LEN: u32 = std::ffi::c_int::MAX as u32; - -#[cfg(feature = "ffi-dataplane")] -#[derive(Clone, Copy, Debug, Eq, PartialEq)] -enum DataPlaneAsyncOpKind { - TcpConnect, - TcpBind, - TcpAccept, - TcpRead, - TcpWrite, - UdpBind, - UdpSendTo, - UdpRecvFrom, -} - -#[cfg(feature = "ffi-dataplane")] -struct DataPlaneAsyncOp { - kind: DataPlaneAsyncOpKind, - instance_id: Option, - target_handle: Option, - cancel_token: CancellationToken, - state: Mutex, - ready: Condvar, -} - -#[cfg(feature = "ffi-dataplane")] -enum DataPlaneAsyncOpState { - Pending, - Ready(Box), - Failed(String), - Consumed, -} - -#[cfg(feature = "ffi-dataplane")] -enum DataPlaneAsyncOpResult { - TcpConnect { - instance_id: Uuid, - runtime: tokio::runtime::Handle, - stream: DataPlaneTcpStream, - local_addr: SocketAddr, - }, - TcpBind { - instance_id: Uuid, - runtime: tokio::runtime::Handle, - listener: DataPlaneTcpListener, - local_addr: SocketAddr, - }, - TcpAccept { - instance_id: Uuid, - runtime: tokio::runtime::Handle, - stream: DataPlaneTcpStream, - local_addr: SocketAddr, - peer_addr: SocketAddr, - }, - TcpRead { - data: Vec, - }, - TcpWrite { - written: usize, - }, - UdpBind { - instance_id: Uuid, - runtime: tokio::runtime::Handle, - socket: DataPlaneUdpSocket, - local_addr: SocketAddr, - }, - UdpSendTo { - sent: usize, - }, - UdpRecvFrom { - data: Vec, - peer_addr: SocketAddr, - }, -} - -#[cfg(feature = "ffi-dataplane")] -fn next_op_handle() -> u64 { - NEXT_DATA_PLANE_OP.fetch_add(1, Ordering::Relaxed) -} - -#[cfg(feature = "ffi-dataplane")] -fn new_op( - kind: DataPlaneAsyncOpKind, - instance_id: Option, - target_handle: Option, -) -> (u64, Arc) { - let handle = next_op_handle(); - let op = Arc::new(DataPlaneAsyncOp { - kind, - instance_id, - target_handle, - cancel_token: CancellationToken::new(), - state: Mutex::new(DataPlaneAsyncOpState::Pending), - ready: Condvar::new(), - }); - DATA_PLANE_OPS.insert(handle, op.clone()); - (handle, op) -} - -#[cfg(feature = "ffi-dataplane")] -fn complete_op(op: &DataPlaneAsyncOp, result: Result) { - let Ok(mut state) = op.state.lock() else { - return; - }; - if !matches!(*state, DataPlaneAsyncOpState::Pending) { - return; - } - *state = match result { - Ok(result) => DataPlaneAsyncOpState::Ready(Box::new(result)), - Err(err) => DataPlaneAsyncOpState::Failed(err), - }; - op.ready.notify_all(); -} - -#[cfg(feature = "ffi-dataplane")] -fn cancel_pending_op(op: &DataPlaneAsyncOp, reason: &str) { - op.cancel_token.cancel(); - let Ok(mut state) = op.state.lock() else { - return; - }; - if matches!(*state, DataPlaneAsyncOpState::Pending) { - *state = DataPlaneAsyncOpState::Failed(reason.to_string()); - op.ready.notify_all(); - } -} - -#[cfg(feature = "ffi-dataplane")] -fn validate_max_len(max_len: u32) -> bool { - if max_len > MAX_ASYNC_READ_LEN { - set_error_msg(&format!( - "max_len exceeds async data plane limit of {} bytes", - MAX_ASYNC_READ_LEN - )); - false - } else { - true - } -} - -#[cfg(feature = "ffi-dataplane")] -fn validate_write_len(len: u32) -> bool { - if len > MAX_ASYNC_WRITE_LEN { - set_error_msg(&format!( - "len exceeds async data plane write limit of {} bytes", - MAX_ASYNC_WRITE_LEN - )); - false - } else { - true - } -} - -#[cfg(feature = "ffi-dataplane")] -fn usize_to_c_int(value: usize, name: &str) -> Option { - if value > std::ffi::c_int::MAX as usize { - set_error_msg(&format!( - "{} exceeds c_int limit of {}", - name, - std::ffi::c_int::MAX - )); - None - } else { - Some(value as std::ffi::c_int) - } -} - -#[cfg(feature = "ffi-dataplane")] -fn spawn_instance_runtime_op( - op: Arc, - instance_id: Uuid, - timeout_ms: u64, - build: F, -) where - Fut: Future> + Send + 'static, - F: FnOnce(tokio::runtime::Handle, Duration) -> Fut + Send + 'static, -{ - let deadline = Instant::now() + timeout_duration(timeout_ms); - ASYNC_RUNTIME.spawn_blocking(move || { - let runtime = loop { - if op.cancel_token.is_cancelled() { - complete_op(&op, Err("data plane async op canceled".to_string())); - return; - } - - let remaining = deadline.saturating_duration_since(Instant::now()); - let wait_for = remaining.min(Duration::from_millis(50)); - if let Some(runtime) = - INSTANCE_MANAGER.data_plane_wait_runtime_handle(&instance_id, wait_for) - { - break runtime; - } - if remaining.is_zero() || Instant::now() >= deadline { - complete_op(&op, Err("instance runtime is not ready".to_string())); - return; - } - }; - if op.cancel_token.is_cancelled() { - complete_op(&op, Err("data plane async op canceled".to_string())); - return; - } - - let runtime_for_task = runtime.clone(); - let op_for_complete = op.clone(); - runtime.spawn(async move { - let remaining = deadline.saturating_duration_since(Instant::now()); - let result = build(runtime_for_task, remaining).await; - complete_op(&op_for_complete, result); - }); - }); -} - -#[cfg(feature = "ffi-dataplane")] -fn state_status(state: &DataPlaneAsyncOpState) -> std::ffi::c_int { - match state { - DataPlaneAsyncOpState::Pending => DATA_PLANE_OP_PENDING, - DataPlaneAsyncOpState::Ready(_) => DATA_PLANE_OP_READY, - DataPlaneAsyncOpState::Failed(_) => DATA_PLANE_OP_FAILED, - DataPlaneAsyncOpState::Consumed => DATA_PLANE_OP_INVALID, - } -} - -#[cfg(feature = "ffi-dataplane")] -fn take_completed_op( - handle: u64, - expected: DataPlaneAsyncOpKind, -) -> Option { - let Some((_, op)) = DATA_PLANE_OPS.remove_if(&handle, |_, op| { - if op.kind != expected { - return false; - } - let Ok(state) = op.state.lock() else { - return false; - }; - match &*state { - DataPlaneAsyncOpState::Ready(_) | DataPlaneAsyncOpState::Failed(_) => true, - DataPlaneAsyncOpState::Pending | DataPlaneAsyncOpState::Consumed => false, - } - }) else { - let Some(op) = DATA_PLANE_OPS.get(&handle).map(|op| op.clone()) else { - set_error_msg("data plane async op not found"); - return None; - }; - if op.kind != expected { - set_error_msg("data plane async op type mismatch"); - return None; - } - let Ok(state) = op.state.lock() else { - set_error_msg("failed to lock data plane async op"); - return None; - }; - match &*state { - DataPlaneAsyncOpState::Pending => { - set_error_msg("data plane async op is still pending"); - return None; - } - DataPlaneAsyncOpState::Consumed => { - set_error_msg("data plane async op already consumed"); - return None; - } - DataPlaneAsyncOpState::Ready(_) | DataPlaneAsyncOpState::Failed(_) => { - set_error_msg("data plane async op was consumed concurrently"); - } - } - return None; - }; - - let completed = { - let Ok(mut state) = op.state.lock() else { - set_error_msg("failed to lock data plane async op"); - return None; - }; - std::mem::replace(&mut *state, DataPlaneAsyncOpState::Consumed) - }; - - match completed { - DataPlaneAsyncOpState::Ready(result) => Some(*result), - DataPlaneAsyncOpState::Failed(err) => { - set_error_msg(&err); - None - } - DataPlaneAsyncOpState::Pending | DataPlaneAsyncOpState::Consumed => None, - } -} - -#[cfg(feature = "ffi-dataplane")] -async fn run_with_cancel( - cancel_token: &CancellationToken, - error_prefix: &str, - op: F, -) -> Result -where - E: std::fmt::Display, - F: Future>, -{ - tokio::select! { - biased; - _ = cancel_token.cancelled() => Err(format!("{}: operation canceled", error_prefix)), - res = op => res.map_err(|err| format!("{}: {}", error_prefix, err)), - } -} - -#[cfg(feature = "ffi-dataplane")] -async fn run_io_with_cancel( - cancel_token: &CancellationToken, - close_token: &CancellationToken, - timeout_ms: u64, - error_prefix: &str, - op: F, -) -> Result -where - F: Future>, -{ - tokio::select! { - biased; - _ = cancel_token.cancelled() => Err(format!("{}: operation canceled", error_prefix)), - _ = close_token.cancelled() => Err(format!("{}: handle closed", error_prefix)), - res = tokio::time::timeout(timeout_duration(timeout_ms), op) => match res { - Ok(Ok(value)) => Ok(value), - Ok(Err(err)) => Err(format!("{}: {}", error_prefix, err)), - Err(_) => Err(format!("{} timed out", error_prefix)), - }, - } -} - -#[cfg(feature = "ffi-dataplane")] -fn leak_bytes(data: Vec) -> (*const std::ffi::c_uchar, u32) { - if data.is_empty() { - return (std::ptr::null(), 0); - } - let len = data.len() as u32; - let boxed = data.into_boxed_slice(); - (Box::into_raw(boxed) as *const std::ffi::c_uchar, len) -} - -#[cfg(feature = "ffi-dataplane")] -unsafe fn write_addr( - addr: SocketAddr, - out_ip: *mut *const std::ffi::c_char, - out_port: *mut std::ffi::c_ushort, -) -> Option<*mut std::ffi::c_char> { - let ip = into_ffi_ip_cstring(addr.ip())?; - unsafe { - *out_ip = ip as *const std::ffi::c_char; - *out_port = addr.port(); - } - Some(ip) -} - -#[cfg(feature = "ffi-dataplane")] -pub(crate) fn has_live_ops() -> bool { - !DATA_PLANE_OPS.is_empty() -} - -#[cfg(feature = "ffi-dataplane")] -pub(crate) fn cancel_ops_for_handle(handle: u64) { - let ops = DATA_PLANE_OPS - .iter() - .filter(|entry| entry.target_handle == Some(handle)) - .map(|entry| entry.value().clone()) - .collect::>(); - for op in ops { - cancel_pending_op( - &op, - "data plane async op canceled because handle was closed", - ); - } -} - -#[cfg(feature = "ffi-dataplane")] -pub(crate) fn remove_ops_by_instance_ids(ids: &[Uuid]) { - if ids.is_empty() { - return; - } - - let op_handles = DATA_PLANE_OPS - .iter() - .filter(|entry| entry.instance_id.is_some_and(|id| ids.contains(&id))) - .map(|entry| *entry.key()) - .collect::>(); - for handle in op_handles { - if let Some((_, op)) = DATA_PLANE_OPS.remove(&handle) { - op.cancel_token.cancel(); - if let Ok(mut state) = op.state.lock() { - *state = DataPlaneAsyncOpState::Consumed; - op.ready.notify_all(); - } - } - } -} - -#[cfg(feature = "ffi-dataplane")] -pub(crate) fn data_plane_async_op_status(handle: u64) -> std::ffi::c_int { - let Some(op) = DATA_PLANE_OPS.get(&handle).map(|op| op.clone()) else { - return DATA_PLANE_OP_INVALID; - }; - let Ok(state) = op.state.lock() else { - return DATA_PLANE_OP_FAILED; - }; - state_status(&state) -} - -#[cfg(feature = "ffi-dataplane")] -pub(crate) fn data_plane_async_op_wait(handle: u64, timeout_ms: u64) -> std::ffi::c_int { - let Some(op) = DATA_PLANE_OPS.get(&handle).map(|op| op.clone()) else { - return DATA_PLANE_OP_INVALID; - }; - let Ok(mut state) = op.state.lock() else { - return DATA_PLANE_OP_FAILED; - }; - if matches!(*state, DataPlaneAsyncOpState::Pending) && timeout_ms > 0 { - let timeout = Duration::from_millis(timeout_ms); - let Ok((next_state, _)) = op.ready.wait_timeout_while(state, timeout, |state| { - matches!(state, DataPlaneAsyncOpState::Pending) - }) else { - return DATA_PLANE_OP_FAILED; - }; - state = next_state; - } - state_status(&state) -} - -#[cfg(feature = "ffi-dataplane")] -pub(crate) fn data_plane_async_op_cancel(handle: u64) -> std::ffi::c_int { - let Some(op) = DATA_PLANE_OPS.get(&handle).map(|op| op.clone()) else { - return DATA_PLANE_OP_INVALID; - }; - cancel_pending_op(&op, "data plane async op canceled"); - 0 -} - -#[cfg(feature = "ffi-dataplane")] -pub(crate) fn data_plane_async_op_free(handle: u64) -> std::ffi::c_int { - let Some((_, op)) = DATA_PLANE_OPS.remove(&handle) else { - return DATA_PLANE_OP_INVALID; - }; - op.cancel_token.cancel(); - if let Ok(mut state) = op.state.lock() { - *state = DataPlaneAsyncOpState::Consumed; - op.ready.notify_all(); - } - 0 -} - -#[cfg(feature = "ffi-dataplane")] -pub(crate) fn data_plane_free_bytes(ptr: *const std::ffi::c_uchar, len: u32) { - if ptr.is_null() { - return; - } - let slice = std::ptr::slice_from_raw_parts_mut(ptr as *mut std::ffi::c_uchar, len as usize); - unsafe { - drop(Box::from_raw(slice)); - } -} - -#[cfg(feature = "ffi-dataplane")] -pub(crate) unsafe fn data_plane_tcp_connect_start( - inst_name: *const std::ffi::c_char, - dst_ip: *const std::ffi::c_char, - dst_port: std::ffi::c_ushort, - timeout_ms: u64, -) -> u64 { - let _data_plane_usage_guard = match enter_data_plane_operation() { - Some(guard) => guard, - None => return 0, - }; - let Some(inst_name) = (unsafe { cstr_to_string(inst_name, "inst_name") }) else { - return 0; - }; - let Some(dst_ip) = (unsafe { cstr_to_string(dst_ip, "dst_ip") }) else { - return 0; - }; - let Some(instance_id) = get_instance_id(&inst_name) else { - set_error_msg("instance not found"); - return 0; - }; - let Some(dst_addr) = parse_socket_addr(&dst_ip, dst_port) else { - return 0; - }; - - let (handle, op) = new_op(DataPlaneAsyncOpKind::TcpConnect, Some(instance_id), None); - let op_for_task = op.clone(); - spawn_instance_runtime_op( - op, - instance_id, - timeout_ms, - move |runtime_for_result, remaining| async move { - run_with_cancel( - &op_for_task.cancel_token, - "failed to connect tcp data plane", - INSTANCE_MANAGER.data_plane_tcp_connect(&instance_id, dst_addr, remaining), - ) - .await - .map(|stream| { - let local_addr = stream.local_addr(); - DataPlaneAsyncOpResult::TcpConnect { - instance_id, - runtime: runtime_for_result, - stream, - local_addr, - } - }) - }, - ); - handle -} - -#[cfg(feature = "ffi-dataplane")] -pub(crate) unsafe fn data_plane_tcp_connect_finish( - op_handle: u64, - out_local_ip: *mut *const std::ffi::c_char, - out_local_port: *mut std::ffi::c_ushort, -) -> u64 { - if out_local_ip.is_null() || out_local_port.is_null() { - set_error_msg("output pointer is null"); - return 0; - } - let _data_plane_usage_guard = match enter_data_plane_operation() { - Some(guard) => guard, - None => return 0, - }; - let Some(result) = take_completed_op(op_handle, DataPlaneAsyncOpKind::TcpConnect) else { - return 0; - }; - let DataPlaneAsyncOpResult::TcpConnect { - instance_id, - runtime, - stream, - local_addr, - } = result - else { - set_error_msg("data plane async op result type mismatch"); - return 0; - }; - let Some(_ip) = (unsafe { write_addr(local_addr, out_local_ip, out_local_port) }) else { - return 0; - }; - insert_tcp_stream_handle(instance_id, runtime, stream) -} - -#[cfg(feature = "ffi-dataplane")] -pub(crate) unsafe fn data_plane_tcp_bind_start( - inst_name: *const std::ffi::c_char, - local_port: std::ffi::c_ushort, - timeout_ms: u64, -) -> u64 { - let _data_plane_usage_guard = match enter_data_plane_operation() { - Some(guard) => guard, - None => return 0, - }; - let Some(inst_name) = (unsafe { cstr_to_string(inst_name, "inst_name") }) else { - return 0; - }; - let Some(instance_id) = get_instance_id(&inst_name) else { - set_error_msg("instance not found"); - return 0; - }; - - let (handle, op) = new_op(DataPlaneAsyncOpKind::TcpBind, Some(instance_id), None); - let op_for_task = op.clone(); - spawn_instance_runtime_op( - op, - instance_id, - timeout_ms, - move |runtime_for_result, remaining| async move { - run_with_cancel( - &op_for_task.cancel_token, - "failed to bind tcp data plane", - INSTANCE_MANAGER.data_plane_tcp_bind(&instance_id, local_port, remaining), - ) - .await - .map(|listener| { - let local_addr = listener.local_addr(); - DataPlaneAsyncOpResult::TcpBind { - instance_id, - runtime: runtime_for_result, - listener, - local_addr, - } - }) - }, - ); - handle -} - -#[cfg(feature = "ffi-dataplane")] -pub(crate) unsafe fn data_plane_tcp_bind_finish( - op_handle: u64, - out_local_ip: *mut *const std::ffi::c_char, - out_local_port: *mut std::ffi::c_ushort, -) -> u64 { - if out_local_ip.is_null() || out_local_port.is_null() { - set_error_msg("output pointer is null"); - return 0; - } - let _data_plane_usage_guard = match enter_data_plane_operation() { - Some(guard) => guard, - None => return 0, - }; - let Some(result) = take_completed_op(op_handle, DataPlaneAsyncOpKind::TcpBind) else { - return 0; - }; - let DataPlaneAsyncOpResult::TcpBind { - instance_id, - runtime, - listener, - local_addr, - } = result - else { - set_error_msg("data plane async op result type mismatch"); - return 0; - }; - let Some(_ip) = (unsafe { write_addr(local_addr, out_local_ip, out_local_port) }) else { - return 0; - }; - insert_tcp_listener_handle(instance_id, runtime, listener) -} - -#[cfg(feature = "ffi-dataplane")] -pub(crate) unsafe fn data_plane_tcp_accept_start(handle: u64, timeout_ms: u64) -> u64 { - let _data_plane_usage_guard = match enter_data_plane_operation() { - Some(guard) => guard, - None => return 0, - }; - let Some((listener, runtime, close_token, instance_id)) = get_tcp_listener(handle) else { - return 0; - }; - - let (op_handle, op) = new_op( - DataPlaneAsyncOpKind::TcpAccept, - Some(instance_id), - Some(handle), - ); - let runtime_for_result = runtime.clone(); - runtime.spawn(async move { - let result = async { - let mut listener = listener.lock().await; - let (stream, peer_addr) = run_io_with_cancel( - &op.cancel_token, - &close_token, - timeout_ms, - "tcp data plane accept", - listener.accept(), - ) - .await?; - let local_addr = stream.local_addr(); - Ok(DataPlaneAsyncOpResult::TcpAccept { - instance_id, - runtime: runtime_for_result, - stream, - local_addr, - peer_addr, - }) - } - .await; - complete_op(&op, result); - }); - op_handle -} - -#[cfg(feature = "ffi-dataplane")] -pub(crate) unsafe fn data_plane_tcp_accept_finish( - op_handle: u64, - out_local_ip: *mut *const std::ffi::c_char, - out_local_port: *mut std::ffi::c_ushort, - out_peer_ip: *mut *const std::ffi::c_char, - out_peer_port: *mut std::ffi::c_ushort, -) -> u64 { - if out_local_ip.is_null() - || out_local_port.is_null() - || out_peer_ip.is_null() - || out_peer_port.is_null() - { - set_error_msg("output pointer is null"); - return 0; - } - let _data_plane_usage_guard = match enter_data_plane_operation() { - Some(guard) => guard, - None => return 0, - }; - let Some(result) = take_completed_op(op_handle, DataPlaneAsyncOpKind::TcpAccept) else { - return 0; - }; - let DataPlaneAsyncOpResult::TcpAccept { - instance_id, - runtime, - stream, - local_addr, - peer_addr, - } = result - else { - set_error_msg("data plane async op result type mismatch"); - return 0; - }; - let Some(local_ip) = (unsafe { write_addr(local_addr, out_local_ip, out_local_port) }) else { - return 0; - }; - if (unsafe { write_addr(peer_addr, out_peer_ip, out_peer_port) }).is_none() { - free_string(local_ip); - return 0; - } - insert_tcp_stream_handle(instance_id, runtime, stream) -} - -#[cfg(feature = "ffi-dataplane")] -pub(crate) unsafe fn data_plane_tcp_read_start(handle: u64, max_len: u32, timeout_ms: u64) -> u64 { - let _data_plane_usage_guard = match enter_data_plane_operation() { - Some(guard) => guard, - None => return 0, - }; - if !validate_max_len(max_len) { - return 0; - } - let Some((halves, runtime, close_token, instance_id)) = get_tcp_stream_with_instance(handle) - else { - return 0; - }; - let (op_handle, op) = new_op( - DataPlaneAsyncOpKind::TcpRead, - Some(instance_id), - Some(handle), - ); - runtime.spawn(async move { - let result = read_tcp(halves, op.clone(), close_token, max_len, timeout_ms).await; - complete_op(&op, result); - }); - op_handle -} - -#[cfg(feature = "ffi-dataplane")] -async fn read_tcp( - halves: Arc, - op: Arc, - close_token: CancellationToken, - max_len: u32, - timeout_ms: u64, -) -> Result { - let mut buf = vec![0; max_len as usize]; - let mut rd = halves.read.lock().await; - let n = run_io_with_cancel( - &op.cancel_token, - &close_token, - timeout_ms, - "failed to read tcp data plane", - rd.read(&mut buf), - ) - .await?; - buf.truncate(n); - Ok(DataPlaneAsyncOpResult::TcpRead { data: buf }) -} - -#[cfg(feature = "ffi-dataplane")] -pub(crate) unsafe fn data_plane_tcp_read_finish( - op_handle: u64, - out_buf: *mut *const std::ffi::c_uchar, - out_len: *mut u32, -) -> std::ffi::c_int { - if out_buf.is_null() || out_len.is_null() { - set_error_msg("output pointer is null"); - return -1; - } - let Some(result) = take_completed_op(op_handle, DataPlaneAsyncOpKind::TcpRead) else { - return -1; - }; - let DataPlaneAsyncOpResult::TcpRead { data } = result else { - set_error_msg("data plane async op result type mismatch"); - return -1; - }; - let (ptr, len) = leak_bytes(data); - unsafe { - *out_buf = ptr; - *out_len = len; - } - len as std::ffi::c_int -} - -#[cfg(feature = "ffi-dataplane")] -pub(crate) unsafe fn data_plane_tcp_write_start( - handle: u64, - buf: *const std::ffi::c_uchar, - len: u32, - timeout_ms: u64, -) -> u64 { - let _data_plane_usage_guard = match enter_data_plane_operation() { - Some(guard) => guard, - None => return 0, - }; - if len > 0 && buf.is_null() { - set_error_msg("buf is null"); - return 0; - } - if !validate_write_len(len) { - return 0; - } - let Some((halves, runtime, close_token, instance_id)) = get_tcp_stream_with_instance(handle) - else { - return 0; - }; - let data = if len == 0 { - Vec::new() - } else { - unsafe { std::slice::from_raw_parts(buf, len as usize) }.to_vec() - }; - let (op_handle, op) = new_op( - DataPlaneAsyncOpKind::TcpWrite, - Some(instance_id), - Some(handle), - ); - runtime.spawn(async move { - let result = write_tcp(halves, op.clone(), close_token, data, timeout_ms).await; - complete_op(&op, result); - }); - op_handle -} - -#[cfg(feature = "ffi-dataplane")] -async fn write_tcp( - halves: Arc, - op: Arc, - close_token: CancellationToken, - data: Vec, - timeout_ms: u64, -) -> Result { - let written = data.len(); - let mut wr = halves.write.lock().await; - run_io_with_cancel( - &op.cancel_token, - &close_token, - timeout_ms, - "failed to write tcp data plane", - wr.write_all(&data), - ) - .await?; - Ok(DataPlaneAsyncOpResult::TcpWrite { written }) -} - -#[cfg(feature = "ffi-dataplane")] -pub(crate) fn data_plane_tcp_write_finish(op_handle: u64) -> std::ffi::c_int { - let Some(result) = take_completed_op(op_handle, DataPlaneAsyncOpKind::TcpWrite) else { - return -1; - }; - let DataPlaneAsyncOpResult::TcpWrite { written } = result else { - set_error_msg("data plane async op result type mismatch"); - return -1; - }; - usize_to_c_int(written, "tcp write byte count").unwrap_or(-1) -} - -#[cfg(feature = "ffi-dataplane")] -pub(crate) unsafe fn data_plane_udp_bind_start( - inst_name: *const std::ffi::c_char, - local_port: std::ffi::c_ushort, - timeout_ms: u64, -) -> u64 { - let _data_plane_usage_guard = match enter_data_plane_operation() { - Some(guard) => guard, - None => return 0, - }; - let Some(inst_name) = (unsafe { cstr_to_string(inst_name, "inst_name") }) else { - return 0; - }; - let Some(instance_id) = get_instance_id(&inst_name) else { - set_error_msg("instance not found"); - return 0; - }; - - let (handle, op) = new_op(DataPlaneAsyncOpKind::UdpBind, Some(instance_id), None); - let op_for_task = op.clone(); - spawn_instance_runtime_op( - op, - instance_id, - timeout_ms, - move |runtime_for_result, remaining| async move { - run_with_cancel( - &op_for_task.cancel_token, - "failed to bind udp data plane", - INSTANCE_MANAGER.data_plane_udp_bind(&instance_id, local_port, remaining), - ) - .await - .map(|socket| { - let local_addr = socket.local_addr(); - DataPlaneAsyncOpResult::UdpBind { - instance_id, - runtime: runtime_for_result, - socket, - local_addr, - } - }) - }, - ); - handle -} - -#[cfg(feature = "ffi-dataplane")] -pub(crate) unsafe fn data_plane_udp_bind_finish( - op_handle: u64, - out_local_ip: *mut *const std::ffi::c_char, - out_local_port: *mut std::ffi::c_ushort, -) -> u64 { - if out_local_ip.is_null() || out_local_port.is_null() { - set_error_msg("output pointer is null"); - return 0; - } - let _data_plane_usage_guard = match enter_data_plane_operation() { - Some(guard) => guard, - None => return 0, - }; - let Some(result) = take_completed_op(op_handle, DataPlaneAsyncOpKind::UdpBind) else { - return 0; - }; - let DataPlaneAsyncOpResult::UdpBind { - instance_id, - runtime, - socket, - local_addr, - } = result - else { - set_error_msg("data plane async op result type mismatch"); - return 0; - }; - let Some(_ip) = (unsafe { write_addr(local_addr, out_local_ip, out_local_port) }) else { - return 0; - }; - insert_udp_socket_handle(instance_id, runtime, socket) -} - -#[cfg(feature = "ffi-dataplane")] -pub(crate) unsafe fn data_plane_udp_send_to_start( - handle: u64, - dst_ip: *const std::ffi::c_char, - dst_port: std::ffi::c_ushort, - buf: *const std::ffi::c_uchar, - len: u32, - timeout_ms: u64, -) -> u64 { - let _data_plane_usage_guard = match enter_data_plane_operation() { - Some(guard) => guard, - None => return 0, - }; - if len > 0 && buf.is_null() { - set_error_msg("buf is null"); - return 0; - } - if !validate_write_len(len) { - return 0; - } - let Some(dst_ip) = (unsafe { cstr_to_string(dst_ip, "dst_ip") }) else { - return 0; - }; - let Some(dst_addr) = parse_socket_addr(&dst_ip, dst_port) else { - return 0; - }; - let Some((socket, runtime, close_token, instance_id)) = get_udp_socket_with_instance(handle) - else { - return 0; - }; - let data = if len == 0 { - Vec::new() - } else { - unsafe { std::slice::from_raw_parts(buf, len as usize) }.to_vec() - }; - - let (op_handle, op) = new_op( - DataPlaneAsyncOpKind::UdpSendTo, - Some(instance_id), - Some(handle), - ); - runtime.spawn(async move { - let result = run_io_with_cancel( - &op.cancel_token, - &close_token, - timeout_ms, - "failed to send udp data plane", - socket.send_to(&data, dst_addr), - ) - .await - .map(|sent| DataPlaneAsyncOpResult::UdpSendTo { sent }); - complete_op(&op, result); - }); - op_handle -} - -#[cfg(feature = "ffi-dataplane")] -pub(crate) fn data_plane_udp_send_to_finish(op_handle: u64) -> std::ffi::c_int { - let Some(result) = take_completed_op(op_handle, DataPlaneAsyncOpKind::UdpSendTo) else { - return -1; - }; - let DataPlaneAsyncOpResult::UdpSendTo { sent } = result else { - set_error_msg("data plane async op result type mismatch"); - return -1; - }; - usize_to_c_int(sent, "udp send byte count").unwrap_or(-1) -} - -#[cfg(feature = "ffi-dataplane")] -pub(crate) unsafe fn data_plane_udp_recv_from_start( - handle: u64, - max_len: u32, - timeout_ms: u64, -) -> u64 { - let _data_plane_usage_guard = match enter_data_plane_operation() { - Some(guard) => guard, - None => return 0, - }; - if !validate_max_len(max_len) { - return 0; - } - let Some((socket, runtime, close_token, instance_id)) = get_udp_socket_with_instance(handle) - else { - return 0; - }; - let (op_handle, op) = new_op( - DataPlaneAsyncOpKind::UdpRecvFrom, - Some(instance_id), - Some(handle), - ); - runtime.spawn(async move { - let mut buf = vec![0; max_len as usize]; - let result = run_io_with_cancel( - &op.cancel_token, - &close_token, - timeout_ms, - "udp data plane receive", - socket.recv_from(&mut buf), - ) - .await - .map(|(n, peer_addr)| { - buf.truncate(n); - DataPlaneAsyncOpResult::UdpRecvFrom { - data: buf, - peer_addr, - } - }); - complete_op(&op, result); - }); - op_handle -} - -#[cfg(feature = "ffi-dataplane")] -pub(crate) unsafe fn data_plane_udp_recv_from_finish( - op_handle: u64, - out_buf: *mut *const std::ffi::c_uchar, - out_len: *mut u32, - out_ip: *mut *const std::ffi::c_char, - out_port: *mut std::ffi::c_ushort, -) -> std::ffi::c_int { - if out_buf.is_null() || out_len.is_null() || out_ip.is_null() || out_port.is_null() { - set_error_msg("output pointer is null"); - return -1; - } - let Some(result) = take_completed_op(op_handle, DataPlaneAsyncOpKind::UdpRecvFrom) else { - return -1; - }; - let DataPlaneAsyncOpResult::UdpRecvFrom { data, peer_addr } = result else { - set_error_msg("data plane async op result type mismatch"); - return -1; - }; - let Some(_ip) = (unsafe { write_addr(peer_addr, out_ip, out_port) }) else { - return -1; - }; - let (ptr, len) = leak_bytes(data); - unsafe { - *out_buf = ptr; - *out_len = len; - } - len as std::ffi::c_int -} - -#[cfg(all(test, feature = "ffi-dataplane"))] -mod tests { - use super::*; - - #[test] - fn cancel_marks_pending_op_failed_and_consumable() { - let (handle, _op) = new_op(DataPlaneAsyncOpKind::TcpRead, None, None); - - assert_eq!(data_plane_async_op_status(handle), DATA_PLANE_OP_PENDING); - assert_eq!(data_plane_async_op_cancel(handle), 0); - assert_eq!(data_plane_async_op_wait(handle, 0), DATA_PLANE_OP_FAILED); - assert!(take_completed_op(handle, DataPlaneAsyncOpKind::TcpRead).is_none()); - assert_eq!(data_plane_async_op_status(handle), DATA_PLANE_OP_INVALID); - } - - #[test] - fn max_len_limit_rejects_oversized_async_reads() { - assert!(validate_max_len(MAX_ASYNC_READ_LEN)); - assert!(!validate_max_len(MAX_ASYNC_READ_LEN + 1)); - } - - #[test] - fn write_len_limit_rejects_values_that_c_int_cannot_return() { - assert!(validate_write_len(MAX_ASYNC_WRITE_LEN)); - assert!(!validate_write_len(MAX_ASYNC_WRITE_LEN + 1)); - } - - #[test] - fn finish_return_count_must_fit_c_int() { - assert_eq!(usize_to_c_int(123, "test byte count"), Some(123)); - assert!(usize_to_c_int(std::ffi::c_int::MAX as usize + 1, "test byte count").is_none()); - } - - #[test] - fn free_consumes_ready_op_before_finish_can_take_it() { - let (handle, op) = new_op(DataPlaneAsyncOpKind::TcpRead, None, None); - complete_op(&op, Ok(DataPlaneAsyncOpResult::TcpRead { data: vec![1] })); - - assert_eq!(data_plane_async_op_free(handle), 0); - assert!(take_completed_op(handle, DataPlaneAsyncOpKind::TcpRead).is_none()); - assert_eq!(data_plane_async_op_status(handle), DATA_PLANE_OP_INVALID); - } -} diff --git a/easytier-contrib/easytier-ffi/src/instance_api.rs b/easytier-contrib/easytier-ffi/src/instance_api.rs index 1ce290d6..285d0eaa 100644 --- a/easytier-contrib/easytier-ffi/src/instance_api.rs +++ b/easytier-contrib/easytier-ffi/src/instance_api.rs @@ -1,18 +1,11 @@ use std::ffi::{CString, c_char, c_int}; -use easytier::common::config::{ConfigFileControl, ConfigLoader as _, TomlConfigLoader}; +use easytier::common::config::{ConfigFileControl, TomlConfigLoader}; use crate::{ - config_server::{ - in_config_server_callback, remove_config_server_tracked_instance_ids, - wait_for_config_server_delivery, - }, - data_plane::remove_data_plane_handles_by_instance_ids, + config_server::{in_config_server_callback, wait_for_config_server_delivery}, error::set_error_msg, - state::{ - INSTANCE_MANAGER, INSTANCE_MUTATION_LOCK, INSTANCE_NAME_ID_MAP, instance_name_exists, - lock_remote_instance_mutation, - }, + state::{ffi_context, resolve_instance_id_by_name}, types::KeyValuePair, }; @@ -25,17 +18,19 @@ pub(crate) unsafe fn set_tun_fd(inst_name: *const c_char, fd: c_int) -> c_int { .to_string_lossy() .into_owned() }; - if !INSTANCE_NAME_ID_MAP.contains_key(&inst_name) { - return -1; - } + let inst_id = match resolve_instance_id_by_name(&inst_name) { + Ok(Some(instance_id)) => instance_id, + Ok(None) => { + set_error_msg("instance not found"); + return -1; + } + Err(error) => { + set_error_msg(&error.to_string()); + return -1; + } + }; - let inst_id = *INSTANCE_NAME_ID_MAP - .get(&inst_name) - .as_ref() - .unwrap() - .value(); - - match INSTANCE_MANAGER.set_tun_fd(&inst_id, fd) { + match ffi_context().manager.attach_tun_fd(inst_id, fd) { Ok(_) => 0, Err(_) => -1, } @@ -81,34 +76,16 @@ pub(crate) unsafe fn run_network_instance(cfg_str: *const std::ffi::c_char) -> s } }; - let inst_name = cfg.get_inst_name(); - wait_for_config_server_delivery(); - let _remote_mutation_guard = lock_remote_instance_mutation(); - let _mutation_guard = match INSTANCE_MUTATION_LOCK.lock() { - Ok(guard) => guard, - Err(err) => { - set_error_msg(&format!("failed to lock instance mutation: {}", err)); - return -1; - } - }; - - if instance_name_exists(&inst_name) { - set_error_msg("instance already exists"); + if let Err(e) = ffi_context().runtime.block_on( + ffi_context() + .process_management + .run_owned_network_instance(cfg, ConfigFileControl::STATIC_CONFIG), + ) { + set_error_msg(&format!("failed to start instance: {}", e)); return -1; } - let instance_id = - match INSTANCE_MANAGER.run_network_instance(cfg, false, ConfigFileControl::STATIC_CONFIG) { - Ok(id) => id, - Err(e) => { - set_error_msg(&format!("failed to start instance: {}", e)); - return -1; - } - }; - - INSTANCE_NAME_ID_MAP.insert(inst_name, instance_id); - 0 } @@ -152,50 +129,24 @@ pub(crate) unsafe fn retain_network_instance( } wait_for_config_server_delivery(); - let _remote_mutation_guard = lock_remote_instance_mutation(); - let _mutation_guard = match INSTANCE_MUTATION_LOCK.lock() { - Ok(guard) => guard, - Err(err) => { - set_error_msg(&format!("failed to lock instance mutation: {}", err)); + let retained_names = if length == 0 { + Vec::new() + } else { + let Some(inst_names) = (unsafe { parse_instance_names(inst_names, length) }) else { return -1; - } + }; + inst_names }; - if length == 0 { - let removed_ids = INSTANCE_MANAGER.list_network_instance_ids(); - if let Err(e) = INSTANCE_MANAGER.delete_network_instance(removed_ids.clone()) { - set_error_msg(&format!("failed to delete instances: {}", e)); - return -1; - } - remove_config_server_tracked_instance_ids(&removed_ids); - remove_data_plane_handles_by_instance_ids(&removed_ids); - INSTANCE_NAME_ID_MAP.clear(); - return 0; - } - - let Some(inst_names) = (unsafe { parse_instance_names(inst_names, length) }) else { - return -1; - }; - - let removed_ids = INSTANCE_MANAGER - .list_network_instance_ids() - .into_iter() - .filter(|id| { - INSTANCE_MANAGER - .get_instance_name(id) - .is_none_or(|name| !inst_names.contains(&name)) - }) - .collect::>(); - - if let Err(e) = INSTANCE_MANAGER.delete_network_instance(removed_ids.clone()) { - set_error_msg(&format!("failed to delete instances: {}", e)); + if let Err(error) = ffi_context().runtime.block_on( + ffi_context() + .process_management + .retain_owned_network_instances_by_name(retained_names), + ) { + set_error_msg(&format!("failed to retain instances: {error}")); return -1; } - remove_config_server_tracked_instance_ids(&removed_ids); - remove_data_plane_handles_by_instance_ids(&removed_ids); - INSTANCE_NAME_ID_MAP.retain(|k, _| inst_names.contains(k)); - 0 } @@ -211,15 +162,6 @@ pub(crate) unsafe fn delete_network_instance( } wait_for_config_server_delivery(); - let _remote_mutation_guard = lock_remote_instance_mutation(); - let _mutation_guard = match INSTANCE_MUTATION_LOCK.lock() { - Ok(guard) => guard, - Err(err) => { - set_error_msg(&format!("failed to lock instance mutation: {}", err)); - return -1; - } - }; - if length == 0 { return 0; } @@ -228,22 +170,15 @@ pub(crate) unsafe fn delete_network_instance( return -1; }; - let removed_ids = inst_names - .iter() - .filter_map(|name| INSTANCE_NAME_ID_MAP.get(name).map(|id| *id.value())) - .collect::>(); - - if let Err(e) = INSTANCE_MANAGER.delete_network_instance(removed_ids.clone()) { - set_error_msg(&format!("failed to delete instances: {}", e)); + if let Err(error) = ffi_context().runtime.block_on( + ffi_context() + .process_management + .delete_owned_network_instances_by_name(inst_names), + ) { + set_error_msg(&format!("failed to delete instances: {error}")); return -1; } - remove_config_server_tracked_instance_ids(&removed_ids); - remove_data_plane_handles_by_instance_ids(&removed_ids); - for name in inst_names { - INSTANCE_NAME_ID_MAP.remove(&name); - } - 0 } @@ -267,7 +202,7 @@ pub(crate) unsafe fn collect_network_infos( std::slice::from_raw_parts_mut(infos, max_length) }; - let collected_infos = match INSTANCE_MANAGER.collect_network_infos_sync() { + let collected_infos = match ffi_context().manager.collect_network_infos_sync() { Ok(infos) => infos, Err(e) => { set_error_msg(&format!("failed to collect network infos: {}", e)); @@ -280,7 +215,11 @@ pub(crate) unsafe fn collect_network_infos( if index >= max_length { break; } - let Some(key) = INSTANCE_MANAGER.get_instance_name(instance_id) else { + let Some(key) = ffi_context() + .manager + .instance(*instance_id) + .map(|instance| instance.instance_name().to_owned()) + else { continue; }; // convert value to json string @@ -320,13 +259,15 @@ pub(crate) unsafe fn list_instance(infos: *mut KeyValuePair, max_length: usize) } let infos = unsafe { std::slice::from_raw_parts_mut(infos, max_length) }; - let mut instances = INSTANCE_MANAGER - .list_network_instance_ids() + let mut instances = ffi_context() + .manager + .instance_ids() .into_iter() .filter_map(|id| { - INSTANCE_MANAGER - .get_instance_name(&id) - .map(|name| (name, id)) + ffi_context() + .manager + .instance(id) + .map(|instance| (instance.instance_name().to_owned(), id)) }) .collect::>(); instances.sort_by(|(left_name, left_id), (right_name, right_id)| { diff --git a/easytier-contrib/easytier-ffi/src/json_rpc.rs b/easytier-contrib/easytier-ffi/src/json_rpc.rs index 52c29043..e5300ec6 100644 --- a/easytier-contrib/easytier-ffi/src/json_rpc.rs +++ b/easytier-contrib/easytier-ffi/src/json_rpc.rs @@ -1,9 +1,12 @@ -use std::ffi::{CString, c_char, c_int}; +use std::{ + ffi::{CString, c_char, c_int}, + sync::Arc, +}; use crate::{ config_server::in_config_server_callback, error::set_error_msg, - state::{ASYNC_RUNTIME, INSTANCE_MANAGER}, + state::ffi_context, strings::{c_str_to_string, optional_c_str_to_string}, }; @@ -65,19 +68,23 @@ pub(crate) unsafe fn call_json_rpc( } }; - let response = match ASYNC_RUNTIME.block_on(easytier::rpc_service::call_json_rpc( - &INSTANCE_MANAGER, - &service_name, - &method_name, - domain_name.as_deref(), - payload, - )) { - Ok(value) => value, - Err(err) => { - set_error_msg(&format!("RPC Error: {}", err)); - return -1; - } - }; + let response = + match ffi_context() + .runtime + .block_on(easytier_core::management::call_management_json_rpc( + &ffi_context().manager, + Arc::new(easytier::rpc_service::logger::NativeLoggerControl), + &service_name, + &method_name, + domain_name.as_deref(), + payload, + )) { + Ok(value) => value, + Err(err) => { + set_error_msg(&format!("RPC Error: {}", err)); + return -1; + } + }; let response_json = match serde_json::to_string(&response) { Ok(value) => value, Err(err) => { diff --git a/easytier-contrib/easytier-ffi/src/lib.rs b/easytier-contrib/easytier-ffi/src/lib.rs index b963793b..b185b602 100644 --- a/easytier-contrib/easytier-ffi/src/lib.rs +++ b/easytier-contrib/easytier-ffi/src/lib.rs @@ -20,19 +20,12 @@ //! - `is_config_server_client_connected`: report whether the client is connected. //! //! Data plane APIs, enabled by the `ffi-dataplane` feature: -//! - `data_plane_tcp_connect`: open an outbound TCP data-plane stream. -//! - `data_plane_tcp_bind`: bind a TCP data-plane listener. -//! - `data_plane_tcp_accept`: accept a TCP data-plane connection. -//! - `data_plane_tcp_read`: read from a TCP data-plane stream. -//! - `data_plane_tcp_write`: write to a TCP data-plane stream. -//! - `data_plane_tcp_close`: close a TCP data-plane stream. -//! - `data_plane_tcp_listener_close`: close a TCP data-plane listener. -//! - `data_plane_udp_bind`: bind a UDP data-plane socket. -//! - `data_plane_udp_send_to`: send one UDP data-plane datagram. -//! - `data_plane_udp_recv_from`: receive one UDP data-plane datagram. -//! - `data_plane_udp_close`: close a UDP data-plane socket. -//! - `data_plane_*_start` / `data_plane_*_finish`: asynchronous data-plane operations. -//! - `data_plane_async_op_*`: poll, wait, cancel, and free asynchronous operations. +//! - `data_plane_session_open` / `data_plane_session_close`: own one instance session. +//! - `data_plane_*_submit`: submit non-blocking TCP and UDP operations. +//! - `data_plane_completion_wait` / `data_plane_completion_drain`: await completions. +//! - `data_plane_*_result_take`: consume typed operation results. +//! - `data_plane_operation_cancel` / `data_plane_operation_free`: control operations. +//! - `data_plane_resource_close`: close streams, listeners, and UDP sockets. //! //! Shared FFI helper APIs: //! - `get_error_msg`: copy the last FFI or config-server callback error message. @@ -40,8 +33,6 @@ mod config_server; mod data_plane; -#[cfg(feature = "ffi-dataplane")] -mod data_plane_async; mod error; mod instance_api; mod json_rpc; @@ -53,11 +44,11 @@ mod types; mod tests; pub use config_server::{in_config_server_callback, validate_config_server_client_options}; -pub use types::{ConfigServerEventCallback, KeyValuePair}; +pub use types::{ + ConfigServerEventCallback, DataPlaneCompletion, DataPlaneSocketAddr, KeyValuePair, +}; use std::ffi::{c_char, c_int, c_void}; -#[cfg(feature = "ffi-dataplane")] -use std::ffi::{c_uchar, c_ushort}; // ===== Network Management API ===== @@ -254,7 +245,7 @@ pub unsafe extern "C" fn call_json_rpc( /// Start the managed config-server client. /// /// The client reuses EasyTier's web-client path and applies remote config -/// changes through the shared `NetworkInstanceManager`. Successful remote run +/// changes through the shared `NativeInstanceManager`. Successful remote run /// and delete operations are delivered to `callback` as JSON event strings, one /// callback per affected instance. The event string is valid only for the /// duration of the callback; callers must copy it if they need to keep it. @@ -319,634 +310,26 @@ pub extern "C" fn is_config_server_client_connected() -> c_int { // ===== Data Plane API ===== -/// Open an outbound TCP stream through an EasyTier instance data plane. -/// -/// On success, writes the local address selected for the connection into -/// `out_local_ip` and `out_local_port`. The returned IP string is allocated by -/// this library and must be released with `free_string`. -/// -/// The data plane is mutually exclusive with the config-server client. This -/// function returns `0` if the config-server client is active or stopping. -/// -/// # Safety -/// `inst_name`, `dst_ip`, `out_local_ip`, and `out_local_port` must be non-null. -/// String pointers must point to null-terminated UTF-8 strings. -/// -/// # Return -/// Returns a non-zero TCP stream handle on success, or `0` on failure. On -/// failure, call `get_error_msg` on the same thread to retrieve details. #[cfg(feature = "ffi-dataplane")] -#[cfg_attr(feature = "c-abi", unsafe(no_mangle))] -pub unsafe extern "C" fn data_plane_tcp_connect( - inst_name: *const c_char, - dst_ip: *const c_char, - dst_port: c_ushort, - timeout_ms: u64, - out_local_ip: *mut *const c_char, - out_local_port: *mut c_ushort, -) -> u64 { - unsafe { - data_plane::data_plane_tcp_connect( - inst_name, - dst_ip, - dst_port, - timeout_ms, - out_local_ip, - out_local_port, - ) - } -} - -/// Bind a TCP listener through an EasyTier instance data plane. -/// -/// On success, writes the bound local address into `out_local_ip` and -/// `out_local_port`. The returned IP string is allocated by this library and -/// must be released with `free_string`. -/// -/// # Safety -/// `inst_name`, `out_local_ip`, and `out_local_port` must be non-null. -/// `inst_name` must point to a null-terminated UTF-8 string. -/// -/// # Return -/// Returns a non-zero TCP listener handle on success, or `0` on failure. On -/// failure, call `get_error_msg` on the same thread to retrieve details. -#[cfg(feature = "ffi-dataplane")] -#[cfg_attr(feature = "c-abi", unsafe(no_mangle))] -pub unsafe extern "C" fn data_plane_tcp_bind( - inst_name: *const c_char, - local_port: c_ushort, - timeout_ms: u64, - out_local_ip: *mut *const c_char, - out_local_port: *mut c_ushort, -) -> u64 { - unsafe { - data_plane::data_plane_tcp_bind( - inst_name, - local_port, - timeout_ms, - out_local_ip, - out_local_port, - ) - } -} - -/// Accept one connection from a TCP data-plane listener. -/// -/// On success, writes both local and peer socket addresses to the output -/// pointers. Returned IP strings are allocated by this library and must be -/// released with `free_string`. -/// -/// # Safety -/// All output pointers must be non-null and writable. `handle` must be a valid -/// TCP listener handle returned by `data_plane_tcp_bind`. -/// -/// # Return -/// Returns a non-zero TCP stream handle on success, or `0` on failure. On -/// failure, call `get_error_msg` on the same thread to retrieve details. -#[cfg(feature = "ffi-dataplane")] -#[cfg_attr(feature = "c-abi", unsafe(no_mangle))] -pub unsafe extern "C" fn data_plane_tcp_accept( - handle: u64, - timeout_ms: u64, - out_local_ip: *mut *const c_char, - out_local_port: *mut c_ushort, - out_peer_ip: *mut *const c_char, - out_peer_port: *mut c_ushort, -) -> u64 { - unsafe { - data_plane::data_plane_tcp_accept( - handle, - timeout_ms, - out_local_ip, - out_local_port, - out_peer_ip, - out_peer_port, - ) - } -} - -/// Read bytes from a TCP data-plane stream. -/// -/// # Safety -/// `handle` must be a valid TCP stream handle returned by -/// `data_plane_tcp_connect` or `data_plane_tcp_accept`. `buf` must be non-null -/// and writable for `len` bytes. -/// -/// # Return -/// Returns the number of bytes read, or `-1` on failure. On failure, call -/// `get_error_msg` on the same thread to retrieve details. -#[cfg(feature = "ffi-dataplane")] -#[cfg_attr(feature = "c-abi", unsafe(no_mangle))] -pub unsafe extern "C" fn data_plane_tcp_read( - handle: u64, - buf: *mut c_uchar, - len: u32, - timeout_ms: u64, -) -> c_int { - unsafe { data_plane::data_plane_tcp_read(handle, buf, len, timeout_ms) } -} - -/// Write bytes to a TCP data-plane stream. -/// -/// This function attempts to write exactly `len` bytes before returning -/// success. -/// -/// # Safety -/// `handle` must be a valid TCP stream handle returned by -/// `data_plane_tcp_connect` or `data_plane_tcp_accept`. `buf` must be non-null -/// and readable for `len` bytes. -/// -/// # Return -/// Returns `len` on success, or `-1` on failure. On failure, call -/// `get_error_msg` on the same thread to retrieve details. -#[cfg(feature = "ffi-dataplane")] -#[cfg_attr(feature = "c-abi", unsafe(no_mangle))] -pub unsafe extern "C" fn data_plane_tcp_write( - handle: u64, - buf: *const c_uchar, - len: u32, - timeout_ms: u64, -) -> c_int { - unsafe { data_plane::data_plane_tcp_write(handle, buf, len, timeout_ms) } -} - -/// Close a TCP data-plane stream handle. -/// -/// # Return -/// Returns `0` on success, or `-1` if the handle is missing, is not a TCP stream -/// handle, or data-plane calls are currently rejected. -#[cfg(feature = "ffi-dataplane")] -#[cfg_attr(feature = "c-abi", unsafe(no_mangle))] -pub extern "C" fn data_plane_tcp_close(handle: u64) -> c_int { - data_plane::data_plane_tcp_close(handle) -} - -/// Close a TCP data-plane listener handle. -/// -/// # Return -/// Returns `0` on success, or `-1` if the handle is missing, is not a TCP -/// listener handle, or data-plane calls are currently rejected. -#[cfg(feature = "ffi-dataplane")] -#[cfg_attr(feature = "c-abi", unsafe(no_mangle))] -pub extern "C" fn data_plane_tcp_listener_close(handle: u64) -> c_int { - data_plane::data_plane_tcp_listener_close(handle) -} - -/// Bind a UDP socket through an EasyTier instance data plane. -/// -/// On success, writes the bound local address into `out_local_ip` and -/// `out_local_port`. The returned IP string is allocated by this library and -/// must be released with `free_string`. -/// -/// # Safety -/// `inst_name`, `out_local_ip`, and `out_local_port` must be non-null. -/// `inst_name` must point to a null-terminated UTF-8 string. -/// -/// # Return -/// Returns a non-zero UDP socket handle on success, or `0` on failure. On -/// failure, call `get_error_msg` on the same thread to retrieve details. -#[cfg(feature = "ffi-dataplane")] -#[cfg_attr(feature = "c-abi", unsafe(no_mangle))] -pub unsafe extern "C" fn data_plane_udp_bind( - inst_name: *const c_char, - local_port: c_ushort, - timeout_ms: u64, - out_local_ip: *mut *const c_char, - out_local_port: *mut c_ushort, -) -> u64 { - unsafe { - data_plane::data_plane_udp_bind( - inst_name, - local_port, - timeout_ms, - out_local_ip, - out_local_port, - ) - } -} - -/// Send one UDP datagram through a data-plane socket. -/// -/// # Safety -/// `handle` must be a valid UDP socket handle returned by -/// `data_plane_udp_bind`. `dst_ip` must be non-null and point to a -/// null-terminated UTF-8 string. `buf` must be non-null and readable for `len` -/// bytes. -/// -/// # Return -/// Returns the number of bytes sent, or `-1` on failure. On failure, call -/// `get_error_msg` on the same thread to retrieve details. -#[cfg(feature = "ffi-dataplane")] -#[cfg_attr(feature = "c-abi", unsafe(no_mangle))] -pub unsafe extern "C" fn data_plane_udp_send_to( - handle: u64, - dst_ip: *const c_char, - dst_port: c_ushort, - buf: *const c_uchar, - len: u32, - timeout_ms: u64, -) -> c_int { - unsafe { data_plane::data_plane_udp_send_to(handle, dst_ip, dst_port, buf, len, timeout_ms) } -} - -/// Receive one UDP datagram from a data-plane socket. -/// -/// On success, writes the peer address into `out_ip` and `out_port`. The -/// returned IP string is allocated by this library and must be released with -/// `free_string`. -/// -/// # Safety -/// `handle` must be a valid UDP socket handle returned by -/// `data_plane_udp_bind`. `buf`, `out_ip`, and `out_port` must be non-null. -/// `buf` must be writable for `len` bytes. -/// -/// # Return -/// Returns the number of bytes received, or `-1` on failure. On failure, call -/// `get_error_msg` on the same thread to retrieve details. -#[cfg(feature = "ffi-dataplane")] -#[cfg_attr(feature = "c-abi", unsafe(no_mangle))] -pub unsafe extern "C" fn data_plane_udp_recv_from( - handle: u64, - buf: *mut c_uchar, - len: u32, - out_ip: *mut *const c_char, - out_port: *mut c_ushort, - timeout_ms: u64, -) -> c_int { - unsafe { data_plane::data_plane_udp_recv_from(handle, buf, len, out_ip, out_port, timeout_ms) } -} - -/// Close a UDP data-plane socket handle. -/// -/// # Return -/// Returns `0` on success, or `-1` if the handle is missing, is not a UDP -/// socket handle, or data-plane calls are currently rejected. -#[cfg(feature = "ffi-dataplane")] -#[cfg_attr(feature = "c-abi", unsafe(no_mangle))] -pub extern "C" fn data_plane_udp_close(handle: u64) -> c_int { - data_plane::data_plane_udp_close(handle) -} - -// ===== Async Data Plane API ===== - -#[cfg(feature = "ffi-dataplane")] -#[cfg_attr(feature = "c-abi", unsafe(no_mangle))] -pub extern "C" fn data_plane_async_op_status(handle: u64) -> c_int { - data_plane_async::data_plane_async_op_status(handle) -} - -#[cfg(feature = "ffi-dataplane")] -#[cfg_attr(feature = "c-abi", unsafe(no_mangle))] -pub extern "C" fn data_plane_async_op_wait(handle: u64, timeout_ms: u64) -> c_int { - data_plane_async::data_plane_async_op_wait(handle, timeout_ms) -} - -#[cfg(feature = "ffi-dataplane")] -#[cfg_attr(feature = "c-abi", unsafe(no_mangle))] -pub extern "C" fn data_plane_async_op_cancel(handle: u64) -> c_int { - data_plane_async::data_plane_async_op_cancel(handle) -} - -#[cfg(feature = "ffi-dataplane")] -#[cfg_attr(feature = "c-abi", unsafe(no_mangle))] -pub extern "C" fn data_plane_async_op_free(handle: u64) -> c_int { - data_plane_async::data_plane_async_op_free(handle) -} - -#[cfg(feature = "ffi-dataplane")] -#[cfg_attr(feature = "c-abi", unsafe(no_mangle))] -pub extern "C" fn data_plane_free_bytes(ptr: *const c_uchar, len: u32) { - data_plane_async::data_plane_free_bytes(ptr, len) -} - -/// Start an asynchronous TCP data-plane connection. -/// -/// # Safety -/// `inst_name` and `dst_ip` must be non-null pointers to null-terminated UTF-8 -/// strings. The strings only need to remain valid for the duration of this -/// call. -/// -/// # Return -/// Returns a non-zero async operation handle on success, or `0` on failure. -#[cfg(feature = "ffi-dataplane")] -#[cfg_attr(feature = "c-abi", unsafe(no_mangle))] -pub unsafe extern "C" fn data_plane_tcp_connect_start( - inst_name: *const c_char, - dst_ip: *const c_char, - dst_port: c_ushort, - timeout_ms: u64, -) -> u64 { - unsafe { - data_plane_async::data_plane_tcp_connect_start(inst_name, dst_ip, dst_port, timeout_ms) - } -} - -/// Finish an asynchronous TCP data-plane connection. -/// -/// On success, writes the stream local address into `out_local_ip` and -/// `out_local_port`. The returned IP string is allocated by this library and -/// must be released with `free_string`. -/// -/// # Safety -/// `out_local_ip` and `out_local_port` must be non-null pointers to writable -/// storage. -/// -/// # Return -/// Returns a non-zero TCP stream handle on success, or `0` on failure. -#[cfg(feature = "ffi-dataplane")] -#[cfg_attr(feature = "c-abi", unsafe(no_mangle))] -pub unsafe extern "C" fn data_plane_tcp_connect_finish( - op_handle: u64, - out_local_ip: *mut *const c_char, - out_local_port: *mut c_ushort, -) -> u64 { - unsafe { - data_plane_async::data_plane_tcp_connect_finish(op_handle, out_local_ip, out_local_port) - } -} - -/// Start an asynchronous TCP data-plane bind. -/// -/// # Safety -/// `inst_name` must be a non-null pointer to a null-terminated UTF-8 string. -/// The string only needs to remain valid for the duration of this call. -/// -/// # Return -/// Returns a non-zero async operation handle on success, or `0` on failure. -#[cfg(feature = "ffi-dataplane")] -#[cfg_attr(feature = "c-abi", unsafe(no_mangle))] -pub unsafe extern "C" fn data_plane_tcp_bind_start( - inst_name: *const c_char, - local_port: c_ushort, - timeout_ms: u64, -) -> u64 { - unsafe { data_plane_async::data_plane_tcp_bind_start(inst_name, local_port, timeout_ms) } -} - -/// Finish an asynchronous TCP data-plane bind. -/// -/// On success, writes the listener local address into `out_local_ip` and -/// `out_local_port`. The returned IP string is allocated by this library and -/// must be released with `free_string`. -/// -/// # Safety -/// `out_local_ip` and `out_local_port` must be non-null pointers to writable -/// storage. -/// -/// # Return -/// Returns a non-zero TCP listener handle on success, or `0` on failure. -#[cfg(feature = "ffi-dataplane")] -#[cfg_attr(feature = "c-abi", unsafe(no_mangle))] -pub unsafe extern "C" fn data_plane_tcp_bind_finish( - op_handle: u64, - out_local_ip: *mut *const c_char, - out_local_port: *mut c_ushort, -) -> u64 { - unsafe { data_plane_async::data_plane_tcp_bind_finish(op_handle, out_local_ip, out_local_port) } -} - -/// Start an asynchronous TCP data-plane accept on a listener handle. -/// -/// # Safety -/// `handle` must be a valid TCP listener handle returned by -/// `data_plane_tcp_bind` or `data_plane_tcp_bind_finish`. -/// -/// # Return -/// Returns a non-zero async operation handle on success, or `0` on failure. -#[cfg(feature = "ffi-dataplane")] -#[cfg_attr(feature = "c-abi", unsafe(no_mangle))] -pub unsafe extern "C" fn data_plane_tcp_accept_start(handle: u64, timeout_ms: u64) -> u64 { - unsafe { data_plane_async::data_plane_tcp_accept_start(handle, timeout_ms) } -} - -/// Finish an asynchronous TCP data-plane accept. -/// -/// On success, writes the accepted stream local address into `out_local_ip` and -/// `out_local_port`, and the peer address into `out_peer_ip` and -/// `out_peer_port`. Returned IP strings are allocated by this library and must -/// be released with `free_string`. -/// -/// # Safety -/// `out_local_ip`, `out_local_port`, `out_peer_ip`, and `out_peer_port` must be -/// non-null pointers to writable storage. -/// -/// # Return -/// Returns a non-zero TCP stream handle on success, or `0` on failure. -#[cfg(feature = "ffi-dataplane")] -#[cfg_attr(feature = "c-abi", unsafe(no_mangle))] -pub unsafe extern "C" fn data_plane_tcp_accept_finish( - op_handle: u64, - out_local_ip: *mut *const c_char, - out_local_port: *mut c_ushort, - out_peer_ip: *mut *const c_char, - out_peer_port: *mut c_ushort, -) -> u64 { - unsafe { - data_plane_async::data_plane_tcp_accept_finish( - op_handle, - out_local_ip, - out_local_port, - out_peer_ip, - out_peer_port, - ) - } -} - -/// Start an asynchronous TCP data-plane read. -/// -/// # Safety -/// `handle` must be a valid TCP stream handle returned by -/// `data_plane_tcp_connect_finish` or `data_plane_tcp_accept_finish`. -/// -/// # Return -/// Returns a non-zero async operation handle on success, or `0` on failure. -#[cfg(feature = "ffi-dataplane")] -#[cfg_attr(feature = "c-abi", unsafe(no_mangle))] -pub unsafe extern "C" fn data_plane_tcp_read_start( - handle: u64, - max_len: u32, - timeout_ms: u64, -) -> u64 { - unsafe { data_plane_async::data_plane_tcp_read_start(handle, max_len, timeout_ms) } -} - -/// Finish an asynchronous TCP data-plane read. -/// -/// On success, writes the received buffer pointer and length into `out_buf` and -/// `out_len`. The returned buffer is allocated by this library and must be -/// released with `data_plane_free_bytes`. -/// -/// # Safety -/// `out_buf` and `out_len` must be non-null pointers to writable storage. -/// -/// # Return -/// Returns the number of bytes read, or `-1` on failure. -#[cfg(feature = "ffi-dataplane")] -#[cfg_attr(feature = "c-abi", unsafe(no_mangle))] -pub unsafe extern "C" fn data_plane_tcp_read_finish( - op_handle: u64, - out_buf: *mut *const c_uchar, - out_len: *mut u32, -) -> c_int { - unsafe { data_plane_async::data_plane_tcp_read_finish(op_handle, out_buf, out_len) } -} - -/// Start an asynchronous TCP data-plane write. -/// -/// The input bytes are copied before this function returns. -/// -/// # Safety -/// `handle` must be a valid TCP stream handle returned by -/// `data_plane_tcp_connect_finish` or `data_plane_tcp_accept_finish`. If `len` -/// is non-zero, `buf` must be non-null and readable for `len` bytes. -/// -/// # Return -/// Returns a non-zero async operation handle on success, or `0` on failure. -#[cfg(feature = "ffi-dataplane")] -#[cfg_attr(feature = "c-abi", unsafe(no_mangle))] -pub unsafe extern "C" fn data_plane_tcp_write_start( - handle: u64, - buf: *const c_uchar, - len: u32, - timeout_ms: u64, -) -> u64 { - unsafe { data_plane_async::data_plane_tcp_write_start(handle, buf, len, timeout_ms) } -} - -#[cfg(feature = "ffi-dataplane")] -#[cfg_attr(feature = "c-abi", unsafe(no_mangle))] -pub extern "C" fn data_plane_tcp_write_finish(op_handle: u64) -> c_int { - data_plane_async::data_plane_tcp_write_finish(op_handle) -} - -/// Start an asynchronous UDP data-plane bind. -/// -/// # Safety -/// `inst_name` must be a non-null pointer to a null-terminated UTF-8 string. -/// The string only needs to remain valid for the duration of this call. -/// -/// # Return -/// Returns a non-zero async operation handle on success, or `0` on failure. -#[cfg(feature = "ffi-dataplane")] -#[cfg_attr(feature = "c-abi", unsafe(no_mangle))] -pub unsafe extern "C" fn data_plane_udp_bind_start( - inst_name: *const c_char, - local_port: c_ushort, - timeout_ms: u64, -) -> u64 { - unsafe { data_plane_async::data_plane_udp_bind_start(inst_name, local_port, timeout_ms) } -} - -/// Finish an asynchronous UDP data-plane bind. -/// -/// On success, writes the socket local address into `out_local_ip` and -/// `out_local_port`. The returned IP string is allocated by this library and -/// must be released with `free_string`. -/// -/// # Safety -/// `out_local_ip` and `out_local_port` must be non-null pointers to writable -/// storage. -/// -/// # Return -/// Returns a non-zero UDP socket handle on success, or `0` on failure. -#[cfg(feature = "ffi-dataplane")] -#[cfg_attr(feature = "c-abi", unsafe(no_mangle))] -pub unsafe extern "C" fn data_plane_udp_bind_finish( - op_handle: u64, - out_local_ip: *mut *const c_char, - out_local_port: *mut c_ushort, -) -> u64 { - unsafe { data_plane_async::data_plane_udp_bind_finish(op_handle, out_local_ip, out_local_port) } -} - -/// Start an asynchronous UDP data-plane send. -/// -/// The input bytes are copied before this function returns. -/// -/// # Safety -/// `handle` must be a valid UDP socket handle returned by -/// `data_plane_udp_bind_finish`. `dst_ip` must be a non-null pointer to a -/// null-terminated UTF-8 string. If `len` is non-zero, `buf` must be non-null -/// and readable for `len` bytes. -/// -/// # Return -/// Returns a non-zero async operation handle on success, or `0` on failure. -#[cfg(feature = "ffi-dataplane")] -#[cfg_attr(feature = "c-abi", unsafe(no_mangle))] -pub unsafe extern "C" fn data_plane_udp_send_to_start( - handle: u64, - dst_ip: *const c_char, - dst_port: c_ushort, - buf: *const c_uchar, - len: u32, - timeout_ms: u64, -) -> u64 { - unsafe { - data_plane_async::data_plane_udp_send_to_start( - handle, dst_ip, dst_port, buf, len, timeout_ms, - ) - } -} - -#[cfg(feature = "ffi-dataplane")] -#[cfg_attr(feature = "c-abi", unsafe(no_mangle))] -pub extern "C" fn data_plane_udp_send_to_finish(op_handle: u64) -> c_int { - data_plane_async::data_plane_udp_send_to_finish(op_handle) -} - -/// Start an asynchronous UDP data-plane receive. -/// -/// # Safety -/// `handle` must be a valid UDP socket handle returned by -/// `data_plane_udp_bind_finish`. -/// -/// # Return -/// Returns a non-zero async operation handle on success, or `0` on failure. -#[cfg(feature = "ffi-dataplane")] -#[cfg_attr(feature = "c-abi", unsafe(no_mangle))] -pub unsafe extern "C" fn data_plane_udp_recv_from_start( - handle: u64, - max_len: u32, - timeout_ms: u64, -) -> u64 { - unsafe { data_plane_async::data_plane_udp_recv_from_start(handle, max_len, timeout_ms) } -} - -/// Finish an asynchronous UDP data-plane receive. -/// -/// On success, writes the received buffer into `out_buf` and `out_len`, and -/// the peer address into `out_ip` and `out_port`. The returned buffer is -/// allocated by this library and must be released with `data_plane_free_bytes`; -/// the returned IP string must be released with `free_string`. -/// -/// # Safety -/// `out_buf`, `out_len`, `out_ip`, and `out_port` must be non-null pointers to -/// writable storage. -/// -/// # Return -/// Returns the number of bytes received, or `-1` on failure. -#[cfg(feature = "ffi-dataplane")] -#[cfg_attr(feature = "c-abi", unsafe(no_mangle))] -pub unsafe extern "C" fn data_plane_udp_recv_from_finish( - op_handle: u64, - out_buf: *mut *const c_uchar, - out_len: *mut u32, - out_ip: *mut *const c_char, - out_port: *mut c_ushort, -) -> c_int { - unsafe { - data_plane_async::data_plane_udp_recv_from_finish( - op_handle, out_buf, out_len, out_ip, out_port, - ) - } -} +pub use data_plane::{ + data_plane_completion_drain, data_plane_completion_wait, data_plane_operation_cancel, + data_plane_operation_free, data_plane_resource_close, data_plane_result_size, + data_plane_session_close, data_plane_session_open, data_plane_tcp_accept_result_take, + data_plane_tcp_accept_submit, data_plane_tcp_bind_result_take, data_plane_tcp_bind_submit, + data_plane_tcp_connect_result_take, data_plane_tcp_connect_submit, + data_plane_tcp_read_result_take, data_plane_tcp_read_submit, data_plane_tcp_write_result_take, + data_plane_tcp_write_submit, data_plane_udp_bind_result_take, data_plane_udp_bind_submit, + data_plane_udp_receive_result_take, data_plane_udp_receive_submit, + data_plane_udp_send_result_take, data_plane_udp_send_submit, +}; // ===== Shared FFI Helper API ===== /// Return the last FFI error message. /// -/// Synchronous API failures are stored in a thread-local buffer, so call this -/// on the same thread that received `-1` or `0` from another API. Config-server +/// API failures are stored in a thread-local buffer, so call this on the same +/// thread that received a negative status or another documented failure +/// sentinel. Config-server /// callback delivery failures may happen on a runtime thread; those are stored /// globally and are included here so direct FFI callers can still retrieve the /// last callback error. If there is no error message, this writes a null pointer diff --git a/easytier-contrib/easytier-ffi/src/state.rs b/easytier-contrib/easytier-ffi/src/state.rs index 706e8b45..749acea7 100644 --- a/easytier-contrib/easytier-ffi/src/state.rs +++ b/easytier-contrib/easytier-ffi/src/state.rs @@ -1,54 +1,66 @@ -use std::sync::{Arc, Mutex}; +use std::sync::Arc; -use dashmap::DashMap; -use easytier::instance_manager::NetworkInstanceManager; +use easytier::instance::factory::{ + NativeInstanceManager, NativeProcessManagement, native_instance_manager_with_runtime, + native_process_management, +}; use tokio::runtime::{Builder, Runtime}; -use uuid::Uuid; -pub(crate) static INSTANCE_NAME_ID_MAP: once_cell::sync::Lazy> = - once_cell::sync::Lazy::new(DashMap::new); -pub(crate) static INSTANCE_MANAGER: once_cell::sync::Lazy> = - once_cell::sync::Lazy::new(|| Arc::new(NetworkInstanceManager::new())); -pub(crate) static ASYNC_RUNTIME: once_cell::sync::Lazy = - once_cell::sync::Lazy::new(|| { - Builder::new_multi_thread() +struct FfiOwnedInstanceHooks; + +#[async_trait::async_trait] +impl easytier_core::management::InstanceMutationHooks for FfiOwnedInstanceHooks { + async fn post_remove_network_instances( + &self, + instance_ids: &[uuid::Uuid], + ) -> Result<(), String> { + crate::config_server::remove_config_server_tracked_instance_ids(instance_ids); + crate::data_plane::remove_data_plane_sessions_by_instance_ids(instance_ids); + Ok(()) + } +} + +pub(crate) struct FfiContext { + pub(crate) runtime: Runtime, + pub(crate) manager: Arc, + pub(crate) process_management: NativeProcessManagement, +} + +impl FfiContext { + fn new() -> Self { + let runtime = Builder::new_multi_thread() .enable_all() .build() - .expect("tokio runtime for easytier-ffi") - }); -pub(crate) static INSTANCE_MUTATION_LOCK: once_cell::sync::Lazy> = - once_cell::sync::Lazy::new(|| Mutex::new(())); - -pub(crate) fn remove_instance_name_ids(ids: &[Uuid]) { - if ids.is_empty() { - return; + .expect("tokio runtime for easytier-ffi"); + let manager = Arc::new(native_instance_manager_with_runtime( + runtime.handle().clone(), + )); + let process_management = + native_process_management(manager.clone(), Arc::new(FfiOwnedInstanceHooks)); + Self { + runtime, + manager, + process_management, + } } - - INSTANCE_NAME_ID_MAP.retain(|_, instance_id| !ids.contains(instance_id)); } -pub(crate) fn lock_remote_instance_mutation() -> tokio::sync::OwnedMutexGuard<()> { - INSTANCE_MANAGER - .remote_mutation_lock() - .blocking_lock_owned() +static FFI_CONTEXT: once_cell::sync::Lazy = once_cell::sync::Lazy::new(FfiContext::new); + +pub(crate) fn ffi_context() -> &'static FfiContext { + &FFI_CONTEXT } -pub(crate) fn instance_name_exists(inst_name: &str) -> bool { - find_instance_id_by_name(inst_name).is_some() +pub(crate) fn resolve_instance_id_by_name(inst_name: &str) -> Result, String> { + easytier_core::management::resolve_optional_instance_by_name( + ffi_context().manager.as_ref(), + inst_name, + ) + .map(|instance| instance.map(|instance| instance.instance_id())) + .map_err(|error| error.to_string()) } -pub(crate) fn find_instance_id_by_name(inst_name: &str) -> Option { - INSTANCE_NAME_ID_MAP - .get(inst_name) - .map(|id| *id) - .or_else(|| { - INSTANCE_MANAGER - .list_network_instance_ids() - .into_iter() - .find(|id| { - INSTANCE_MANAGER - .get_instance_name(id) - .is_some_and(|name| name == inst_name) - }) - }) +#[cfg(test)] +pub(crate) fn find_instance_id_by_name(inst_name: &str) -> Option { + resolve_instance_id_by_name(inst_name).ok().flatten() } diff --git a/easytier-contrib/easytier-ffi/src/tests.rs b/easytier-contrib/easytier-ffi/src/tests.rs index 8bbf01a7..03a78383 100644 --- a/easytier-contrib/easytier-ffi/src/tests.rs +++ b/easytier-contrib/easytier-ffi/src/tests.rs @@ -2,10 +2,7 @@ use crate::{ config_server::{ ConfigServerCallbackScope, ManagedConfigServerClientHooks, set_active_for_test, }, - state::{ - INSTANCE_MANAGER, INSTANCE_NAME_ID_MAP, find_instance_id_by_name, - lock_remote_instance_mutation, remove_instance_name_ids, - }, + state::{ffi_context, find_instance_id_by_name}, *, }; use easytier::{ @@ -15,7 +12,7 @@ use easytier::{ use serde_json::Value; use std::{ collections::HashSet, - ffi::{CStr, CString, c_char, c_void}, + ffi::{CStr, CString, c_char, c_int, c_void}, sync::{Mutex, mpsc}, time::Duration, }; @@ -101,10 +98,10 @@ fn list_instance_returns_instance_names_and_ids() { let cfg = TomlConfigLoader::default(); cfg.set_id(instance_id); cfg.set_inst_name(instance_name.clone()); - INSTANCE_MANAGER - .run_network_instance(cfg, false, ConfigFileControl::STATIC_CONFIG) + ffi_context() + .manager + .run_network_instance(cfg, ConfigFileControl::STATIC_CONFIG) .unwrap(); - INSTANCE_NAME_ID_MAP.insert(instance_name.clone(), instance_id); let mut infos = vec![ KeyValuePair { @@ -127,10 +124,14 @@ fn list_instance_returns_instance_names_and_ids() { } free_key_value_pairs(&infos[..count as usize]); - INSTANCE_MANAGER - .delete_network_instance(vec![instance_id]) + ffi_context() + .runtime + .block_on( + ffi_context() + .manager + .delete_network_instances([instance_id]), + ) .unwrap(); - remove_instance_name_ids(&[instance_id]); assert!(found); } @@ -261,8 +262,9 @@ async fn config_server_hooks_emit_run_event() { let inst_name = format!("test-{}", instance_id); cfg.set_inst_name(inst_name.clone()); hooks.pre_run_network_instance(&cfg).await.unwrap(); - INSTANCE_MANAGER - .run_network_instance(cfg, false, ConfigFileControl::STATIC_CONFIG) + ffi_context() + .manager + .run_network_instance(cfg, ConfigFileControl::STATIC_CONFIG) .unwrap(); hooks.post_run_network_instance(&instance_id).await.unwrap(); @@ -278,17 +280,18 @@ async fn config_server_hooks_emit_run_event() { ); assert_eq!(hooks.tracked_instance_ids(), vec![instance_id]); - let events = events.lock().unwrap(); + let events = events.lock().unwrap().clone(); assert_eq!(events.len(), 1); let event: Value = serde_json::from_str(&events[0]).unwrap(); assert_eq!(event["event"], "run_network_instance"); assert_eq!(event["success"], true); assert_eq!(event["instance_id"], instance_id.to_string()); assert!(event["error"].is_null()); - INSTANCE_MANAGER - .delete_network_instance(vec![instance_id]) + ffi_context() + .manager + .delete_network_instances([instance_id]) + .await .unwrap(); - remove_instance_name_ids(&[instance_id]); } #[tokio::test] @@ -306,8 +309,9 @@ async fn config_server_hooks_emit_delete_events_for_tracked_instances() { cfg.set_id(id); cfg.set_inst_name(format!("test-{}", id)); hooks.pre_run_network_instance(&cfg).await.unwrap(); - INSTANCE_MANAGER - .run_network_instance(cfg, false, ConfigFileControl::STATIC_CONFIG) + ffi_context() + .manager + .run_network_instance(cfg, ConfigFileControl::STATIC_CONFIG) .unwrap(); } @@ -327,7 +331,7 @@ async fn config_server_hooks_emit_delete_events_for_tracked_instances() { .unwrap(); assert!(hooks.tracked_instance_ids().is_empty()); - let events = events.lock().unwrap(); + let events = events.lock().unwrap().clone(); assert_eq!(events.len(), 2); let event_ids = events .iter() @@ -343,29 +347,27 @@ async fn config_server_hooks_emit_delete_events_for_tracked_instances() { event_ids, HashSet::from([instance_id_1.to_string(), instance_id_2.to_string()]) ); - INSTANCE_MANAGER - .delete_network_instance(vec![instance_id_1, instance_id_2]) + ffi_context() + .manager + .delete_network_instances([instance_id_1, instance_id_2]) + .await .unwrap(); - remove_instance_name_ids(&[instance_id_1, instance_id_2]); } #[tokio::test] -async fn config_server_hooks_remove_untracked_name_mapping_without_event() { +async fn config_server_hooks_ignore_untracked_instance_without_event() { let events: Mutex> = Mutex::new(Vec::new()); let hooks = ManagedConfigServerClientHooks::new( Some(record_config_server_event), &events as *const _ as *mut c_void, ); let local_id = Uuid::new_v4(); - let inst_name = format!("local-{}", local_id); - INSTANCE_NAME_ID_MAP.insert(inst_name.clone(), local_id); hooks .post_remove_network_instances(&[local_id]) .await .unwrap(); - assert!(INSTANCE_NAME_ID_MAP.get(&inst_name).is_none()); assert!(events.lock().unwrap().is_empty()); } @@ -375,15 +377,25 @@ async fn config_server_hooks_reject_duplicate_instance_name() { let inst_name = format!("test-{}", Uuid::new_v4()); let existing_id = Uuid::new_v4(); let new_id = Uuid::new_v4(); - INSTANCE_NAME_ID_MAP.insert(inst_name.clone(), existing_id); + let existing_cfg = TomlConfigLoader::default(); + existing_cfg.set_inst_name(inst_name.clone()); + existing_cfg.set_id(existing_id); + ffi_context() + .manager + .run_network_instance(existing_cfg, ConfigFileControl::STATIC_CONFIG) + .unwrap(); let cfg = TomlConfigLoader::default(); cfg.set_inst_name(inst_name.clone()); cfg.set_id(new_id); assert!(hooks.pre_run_network_instance(&cfg).await.is_err()); - assert_eq!(*INSTANCE_NAME_ID_MAP.get(&inst_name).unwrap(), existing_id); - INSTANCE_NAME_ID_MAP.remove(&inst_name); + assert_eq!(find_instance_id_by_name(&inst_name), Some(existing_id)); + ffi_context() + .manager + .delete_network_instances([existing_id]) + .await + .unwrap(); } #[tokio::test] @@ -398,8 +410,23 @@ async fn config_server_hooks_remove_overwritten_id_before_duplicate_name_error() let overwritten_id = Uuid::new_v4(); let duplicate_id = Uuid::new_v4(); hooks.instance_ids.lock().unwrap().insert(overwritten_id); - INSTANCE_NAME_ID_MAP.insert(old_name.clone(), overwritten_id); - INSTANCE_NAME_ID_MAP.insert(duplicate_name.clone(), duplicate_id); + for (id, name) in [ + (overwritten_id, old_name.clone()), + (duplicate_id, duplicate_name.clone()), + ] { + let cfg = TomlConfigLoader::default(); + cfg.set_id(id); + cfg.set_inst_name(name); + ffi_context() + .manager + .run_network_instance(cfg, ConfigFileControl::STATIC_CONFIG) + .unwrap(); + } + ffi_context() + .manager + .delete_network_instances([overwritten_id]) + .await + .unwrap(); hooks .post_remove_network_instances(&[overwritten_id]) @@ -412,13 +439,17 @@ async fn config_server_hooks_remove_overwritten_id_before_duplicate_name_error() assert!(hooks.pre_run_network_instance(&cfg).await.is_err()); assert!(hooks.tracked_instance_ids().is_empty()); - assert!(INSTANCE_NAME_ID_MAP.get(&old_name).is_none()); + assert!(find_instance_id_by_name(&old_name).is_none()); assert_eq!( - *INSTANCE_NAME_ID_MAP.get(&duplicate_name).unwrap(), - duplicate_id + find_instance_id_by_name(&duplicate_name), + Some(duplicate_id) ); assert_eq!(events.lock().unwrap().len(), 1); - INSTANCE_NAME_ID_MAP.remove(&duplicate_name); + ffi_context() + .manager + .delete_network_instances([duplicate_id]) + .await + .unwrap(); } #[tokio::test] @@ -427,11 +458,19 @@ async fn config_server_hooks_remove_tracked_state_before_overwrite_retry() { let inst_name = format!("test-{}", Uuid::new_v4()); let instance_id = Uuid::new_v4(); hooks.instance_ids.lock().unwrap().insert(instance_id); - INSTANCE_NAME_ID_MAP.insert(inst_name.clone(), instance_id); let cfg = TomlConfigLoader::default(); cfg.set_inst_name(inst_name.clone()); cfg.set_id(instance_id); + ffi_context() + .manager + .run_network_instance(cfg.clone(), ConfigFileControl::STATIC_CONFIG) + .unwrap(); + ffi_context() + .manager + .delete_network_instances([instance_id]) + .await + .unwrap(); hooks .post_remove_network_instances(&[instance_id]) @@ -440,7 +479,7 @@ async fn config_server_hooks_remove_tracked_state_before_overwrite_retry() { hooks.pre_run_network_instance(&cfg).await.unwrap(); assert!(hooks.tracked_instance_ids().is_empty()); - assert!(INSTANCE_NAME_ID_MAP.get(&inst_name).is_none()); + assert!(find_instance_id_by_name(&inst_name).is_none()); } #[tokio::test] @@ -451,11 +490,14 @@ async fn config_server_hooks_reject_post_run_after_external_delete() { cfg.set_id(instance_id); cfg.set_inst_name(format!("test-{}", instance_id)); hooks.pre_run_network_instance(&cfg).await.unwrap(); - INSTANCE_MANAGER - .run_network_instance(cfg, false, ConfigFileControl::STATIC_CONFIG) + ffi_context() + .manager + .run_network_instance(cfg, ConfigFileControl::STATIC_CONFIG) .unwrap(); - INSTANCE_MANAGER - .delete_network_instance(vec![instance_id]) + ffi_context() + .manager + .delete_network_instances([instance_id]) + .await .unwrap(); assert!(hooks.post_run_network_instance(&instance_id).await.is_err()); @@ -468,15 +510,20 @@ fn find_instance_id_by_name_resolves_uncommitted_manager_instance_name() { let cfg = TomlConfigLoader::default(); cfg.set_id(instance_id); cfg.set_inst_name(inst_name.clone()); - INSTANCE_MANAGER - .run_network_instance(cfg, false, ConfigFileControl::STATIC_CONFIG) + ffi_context() + .manager + .run_network_instance(cfg, ConfigFileControl::STATIC_CONFIG) .unwrap(); assert_eq!(find_instance_id_by_name(&inst_name), Some(instance_id)); - INSTANCE_MANAGER - .delete_network_instance(vec![instance_id]) + ffi_context() + .runtime + .block_on( + ffi_context() + .manager + .delete_network_instances([instance_id]), + ) .unwrap(); - remove_instance_name_ids(&[instance_id]); } #[test] @@ -493,10 +540,10 @@ fn delete_network_instance_removes_only_named_instances() { let cfg = TomlConfigLoader::default(); cfg.set_id(id); cfg.set_inst_name(name.clone()); - INSTANCE_MANAGER - .run_network_instance(cfg, false, ConfigFileControl::STATIC_CONFIG) + ffi_context() + .manager + .run_network_instance(cfg, ConfigFileControl::STATIC_CONFIG) .unwrap(); - INSTANCE_NAME_ID_MAP.insert(name, id); } let delete_name = CString::new(delete_name.clone()).unwrap(); @@ -509,10 +556,10 @@ fn delete_network_instance_removes_only_named_instances() { assert_eq!(find_instance_id_by_name(&keep_name), Some(keep_id)); assert!(find_instance_id_by_name(delete_name.to_str().unwrap()).is_none()); - INSTANCE_MANAGER - .delete_network_instance(vec![keep_id]) + ffi_context() + .runtime + .block_on(ffi_context().manager.delete_network_instances([keep_id])) .unwrap(); - remove_instance_name_ids(&[keep_id]); } #[test] @@ -532,13 +579,18 @@ fn retain_and_delete_network_instance_reject_invalid_name_pointers() { } #[test] -fn ffi_remote_mutation_lock_uses_manager_lock() { - let manager_guard = INSTANCE_MANAGER - .remote_mutation_lock() - .blocking_lock_owned(); +fn ffi_process_management_uses_manager_mutation_lock() { + let manager_guard = ffi_context().manager.mutation_lock().blocking_lock_owned(); let (done_tx, done_rx) = mpsc::channel(); let waiter = std::thread::spawn(move || { - let _ffi_guard = lock_remote_instance_mutation(); + ffi_context() + .runtime + .block_on( + ffi_context() + .process_management + .delete_owned_network_instances(Vec::new()), + ) + .unwrap(); done_tx.send(()).unwrap(); }); @@ -549,7 +601,7 @@ fn ffi_remote_mutation_lock_uses_manager_lock() { } #[tokio::test] -async fn config_server_hooks_suppress_late_run_events_while_stopping() { +async fn config_server_hooks_reject_late_runs_for_core_rollback() { let events: Mutex> = Mutex::new(Vec::new()); let hooks = ManagedConfigServerClientHooks::new( Some(record_config_server_event), @@ -557,15 +609,54 @@ async fn config_server_hooks_suppress_late_run_events_while_stopping() { ); hooks.start_stopping(); - hooks - .post_run_network_instance(&Uuid::new_v4()) - .await - .unwrap(); + assert!( + hooks + .post_run_network_instance(&Uuid::new_v4()) + .await + .is_err() + ); assert!(hooks.tracked_instance_ids().is_empty()); assert!(events.lock().unwrap().is_empty()); } +#[test] +fn delete_network_instance_rejects_an_ambiguous_name() { + let duplicate_name = format!("duplicate-{}", Uuid::new_v4()); + let instance_ids = [Uuid::new_v4(), Uuid::new_v4()]; + for instance_id in instance_ids { + let config = TomlConfigLoader::default(); + config.set_id(instance_id); + config.set_inst_name(duplicate_name.clone()); + ffi_context() + .manager + .run_network_instance(config, ConfigFileControl::STATIC_CONFIG) + .unwrap(); + } + + let duplicate_name = CString::new(duplicate_name).unwrap(); + let names = [duplicate_name.as_ptr()]; + assert_eq!( + unsafe { delete_network_instance(names.as_ptr(), names.len()) }, + -1 + ); + assert!(take_last_error().unwrap().contains("2 instances match")); + assert!( + instance_ids + .iter() + .all(|id| ffi_context().manager.instance(*id).is_some()) + ); + + ffi_context() + .runtime + .block_on( + ffi_context() + .process_management + .delete_owned_network_instances(instance_ids.to_vec()), + ) + .unwrap(); +} + #[test] fn config_server_callback_context_rejects_nested_blocking_ffi_calls() { let _callback_scope = ConfigServerCallbackScope::enter(); @@ -615,112 +706,12 @@ fn config_server_callback_context_rejects_nested_blocking_ffi_calls() { #[cfg(feature = "ffi-dataplane")] { + let mut session = 0; assert_eq!( - unsafe { - data_plane_tcp_connect( - std::ptr::null(), - std::ptr::null(), - 0, - 0, - std::ptr::null_mut(), - std::ptr::null_mut(), - ) - }, - 0 + unsafe { data_plane_session_open(std::ptr::null(), &mut session) }, + -(easytier_core::gateway::DataPlaneErrorKind::Io as c_int) ); - assert_eq!( - unsafe { - data_plane_tcp_bind( - std::ptr::null(), - 0, - 0, - std::ptr::null_mut(), - std::ptr::null_mut(), - ) - }, - 0 - ); - assert_eq!( - unsafe { - data_plane_tcp_accept( - 0, - 0, - std::ptr::null_mut(), - std::ptr::null_mut(), - std::ptr::null_mut(), - std::ptr::null_mut(), - ) - }, - 0 - ); - assert_eq!( - unsafe { data_plane_tcp_read(0, std::ptr::null_mut(), 0, 0) }, - -1 - ); - assert_eq!( - unsafe { data_plane_tcp_write(0, std::ptr::null(), 0, 0) }, - -1 - ); - assert_eq!(data_plane_tcp_close(0), -1); - assert_eq!(data_plane_tcp_listener_close(0), -1); - assert_eq!( - unsafe { - data_plane_udp_bind( - std::ptr::null(), - 0, - 0, - std::ptr::null_mut(), - std::ptr::null_mut(), - ) - }, - 0 - ); - assert_eq!( - unsafe { data_plane_udp_send_to(0, std::ptr::null(), 0, std::ptr::null(), 0, 0) }, - -1 - ); - assert_eq!( - unsafe { - data_plane_udp_recv_from( - 0, - std::ptr::null_mut(), - 0, - std::ptr::null_mut(), - std::ptr::null_mut(), - 0, - ) - }, - -1 - ); - assert_eq!(data_plane_udp_close(0), -1); - assert_eq!(data_plane_async_op_status(0), -2); - assert_eq!(data_plane_async_op_wait(0, 0), -2); - assert_eq!(data_plane_async_op_cancel(0), -2); - assert_eq!(data_plane_async_op_free(0), -2); - data_plane_free_bytes(std::ptr::null(), 0); - assert_eq!( - unsafe { data_plane_tcp_connect_start(std::ptr::null(), std::ptr::null(), 0, 0) }, - 0 - ); - assert_eq!( - unsafe { data_plane_tcp_bind_start(std::ptr::null(), 0, 0) }, - 0 - ); - assert_eq!(unsafe { data_plane_tcp_accept_start(0, 0) }, 0); - assert_eq!(unsafe { data_plane_tcp_read_start(0, 0, 0) }, 0); - assert_eq!( - unsafe { data_plane_tcp_write_start(0, std::ptr::null(), 0, 0) }, - 0 - ); - assert_eq!( - unsafe { data_plane_udp_bind_start(std::ptr::null(), 0, 0) }, - 0 - ); - assert_eq!( - unsafe { data_plane_udp_send_to_start(0, std::ptr::null(), 0, std::ptr::null(), 0, 0) }, - 0 - ); - assert_eq!(unsafe { data_plane_udp_recv_from_start(0, 0, 0) }, 0); + assert_eq!(session, 0); } } @@ -729,38 +720,23 @@ fn config_server_callback_context_rejects_nested_blocking_ffi_calls() { fn active_config_server_rejects_data_plane() { set_active_for_test(true); + let name = CString::new("missing").unwrap(); + let mut session = 0; assert_eq!( - unsafe { - data_plane_tcp_connect( - std::ptr::null(), - std::ptr::null(), - 0, - 0, - std::ptr::null_mut(), - std::ptr::null_mut(), - ) - }, - 0 + unsafe { data_plane_session_open(name.as_ptr(), &mut session) }, + -(easytier_core::gateway::DataPlaneErrorKind::Io as c_int) ); - assert_eq!( - unsafe { data_plane_tcp_read(0, std::ptr::null_mut(), 0, 0) }, - -1 - ); - assert_eq!( - unsafe { data_plane_tcp_connect_start(std::ptr::null(), std::ptr::null(), 0, 0) }, - 0 - ); - assert_eq!(unsafe { data_plane_tcp_read_start(0, 0, 0) }, 0); + assert_eq!(session, 0); set_active_for_test(false); } #[cfg(feature = "ffi-dataplane")] #[test] -fn async_op_invalid_handle_helpers_are_stable() { - assert_eq!(data_plane_async_op_status(u64::MAX), -2); - assert_eq!(data_plane_async_op_wait(u64::MAX, 1), -2); - assert_eq!(data_plane_async_op_cancel(u64::MAX), -2); - assert_eq!(data_plane_async_op_free(u64::MAX), -2); - data_plane_free_bytes(std::ptr::null(), 0); +fn data_plane_invalid_handle_errors_are_stable() { + let closed = -(easytier_core::gateway::DataPlaneErrorKind::HandleClosed as c_int); + assert_eq!(data_plane_completion_wait(u64::MAX, 0), closed); + assert_eq!(data_plane_operation_cancel(u64::MAX, 1), closed); + assert_eq!(data_plane_operation_free(u64::MAX, 1), closed); + assert_eq!(data_plane_resource_close(u64::MAX, 1), closed); } diff --git a/easytier-contrib/easytier-ffi/src/types.rs b/easytier-contrib/easytier-ffi/src/types.rs index 7701bc97..bc554f15 100644 --- a/easytier-contrib/easytier-ffi/src/types.rs +++ b/easytier-contrib/easytier-ffi/src/types.rs @@ -8,3 +8,23 @@ pub struct KeyValuePair { } pub type ConfigServerEventCallback = Option; + +#[repr(C)] +#[derive(Clone, Copy, Debug, Default, Eq, PartialEq)] +pub struct DataPlaneSocketAddr { + /// `4` for IPv4. Other families are reserved for later ABI versions. + pub family: u16, + /// Native-endian port number. + pub port: u16, + /// Network-order address bytes. IPv4 uses the first four bytes. + pub address: [u8; 16], +} + +#[repr(C)] +#[derive(Clone, Copy, Debug, Default, Eq, PartialEq)] +pub struct DataPlaneCompletion { + pub operation_id: u64, + pub operation_kind: u16, + /// `0` for success, otherwise a stable `DataPlaneErrorKind` value. + pub status: u16, +} diff --git a/easytier-contrib/easytier-ohrs/Cargo.lock b/easytier-contrib/easytier-ohrs/Cargo.lock index bd13d008..b257be62 100644 --- a/easytier-contrib/easytier-ohrs/Cargo.lock +++ b/easytier-contrib/easytier-ohrs/Cargo.lock @@ -150,6 +150,16 @@ version = "1.7.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "69f7f8c3906b62b754cd5326047894316021dcfe5a194c8ea52bdd94934a3457" +[[package]] +name = "ariadne" +version = "0.5.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "36f5e3dca4e09a6f340a61a0e9c7b61e030c69fc27bf29d73218f7e5e3b7638f" +dependencies = [ + "unicode-width 0.1.11", + "yansi", +] + [[package]] name = "arrayvec" version = "0.7.6" @@ -188,28 +198,6 @@ dependencies = [ "ringbuf", ] -[[package]] -name = "async-stream" -version = "0.3.6" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "0b5a71a6f37880a80d1d7f19efd781e4b5de42c88f0722cc13bcb6cc2cfe8476" -dependencies = [ - "async-stream-impl", - "futures-core", - "pin-project-lite", -] - -[[package]] -name = "async-stream-impl" -version = "0.3.6" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "c7c24de15d275a1ecfd47a380fb4d5ec9bfe0933f309ed5e705b775596a3574d" -dependencies = [ - "proc-macro2", - "quote", - "syn 2.0.106", -] - [[package]] name = "async-trait" version = "0.1.89" @@ -946,17 +934,6 @@ version = "2.9.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "2a2330da5de22e8a3cb63252ce2abb30116bf5265e89c0e01bc17015ce30a476" -[[package]] -name = "dbus" -version = "0.9.9" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "190b6255e8ab55a7b568df5a883e9497edc3e4821c06396612048b430e5ad1e9" -dependencies = [ - "libc", - "libdbus-sys", - "windows-sys 0.59.0", -] - [[package]] name = "deflate64" version = "0.1.9" @@ -1174,18 +1151,15 @@ version = "2.6.4" dependencies = [ "anyhow", "arc-swap", + "ariadne", "async-recursion", - "async-ringbuf", - "async-stream", "async-trait", "atomic-shim", "atomic_refcell", "auto_impl", "base64 0.22.1", - "bitflags 2.9.4", "bon", "boringtun-easytier", - "bytecodec", "byteorder", "bytes", "cfg_aliases", @@ -1196,11 +1170,10 @@ dependencies = [ "clap_complete_nushell", "crossbeam", "dashmap", - "dbus", - "delegate", - "derivative", "derive_builder", "derive_more", + "easytier-core", + "easytier-proto", "encoding", "flume", "forwarded-header-value", @@ -1213,75 +1186,52 @@ dependencies = [ "hickory-proto", "hickory-resolver", "hickory-server", - "hmac", "http", - "http_req", "humansize", "humantime-serde", - "idna", "igd-next", "indoc", - "itertools 0.14.0", "kcp-sys", "machine-uid", "moka", - "multimap", "natpmp", "netlink-packet-core", "netlink-packet-route 0.21.0", - "netlink-packet-utils", "netlink-sys", "network-interface", "nix 0.29.0", "once_cell", - "ordered_hash_map", "parking_lot", "paste", - "pbjson", - "pbjson-build", "percent-encoding", - "petgraph", "pin-project-lite", "pnet", - "prefix-trie", - "proc-macro2", "prost 0.14.3", - "prost-build", "prost-reflect 0.16.4", - "prost-reflect-build", - "prost-wkt-types", + "quanta", "quinn", - "quinn-plaintext", - "quote", + "quinn-proto", "rand 0.8.5", "rcgen", "regex", - "reqwest", - "resolv-conf", "ring", - "ringbuf", "rust-i18n", "rustls", + "seahash", "serde", "serde_json", "service-manager", - "sha2", "shellexpand", - "smoltcp", - "snow", "socket2 0.5.10", "strum", - "stun_codec", "sys-locale", "tabled", "terminal_size", "thiserror 1.0.69", "thunk-rs", "time", - "timedmap", "tokio", "tokio-rustls", - "tokio-stream", "tokio-util", "tokio-websockets", "toml", @@ -1291,17 +1241,75 @@ dependencies = [ "unicode-width 0.1.11", "url", "uuid", - "version-compare", - "which 7.0.3", - "wildmatch", "winapi", "windivert", "windows 0.62.2", "windows-service", "winreg 0.52.0", + "zerocopy 0.7.35", +] + +[[package]] +name = "easytier-core" +version = "2.6.4" +dependencies = [ + "aes-gcm", + "anyhow", + "arc-swap", + "ariadne", + "async-ringbuf", + "async-trait", + "atomic-shim", + "auto_impl", + "base64 0.22.1", + "bitflags 2.9.4", + "bytecodec", + "bytes", + "chacha20poly1305", + "cidr", + "crossbeam", + "dashmap", + "derive_builder", + "easytier-proto", + "futures", + "guarden", + "hmac", + "http-body-util", + "hyper", + "hyper-util", + "idna", + "ordered_hash_map", + "parking_lot", + "percent-encoding", + "petgraph", + "pin-project-lite", + "pnet_packet", + "prefix-trie", + "prost 0.14.3", + "prost-reflect 0.16.4", + "prost-wkt-types", + "quanta", + "rand 0.8.5", + "rustls", + "serde", + "serde_json", + "sha2", + "smoltcp", + "snow", + "strum", + "stun_codec", + "thiserror 1.0.69", + "tokio", + "tokio-rustls", + "tokio-util", + "toml", + "tracing", + "url", + "uuid", + "webpki-roots 0.26.11", + "wildmatch", "x25519-dalek", "zerocopy 0.7.35", - "zip", "zstd", ] @@ -1331,6 +1339,43 @@ dependencies = [ "uuid", ] +[[package]] +name = "easytier-proto" +version = "2.6.4" +dependencies = [ + "anyhow", + "async-trait", + "auto_impl", + "base64 0.22.1", + "bytes", + "chrono", + "cidr", + "delegate", + "derivative", + "derive_more", + "hmac", + "indoc", + "pbjson", + "pbjson-build", + "proc-macro2", + "prost 0.14.3", + "prost-build", + "prost-reflect 0.16.4", + "prost-reflect-build", + "prost-wkt-types", + "quote", + "reqwest", + "serde", + "serde_json", + "sha2", + "thiserror 1.0.69", + "tokio", + "url", + "uuid", + "x25519-dalek", + "zip", +] + [[package]] name = "either" version = "1.15.0" @@ -1439,12 +1484,6 @@ dependencies = [ "syn 2.0.106", ] -[[package]] -name = "env_home" -version = "0.1.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "c7f84e12ccf0a7ddc17a6c41c93326024c42920d7ee630d04950e6926645c0fe" - [[package]] name = "equivalent" version = "1.0.2" @@ -1858,21 +1897,22 @@ dependencies = [ [[package]] name = "guarden" -version = "0.1.3" +version = "0.2.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "31c7272e004bec8ea7fe50b2ec5451858695bb2743e897c353753fcb3415f4ef" +checksum = "b8408903291a7d0cc74169d5de4dd1919a9a402a2f67fcd7df3303ed045fae73" dependencies = [ - "futures", + "futures-core", "guarden-macros", "tokio", ] [[package]] name = "guarden-macros" -version = "0.1.3" +version = "0.2.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "2d291d94f41471fe84384a426b3e2c9d22f960a351a5bf26aaa7cd75fbc02c88" +checksum = "1e0ef28f1077c259f9e7e238e234a78ce18cedbf0251fd2135f5fc23c40e79fe" dependencies = [ + "proc-macro-crate", "proc-macro2", "quote", "syn 2.0.106", @@ -1928,9 +1968,9 @@ dependencies = [ [[package]] name = "hashbrown" -version = "0.16.0" +version = "0.17.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "5419bdc4f6a9207fbeba6d11b604d481addf78ecd10c11ad51e76c2f6482748d" +checksum = "ed5909b6e89a2db4456e54cd5f673791d7eca6732202bbf2a9cc504fe2f9b84a" [[package]] name = "hashlink" @@ -2112,22 +2152,6 @@ dependencies = [ "pin-project-lite", ] -[[package]] -name = "http_req" -version = "0.13.1" -source = "git+https://github.com/EasyTier/http_req.git#b10aa9fc0db3067cc3d2174683a87250b80a1ea9" -dependencies = [ - "base64 0.22.1", - "rand 0.8.5", - "rustls", - "rustls-pemfile", - "rustls-pki-types", - "unicase", - "webpki", - "webpki-roots 0.26.11", - "zeroize", -] - [[package]] name = "httparse" version = "1.10.1" @@ -2420,12 +2444,12 @@ dependencies = [ [[package]] name = "indexmap" -version = "2.11.4" +version = "2.14.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "4b0f83760fb341a774ed326568e19f5a863af4a952def8c39f9ab92fd95b88e5" +checksum = "d466e9454f08e4a911e14806c24e16fba1b4c121d1ea474396f396069cf949d9" dependencies = [ "equivalent", - "hashbrown 0.16.0", + "hashbrown 0.17.1", "serde", "serde_core", ] @@ -2637,16 +2661,6 @@ version = "0.2.186" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "68ab91017fe16c622486840e4c83c9a37afeff978bd239b5293d61ece587de66" -[[package]] -name = "libdbus-sys" -version = "0.2.6" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "5cbe856efeb50e4681f010e9aaa2bf0a644e10139e54cde10fc83a307c23bd9f" -dependencies = [ - "cc", - "pkg-config", -] - [[package]] name = "libloading" version = "0.8.9" @@ -2871,9 +2885,6 @@ name = "multimap" version = "0.10.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "1d87ecb2933e8aeadb3e3a02b828fed80a7528047e68b4f424523a0981a3a084" -dependencies = [ - "serde", -] [[package]] name = "napi-build-ohos" @@ -3236,6 +3247,15 @@ dependencies = [ "vcpkg", ] +[[package]] +name = "ordered-float" +version = "2.10.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "68f19d67e5a2795c94e73e0bb1cc1a7edeb2e28efd39e2e1c9b7a40c1108b11c" +dependencies = [ + "num-traits", +] + [[package]] name = "ordered_hash_map" version = "0.5.0" @@ -3564,6 +3584,15 @@ dependencies = [ "syn 2.0.106", ] +[[package]] +name = "proc-macro-crate" +version = "3.5.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e67ba7e9b2b56446f1d419b1d807906278ffa1a658a8a5d8a39dcb1f5a78614f" +dependencies = [ + "toml_edit 0.25.8+spec-1.1.0", +] + [[package]] name = "proc-macro-error" version = "1.0.4" @@ -3702,9 +3731,12 @@ version = "0.16.4" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "590aa145fee8f7a26b5a6055365e7c5e89a5c1caae9869de76ec0ee73181a2f9" dependencies = [ + "base64 0.22.1", "prost 0.14.3", "prost-reflect-derive 0.16.0", "prost-types 0.14.3", + "serde", + "serde-value", ] [[package]] @@ -3803,6 +3835,21 @@ dependencies = [ "serde_json", ] +[[package]] +name = "quanta" +version = "0.12.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f3ab5a9d756f0d97bdc89019bd2e4ea098cf9cde50ee7564dde6b81ccc8f06c7" +dependencies = [ + "crossbeam-utils", + "libc", + "once_cell", + "raw-cpuid", + "wasi 0.11.1+wasi-snapshot-preview1", + "web-sys", + "winapi", +] + [[package]] name = "quick-xml" version = "0.38.3" @@ -3832,18 +3879,6 @@ dependencies = [ "web-time", ] -[[package]] -name = "quinn-plaintext" -version = "0.3.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "f3e617feaeb6493018fa35fc47ae8b630ac8903d8159e9e747018841b99bad3d" -dependencies = [ - "bytes", - "quinn-proto", - "seahash", - "tracing", -] - [[package]] name = "quinn-proto" version = "0.11.14" @@ -3988,6 +4023,15 @@ version = "0.10.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "63b8176103e19a2643978565ca18b50549f6101881c443590420e4dc998a3c69" +[[package]] +name = "raw-cpuid" +version = "11.6.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "498cd0dc59d73224351ee52a95fee0f1a617a2eae0e7d9d720cc622c73a54186" +dependencies = [ + "bitflags 2.9.4", +] + [[package]] name = "rcgen" version = "0.12.1" @@ -4257,15 +4301,6 @@ dependencies = [ "security-framework 3.5.1", ] -[[package]] -name = "rustls-pemfile" -version = "2.2.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "dce314e5fee3f39953d46bb63bb8a46d40c2f8fb7cc5a3b6cab2bde9721d6e50" -dependencies = [ - "rustls-pki-types", -] - [[package]] name = "rustls-pki-types" version = "1.12.0" @@ -4414,6 +4449,16 @@ dependencies = [ "serde_derive", ] +[[package]] +name = "serde-value" +version = "0.7.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f3a1a3341211875ef120e117ea7fd5228530ae7e7036a779fdc9117be6b3282c" +dependencies = [ + "ordered-float", + "serde", +] + [[package]] name = "serde_core" version = "1.0.226" @@ -4492,7 +4537,7 @@ dependencies = [ "encoding_rs", "plist", "sys-info", - "which 4.4.2", + "which", "xml-rs", ] @@ -4915,12 +4960,6 @@ dependencies = [ "time-core", ] -[[package]] -name = "timedmap" -version = "1.0.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "825f6c8a18bc36d56a62f66af7296385b628c9c5543a8663d4c217fc920bfefd" - [[package]] name = "tinystr" version = "0.8.1" @@ -5022,8 +5061,7 @@ dependencies = [ [[package]] name = "tokio-websockets" version = "0.13.2" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "dad543404f98bfc969aeb71994105c592acfc6c43323fddcd016bb208d1c65cb" +source = "git+https://github.com/EasyTier/tokio-websockets#dc9771c7c215882349c3cb328877550a3593df21" dependencies = [ "base64 0.22.1", "bytes", @@ -5049,8 +5087,8 @@ checksum = "dc1beb996b9d83529a9e75c17a1686767d148d70663143c7854d8b4a09ced362" dependencies = [ "serde", "serde_spanned", - "toml_datetime", - "toml_edit", + "toml_datetime 0.6.11", + "toml_edit 0.22.27", ] [[package]] @@ -5062,6 +5100,15 @@ dependencies = [ "serde", ] +[[package]] +name = "toml_datetime" +version = "1.1.0+spec-1.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "97251a7c317e03ad83774a8752a7e81fb6067740609f75ea2b585b569a59198f" +dependencies = [ + "serde_core", +] + [[package]] name = "toml_edit" version = "0.22.27" @@ -5071,9 +5118,30 @@ dependencies = [ "indexmap", "serde", "serde_spanned", - "toml_datetime", + "toml_datetime 0.6.11", "toml_write", - "winnow", + "winnow 0.7.13", +] + +[[package]] +name = "toml_edit" +version = "0.25.8+spec-1.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "16bff38f1d86c47f9ff0647e6838d7bb362522bdf44006c7068c2b1e606f1f3c" +dependencies = [ + "indexmap", + "toml_datetime 1.1.0+spec-1.1.0", + "toml_parser", + "winnow 1.0.3", +] + +[[package]] +name = "toml_parser" +version = "1.1.2+spec-1.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a2abe9b86193656635d2411dc43050282ca48aa31c2451210f4202550afb7526" +dependencies = [ + "winnow 1.0.3", ] [[package]] @@ -5411,12 +5479,6 @@ version = "0.2.15" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "accd4ea62f7bb7a82fe23066fb0957d48ef677f6eeb8215f372f52e48bb32426" -[[package]] -name = "version-compare" -version = "0.2.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "852e951cb7832cb45cb1169900d19760cfa39b82bc0ea9c0e5a14ae88411c98b" - [[package]] name = "version_check" version = "0.9.5" @@ -5601,16 +5663,6 @@ dependencies = [ "wasm-bindgen", ] -[[package]] -name = "webpki" -version = "0.22.4" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "ed63aea5ce73d0ff405984102c42de94fc55a6b75765d621c65262469b3c9b53" -dependencies = [ - "ring", - "untrusted", -] - [[package]] name = "webpki-root-certs" version = "1.0.5" @@ -5650,18 +5702,6 @@ dependencies = [ "rustix 0.38.44", ] -[[package]] -name = "which" -version = "7.0.3" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "24d643ce3fd3e5b54854602a080f34fb10ab75e0b813ee32d00ca2b44fa74762" -dependencies = [ - "either", - "env_home", - "rustix 1.1.2", - "winsafe", -] - [[package]] name = "widestring" version = "1.2.0" @@ -6262,6 +6302,15 @@ dependencies = [ "memchr", ] +[[package]] +name = "winnow" +version = "1.0.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0592e1c9d151f854e6fd382574c3a0855250e1d9b2f99d9281c6e6391af352f1" +dependencies = [ + "memchr", +] + [[package]] name = "winreg" version = "0.50.0" @@ -6282,12 +6331,6 @@ dependencies = [ "windows-sys 0.48.0", ] -[[package]] -name = "winsafe" -version = "0.0.19" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "d135d17ab770252ad95e9a872d365cf3090e3be864a34ab46f48555993efc904" - [[package]] name = "wintun" version = "0.5.1" @@ -6428,6 +6471,12 @@ dependencies = [ "xml-rs", ] +[[package]] +name = "yansi" +version = "1.0.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "cfe53a6657fd280eaa890a3bc59152892ffa3e30101319d168b781ed6529b049" + [[package]] name = "yasna" version = "0.5.2" diff --git a/easytier-contrib/easytier-ohrs/src/config_repo/import_export.rs b/easytier-contrib/easytier-ohrs/src/config_repo/import_export.rs index 7f698aa4..16c07dd1 100644 --- a/easytier-contrib/easytier-ohrs/src/config_repo/import_export.rs +++ b/easytier-contrib/easytier-ohrs/src/config_repo/import_export.rs @@ -1,4 +1,5 @@ use crate::config::types::stored_config::{ExportTomlResult, StoredConfigRecord}; +use easytier::common::config::NetworkConfigExt; use easytier::common::config::{ConfigLoader, TomlConfigLoader}; use easytier::proto::api::manage::NetworkConfig; diff --git a/easytier-contrib/easytier-ohrs/src/config_repo/validation.rs b/easytier-contrib/easytier-ohrs/src/config_repo/validation.rs index 7e91fd6b..6bcbfae5 100644 --- a/easytier-contrib/easytier-ohrs/src/config_repo/validation.rs +++ b/easytier-contrib/easytier-ohrs/src/config_repo/validation.rs @@ -1,3 +1,4 @@ +use easytier::common::config::NetworkConfigExt; use easytier::proto::api::manage::NetworkConfig; use serde_json::{Map, Value}; use uuid::Uuid; diff --git a/easytier-contrib/easytier-ohrs/src/exports/runtime_api.rs b/easytier-contrib/easytier-ohrs/src/exports/runtime_api.rs index f1792a03..86adc30d 100644 --- a/easytier-contrib/easytier-ohrs/src/exports/runtime_api.rs +++ b/easytier-contrib/easytier-ohrs/src/exports/runtime_api.rs @@ -36,8 +36,8 @@ pub(crate) fn stop_kernel( return false; }; - let ret = INSTANCE_MANAGER - .delete_network_instance(vec![instance_id]) + let ret = ASYNC_RUNTIME + .block_on(INSTANCE_MANAGER.delete_network_instances([instance_id])) .map(|_| true) .unwrap_or_else(|err| { ohrs_log_error!("[Rust] stop_kernel failed {}: {}", config_id, err); @@ -46,7 +46,7 @@ pub(crate) fn stop_kernel( if ret { clear_runtime_config_snapshot(&config_id); } - let has_active_instances = !INSTANCE_MANAGER.list_network_instance_ids().is_empty(); + let has_active_instances = !INSTANCE_MANAGER.instance_ids().is_empty(); let has_web_clients = WEB_CLIENTS .lock() .map(|guard| !guard.is_empty()) @@ -102,7 +102,7 @@ pub(crate) fn set_tun_fd( }; INSTANCE_MANAGER - .set_tun_fd(&instance_id, fd) + .attach_tun_fd(instance_id, fd) .map(|_| { mark_tun_attached(&config_id); ohrs_log_info!( diff --git a/easytier-contrib/easytier-ohrs/src/kernel_bridge/socket_server.rs b/easytier-contrib/easytier-ohrs/src/kernel_bridge/socket_server.rs index 69733bfe..7ae0771d 100644 --- a/easytier-contrib/easytier-ohrs/src/kernel_bridge/socket_server.rs +++ b/easytier-contrib/easytier-ohrs/src/kernel_bridge/socket_server.rs @@ -9,8 +9,7 @@ use crate::runtime::state::runtime_state::{ }; use crate::{ASYNC_RUNTIME, INSTANCE_MANAGER}; use easytier::common::global_ctx::{EventBusSubscriber, GlobalCtxEvent}; -use easytier::proto::api::instance::ListPeerRequest; -use easytier::proto::rpc_types::controller::BaseController; +use easytier::instance::factory::subscribe_native_instance_event; use once_cell::sync::Lazy; use serde::Serialize; use std::collections::{HashMap, HashSet}; @@ -103,11 +102,11 @@ fn shrink_hash_set_if_sparse(set: &mut HashSet) { fn sync_tun_event_receivers(receivers: &mut HashMap) { let mut active_instance_ids = HashSet::new(); - for instance in INSTANCE_MANAGER.iter() { - let instance_id = instance.key().to_string(); + for instance in INSTANCE_MANAGER.instances() { + let instance_id = instance.instance_id().to_string(); active_instance_ids.insert(instance_id.clone()); if !receivers.contains_key(&instance_id) - && let Some(receiver) = instance.value().subscribe_event() + && let Some(receiver) = subscribe_native_instance_event(&instance) { receivers.insert(instance_id, receiver); } @@ -226,34 +225,17 @@ fn tun_candidate_ids(snapshot: &RuntimeAggregateState) -> HashSet { } fn collect_traffic_stats() -> TrafficStatsPayload { - let services = INSTANCE_MANAGER - .iter() - .filter_map(|instance| { - instance - .value() - .get_api_service() - .map(|api_service| (instance.key().to_string(), api_service)) - }) + let running_instances = INSTANCE_MANAGER + .instances() + .into_iter() + .filter(|instance| instance.is_ready()) .collect::>(); let instances = ASYNC_RUNTIME.block_on(async { let mut instances = Vec::new(); - for (instance_id, api_service) in services { - let peers = match api_service - .get_peer_manage_service() - .list_peer(BaseController::default(), ListPeerRequest::default()) - .await - { - Ok(response) => response.peer_infos, - Err(err) => { - ohrs_log_debug!( - "[Rust] collect traffic stats list_peer failed instance={}: {}", - instance_id, - err - ); - continue; - } - }; + for instance in running_instances { + let instance_id = instance.instance_id().to_string(); + let peers = instance.peer_snapshots().await; let mut instance_rx_bytes = 0i64; let mut instance_tx_bytes = 0i64; diff --git a/easytier-contrib/easytier-ohrs/src/lib.rs b/easytier-contrib/easytier-ohrs/src/lib.rs index 3fa5f7fd..47ea077f 100644 --- a/easytier-contrib/easytier-ohrs/src/lib.rs +++ b/easytier-contrib/easytier-ohrs/src/lib.rs @@ -53,12 +53,13 @@ use config::services::share_link_service::{ }; use config::storage::config_meta::get_config_display_name; use config::types::stored_config::{KeyValuePair, SharedConfigLinkPayload, SnapshotImportResult}; +use easytier::common::config::NetworkConfigExt; use easytier::common::constants::EASYTIER_VERSION; use easytier::common::{ MachineIdOptions, config::{ConfigFileControl, ConfigLoader, TomlConfigLoader}, }; -use easytier::instance_manager::NetworkInstanceManager; +use easytier::instance::factory::{NativeInstanceManager, native_instance_manager_with_runtime}; use easytier::proto::api::manage::NetworkConfig; use easytier::proto::api::manage::NetworkingMethod; use easytier::web_client::{WebClient, WebClientHooks, run_web_client}; @@ -74,14 +75,18 @@ use std::sync::{Arc, Mutex}; use tokio::runtime::{Builder, Runtime}; use uuid::Uuid; -pub(crate) static INSTANCE_MANAGER: once_cell::sync::Lazy> = - once_cell::sync::Lazy::new(|| Arc::new(NetworkInstanceManager::new())); static ASYNC_RUNTIME: once_cell::sync::Lazy = once_cell::sync::Lazy::new(|| { Builder::new_multi_thread() .enable_all() .build() .expect("tokio runtime for easytier-ohrs") }); +pub(crate) static INSTANCE_MANAGER: once_cell::sync::Lazy> = + once_cell::sync::Lazy::new(|| { + Arc::new(native_instance_manager_with_runtime( + ASYNC_RUNTIME.handle().clone(), + )) + }); static WEB_CLIENTS: once_cell::sync::Lazy>> = once_cell::sync::Lazy::new(|| Mutex::new(HashMap::new())); @@ -151,8 +156,8 @@ fn stop_web_client(config_id: &str) -> bool { return true; } - let ret = INSTANCE_MANAGER - .delete_network_instance(tracked_ids) + let ret = ASYNC_RUNTIME + .block_on(INSTANCE_MANAGER.delete_network_instances(tracked_ids)) .map(|_| true) .unwrap_or_else(|err| { ohrs_log_error!( @@ -171,7 +176,7 @@ fn ensure_local_socket_server_started() -> bool { } fn maybe_stop_local_socket_server() { - let no_local_instances = INSTANCE_MANAGER.list_network_instance_ids().is_empty(); + let no_local_instances = INSTANCE_MANAGER.instance_ids().is_empty(); let no_web_clients = WEB_CLIENTS .lock() .map(|guard| guard.is_empty()) @@ -182,12 +187,7 @@ fn maybe_stop_local_socket_server() { } fn run_config_server_instance(config_id: &str, config: &NetworkConfig) -> bool { - if INSTANCE_MANAGER - .list_network_instance_ids() - .iter() - .next() - .is_some() - { + if INSTANCE_MANAGER.instance_ids().iter().next().is_some() { ohrs_log_error!("[Rust] there is a running instance!"); return false; } @@ -293,7 +293,7 @@ pub(crate) fn run_network_instance_from_json(cfg_json: &str) -> bool { } }; - if !INSTANCE_MANAGER.list_network_instance_ids().is_empty() { + if !INSTANCE_MANAGER.instance_ids().is_empty() { ohrs_log_error!("[Rust] there is a running instance!"); return false; } @@ -303,15 +303,12 @@ pub(crate) fn run_network_instance_from_json(cfg_json: &str) -> bool { } let inst_id = cfg.get_id(); - if INSTANCE_MANAGER - .list_network_instance_ids() - .contains(&inst_id) - { + if INSTANCE_MANAGER.instance_ids().contains(&inst_id) { ohrs_log_error!("[Rust] instance {} already exists", inst_id); return false; } - match INSTANCE_MANAGER.run_network_instance(cfg, false, ConfigFileControl::STATIC_CONFIG) { + match INSTANCE_MANAGER.run_network_instance(cfg, ConfigFileControl::STATIC_CONFIG) { Ok(_) => { cache_runtime_config_snapshot(inst_id.to_string(), inst_id.to_string(), config); true diff --git a/easytier-contrib/easytier-uptime/src/health_checker.rs b/easytier-contrib/easytier-uptime/src/health_checker.rs index 13c0f91e..fa5bda46 100644 --- a/easytier-contrib/easytier-uptime/src/health_checker.rs +++ b/easytier-contrib/easytier-uptime/src/health_checker.rs @@ -10,9 +10,8 @@ use easytier::{ common::config::{ ConfigFileControl, ConfigLoader, NetworkIdentity, PeerConfig, TomlConfigLoader, }, - instance_manager::NetworkInstanceManager, + instance::factory::{NativeInstanceManager, native_instance_manager}, }; -use guarden::defer; use serde::{Deserialize, Serialize}; use sqlx::any; use tokio_util::task::AbortOnDropHandle; @@ -28,6 +27,32 @@ pub struct HealthCheckOneNode { node_id: String, } +struct InstanceCleanupGuard { + manager: Arc, + instance_id: Option, + runtime: tokio::runtime::Handle, +} + +impl InstanceCleanupGuard { + async fn cleanup(mut self) { + let instance_id = self.instance_id.unwrap(); + let _ = self.manager.delete_network_instances([instance_id]).await; + self.instance_id = None; + } +} + +impl Drop for InstanceCleanupGuard { + fn drop(&mut self) { + let Some(instance_id) = self.instance_id.take() else { + return; + }; + let manager = self.manager.clone(); + self.runtime.spawn(async move { + let _ = manager.delete_network_instances([instance_id]).await; + }); + } +} + const HEALTH_CHECK_RING_GRANULARITY_SEC: usize = 60 * 15; // 15分钟 const HEALTH_CHECK_RING_MAX_DURATION_SEC: usize = 60 * 60 * 24; // 最多一天 @@ -238,7 +263,7 @@ impl HealthyMemRecord { pub struct HealthChecker { db: Db, - instance_mgr: Arc, + instance_mgr: Arc, inst_id_map: DashMap, node_tasks: DashMap>, node_records: Arc>, @@ -247,7 +272,7 @@ pub struct HealthChecker { impl HealthChecker { pub fn new(db: Db) -> Self { - let instance_mgr = Arc::new(NetworkInstanceManager::new()); + let instance_mgr = Arc::new(native_instance_manager()); Self { db, instance_mgr, @@ -387,33 +412,38 @@ impl HealthChecker { max_time: Duration, ) -> anyhow::Result<()> { let cfg = self.get_node_cfg_with_model(node_info, None).await?; - defer!({ - let _ = self - .instance_mgr - .delete_network_instance(vec![cfg.get_id()]); - }); self.instance_mgr - .run_network_instance(cfg.clone(), false, ConfigFileControl::STATIC_CONFIG) + .run_network_instance(cfg.clone(), ConfigFileControl::STATIC_CONFIG) .with_context(|| "failed to run network instance")?; + let cleanup = InstanceCleanupGuard { + manager: self.instance_mgr.clone(), + instance_id: Some(cfg.get_id()), + runtime: tokio::runtime::Handle::current(), + }; - let now = Instant::now(); - let mut err = None; - while now.elapsed() < max_time { - match Self::test_node_healthy(cfg.get_id(), self.instance_mgr.clone()).await { - Ok(_) => { - return Ok(()); - } - Err(e) => { - warn!( - "test node healthy failed, node_info: {:?}, err: {}", - node_info, e - ); - err = Some(e); + let result = async { + let now = Instant::now(); + let mut err = None; + while now.elapsed() < max_time { + match Self::test_node_healthy(cfg.get_id(), self.instance_mgr.clone()).await { + Ok(_) => { + return Ok(()); + } + Err(e) => { + warn!( + "test node healthy failed, node_info: {:?}, err: {}", + node_info, e + ); + err = Some(e); + } } + tokio::time::sleep(Duration::from_millis(100)).await; } - tokio::time::sleep(Duration::from_millis(100)).await; + Err(anyhow::anyhow!("test node healthy failed, err: {:?}", err)) } - Err(anyhow::anyhow!("test node healthy failed, err: {:?}", err)) + .await; + cleanup.cleanup().await; + result } async fn get_node_cfg( @@ -437,7 +467,7 @@ impl HealthChecker { ); self.instance_mgr - .run_network_instance(cfg.clone(), true, ConfigFileControl::STATIC_CONFIG) + .run_network_instance(cfg.clone(), ConfigFileControl::STATIC_CONFIG) .with_context(|| "failed to run network instance")?; self.inst_id_map.insert(node_id, cfg.get_id()); @@ -481,7 +511,10 @@ impl HealthChecker { pub async fn remove_node(&self, node_id: i32) -> anyhow::Result<()> { self.node_tasks.remove(&node_id); if let Some(inst_id) = self.inst_id_map.remove(&node_id) { - let _ = self.instance_mgr.delete_network_instance(vec![inst_id.1]); + let _ = self + .instance_mgr + .delete_network_instances([inst_id.1]) + .await; } self.node_cfg.remove(&node_id); // 保留内存记录,不删除,以便后续查询历史数据 @@ -495,10 +528,10 @@ impl HealthChecker { #[instrument(err, ret, skip(instance_mgr))] async fn test_node_healthy( inst_id: uuid::Uuid, - instance_mgr: Arc, + instance_mgr: Arc, // return version, response time on healthy, conn_count ) -> anyhow::Result<(String, u64, u32)> { - let Some(instance) = instance_mgr.get_network_info(&inst_id).await else { + let Some(instance) = instance_mgr.network_info(inst_id).await else { anyhow::bail!("healthy check node is not started"); }; @@ -566,7 +599,7 @@ impl HealthChecker { async fn node_health_check_task( node_id: i32, inst_id: uuid::Uuid, - instance_mgr: Arc, + instance_mgr: Arc, db: Db, node_records: Arc>, ) { diff --git a/easytier-core/Cargo.toml b/easytier-core/Cargo.toml new file mode 100644 index 00000000..6989d924 --- /dev/null +++ b/easytier-core/Cargo.toml @@ -0,0 +1,137 @@ +[package] +name = "easytier-core" +description = "EasyTier OS-free control-plane core primitives." +homepage = "https://github.com/EasyTier/EasyTier" +repository = "https://github.com/EasyTier/EasyTier" +version = "2.6.4" +edition.workspace = true +rust-version.workspace = true +authors = ["kkrainbow"] +keywords = ["vpn", "p2p", "network", "easytier"] +categories = ["network-programming"] +license-file = "../LICENSE" + +[lib] +crate-type = ["rlib", "cdylib"] + +[dependencies] +anyhow = "1.0" +ariadne = { version = "0.5", optional = true } +arc-swap = "1.7" +async-ringbuf = "0.3.1" +async-trait = "0.1.74" +auto_impl = "1.1.0" +base64 = "0.22" +bitflags = "2.5" +bytecodec = "0.4.15" +bytes = "1.5.0" +chrono = { version = "0.4.37", features = ["clock"] } +cidr = { version = "0.3.1", features = ["serde"] } +crossbeam = "0.8.4" +dashmap = "6.0" +derive_builder = "0.20.2" +easytier-proto = { path = "../easytier-proto", default-features = false, features = ["core"] } +futures = "0.3" +guarden = "0.2" +hmac = "0.12.1" +http-body-util = { version = "0.1", optional = true } +hyper = { version = "1", default-features = false, features = ["client", "http1"], optional = true } +hyper-util = { version = "0.1", default-features = false, features = ["tokio"], optional = true } +idna = "1.0" +atomic-shim = "0.2.0" +ordered_hash_map = "0.5.0" +parking_lot = "0.12.1" +percent-encoding = "2.3.1" +petgraph = "0.8.1" +pin-project-lite = "0.2.13" +pnet_packet = { version = "0.35.0", optional = true } +prefix-trie = { version = "0.7.0", features = ["cidr"] } +prost = "0.14.3" +prost-types = "0.14.3" +rand = "0.8.5" +quanta = "0.12" +rustls = { version = "0.23", default-features = false, features = ["ring", "std", "tls12"], optional = true } +serde = { version = "1.0", features = ["derive"] } +serde_json = "1" +sha2 = "0.10.8" +smoltcp = { git = "https://github.com/smoltcp-rs/smoltcp.git", rev = "0a926767a68bc88d5512afefa7529c5ecdade4ea", optional = true, default-features = false } +snow = "0.10.0" +stun_codec = "0.3.4" +thiserror = "1.0" +tracing = "0.1" +strum = { version = "0.27.2", features = ["derive"] } +toml = "0.8.12" +tokio = { version = "1", default-features = false, features = [ + "rt", + "time", + "sync", + "macros", + "io-util", +] } +tokio-util = { version = "0.7", features = ["io", "rt"] } +tokio-rustls = { version = "0.26", default-features = false, optional = true } +url = { version = "2.5", features = ["serde"] } +wildmatch = "2.3.4" +uuid = { version = "1.5.0", features = ["v4", "fast-rng", "serde"] } +webpki-roots = { version = "0.26", optional = true } +x25519-dalek = { version = "2.0", features = ["static_secrets"] } +zerocopy = { version = "0.7.32", features = ["derive", "simd"] } +zstd = { version = "0.13", optional = true } +aes-gcm = { version = "0.10.3", optional = true } +chacha20poly1305 = { version = "0.10.1", optional = true } + +[features] +default = ["aes-gcm", "endpoint-discovery", "extended-services", "management", "tcp-hole-punch"] +aes-gcm = ["dep:aes-gcm"] +chacha20 = ["dep:chacha20poly1305"] +config-write = [] +endpoint-discovery = [ + "dep:http-body-util", + "dep:hyper", + "dep:hyper-util", + "dep:rustls", + "dep:tokio-rustls", + "dep:webpki-roots", +] +dhcp-ipv4 = [] +public-ipv6-provider = [] +vpn-portal = [] +wrapped-transport = [] +extended-services = [ + "dhcp-ipv4", + "public-ipv6-provider", + "vpn-portal", + "wrapped-transport", + "proxy-cidr-monitor", +] +management = ["management-rpc", "config-write", "extended-services", "rich-config-errors", "easytier-proto/json-rpc"] +management-rpc = ["easytier-proto/api"] +proxy-cidr-monitor = [] +rich-config-errors = ["dep:ariadne"] +tcp-hole-punch = [] +proxy-packet = [ + "wrapped-transport", + "dep:pnet_packet", + "dep:smoltcp", + "smoltcp/std", + "smoltcp/proto-ipv4", + "smoltcp/proto-ipv4-fragmentation", + "smoltcp/fragmentation-buffer-size-8192", + "smoltcp/assembler-max-segment-count-16", + "smoltcp/reassembly-buffer-size-8192", + "smoltcp/reassembly-buffer-count-16", +] +proxy-smoltcp-stack = [ + "proxy-packet", + "smoltcp/medium-ip", + "smoltcp/socket-tcp", + "smoltcp/socket-udp", + "smoltcp/proto-ipv6", + "smoltcp/async", +] +test-utils = [] +tracing-log = ["tracing/log"] +zstd = ["dep:zstd"] + +[target.'cfg(not(target_os = "wasi"))'.dev-dependencies] +tokio = { version = "1", default-features = false, features = ["rt-multi-thread"] } diff --git a/easytier-core/src/config/api.rs b/easytier-core/src/config/api.rs new file mode 100644 index 00000000..515b3af1 --- /dev/null +++ b/easytier-core/src/config/api.rs @@ -0,0 +1,166 @@ +//! Portable conversion between the shared TOML model and management schema. + +use easytier_proto::api::manage::{ + self, NetworkConfig, NetworkingMethod, PortForwardConfig as ApiPortForwardConfig, +}; + +use super::toml::{ConfigLoader as _, TomlConfig}; + +pub fn network_config_from_toml(config: &TomlConfig) -> NetworkConfig { + let default_config = TomlConfig::default(); + let mut result = NetworkConfig { + instance_id: Some(config.get_id().to_string()), + dhcp: Some(config.get_dhcp()), + ..Default::default() + }; + + if config.get_hostname() != default_config.get_hostname() { + result.hostname = Some(config.get_hostname()); + } + + let network_identity = config.get_network_identity(); + result.network_name = Some(network_identity.network_name); + result.network_secret = network_identity.network_secret; + + if let Some(ipv4) = config.get_ipv4() { + result.virtual_ipv4 = Some(ipv4.address().to_string()); + result.network_length = Some(ipv4.network_length() as i32); + } + + if config.get_ipv6_public_addr_provider() != default_config.get_ipv6_public_addr_provider() { + result.ipv6_public_addr_provider = Some(config.get_ipv6_public_addr_provider()); + } + if config.get_ipv6_public_addr_auto() != default_config.get_ipv6_public_addr_auto() { + result.ipv6_public_addr_auto = Some(config.get_ipv6_public_addr_auto()); + } + result.ipv6_public_addr_prefix = config + .get_ipv6_public_addr_prefix() + .map(|prefix| prefix.to_string()); + + let peers = config.get_peers(); + result.networking_method = Some(NetworkingMethod::Manual as i32); + if !peers.is_empty() { + result.peer_urls = peers.iter().map(|peer| peer.uri.to_string()).collect(); + result.peers = peers + .iter() + .map(|peer| manage::NetworkPeerConfig { + uri: peer.uri.to_string(), + peer_public_key: peer.peer_public_key.clone(), + }) + .collect(); + } + + result.listener_urls = config + .get_listeners() + .unwrap_or_default() + .iter() + .map(ToString::to_string) + .collect(); + result.proxy_cidrs = config + .get_proxy_cidrs() + .iter() + .map(|proxy| match proxy.mapped_cidr { + Some(mapped) => format!("{}->{}", proxy.cidr, mapped), + None => proxy.cidr.to_string(), + }) + .collect(); + + let port_forwards = config.get_port_forwards(); + if !port_forwards.is_empty() { + result.port_forwards = port_forwards + .iter() + .map(|forward| ApiPortForwardConfig { + proto: forward.proto.clone(), + bind_ip: forward.bind_addr.ip().to_string(), + bind_port: forward.bind_addr.port() as u32, + dst_ip: forward.dst_addr.ip().to_string(), + dst_port: forward.dst_addr.port() as u32, + }) + .collect(); + } + + if let Some(vpn_config) = config.get_vpn_portal_config() { + result.enable_vpn_portal = Some(true); + result.vpn_portal_client_network_addr = + Some(vpn_config.client_cidr.first_address().to_string()); + result.vpn_portal_client_network_len = Some(vpn_config.client_cidr.network_length() as i32); + result.vpn_portal_listen_port = Some(vpn_config.wireguard_listen.port() as i32); + } + + if let Some(routes) = config.get_routes() + && !routes.is_empty() + { + result.enable_manual_routes = Some(true); + result.routes = routes.iter().map(ToString::to_string).collect(); + } + let exit_nodes = config.get_exit_nodes(); + if !exit_nodes.is_empty() { + result.exit_nodes = exit_nodes.iter().map(ToString::to_string).collect(); + } + if let Some(socks5_portal) = config.get_socks5_portal() { + result.enable_socks5 = Some(true); + result.socks5_port = socks5_portal.port().map(|port| port as i32); + } + let mapped_listeners = config.get_mapped_listeners(); + if !mapped_listeners.is_empty() { + result.mapped_listeners = mapped_listeners.iter().map(ToString::to_string).collect(); + } + + result.secure_mode = config.get_secure_mode(); + result.credential_file = config + .get_credential_file() + .map(|path| path.to_string_lossy().into_owned()); + + let flags = config.get_flags(); + let default_flags = default_config.get_flags(); + result.latency_first = Some(flags.latency_first); + result.dev_name = Some(flags.dev_name.clone()); + result.use_smoltcp = Some(flags.use_smoltcp); + result.disable_ipv6 = Some(!flags.enable_ipv6); + result.enable_kcp_proxy = Some(flags.enable_kcp_proxy); + result.disable_kcp_input = Some(flags.disable_kcp_input); + result.enable_quic_proxy = Some(flags.enable_quic_proxy); + result.disable_quic_input = Some(flags.disable_quic_input); + result.disable_p2p = Some(flags.disable_p2p); + result.p2p_only = Some(flags.p2p_only); + result.lazy_p2p = Some(flags.lazy_p2p); + result.bind_device = Some(flags.bind_device); + result.socket_mark = flags.socket_mark; + result.no_tun = Some(flags.no_tun); + result.enable_exit_node = Some(flags.enable_exit_node); + result.relay_all_peer_rpc = Some(flags.relay_all_peer_rpc); + result.need_p2p = Some(flags.need_p2p); + result.multi_thread = Some(flags.multi_thread); + result.proxy_forward_by_system = Some(flags.proxy_forward_by_system); + result.disable_encryption = Some(!flags.enable_encryption); + result.disable_tcp_hole_punching = Some(flags.disable_tcp_hole_punching); + result.disable_udp_hole_punching = Some(flags.disable_udp_hole_punching); + result.disable_upnp = Some(flags.disable_upnp); + result.disable_relay_data = Some(flags.disable_relay_data); + result.enable_udp_broadcast_relay = Some(flags.enable_udp_broadcast_relay); + result.disable_sym_hole_punching = Some(flags.disable_sym_hole_punching); + result.enable_magic_dns = Some(flags.accept_dns); + result.mtu = Some(flags.mtu as i32); + result.data_compress_algo = (flags.data_compress_algo != default_flags.data_compress_algo) + .then_some(flags.data_compress_algo); + result.encryption_algorithm = (flags.encryption_algorithm + != default_flags.encryption_algorithm) + .then_some(flags.encryption_algorithm); + result.instance_recv_bps_limit = + (flags.instance_recv_bps_limit != u64::MAX).then_some(flags.instance_recv_bps_limit); + result.enable_private_mode = Some(flags.private_mode); + result.acl = config.get_acl(); + + if flags.relay_network_whitelist == "*" { + result.enable_relay_network_whitelist = Some(false); + } else { + result.enable_relay_network_whitelist = Some(true); + result.relay_network_whitelist = flags + .relay_network_whitelist + .split_whitespace() + .map(ToOwned::to_owned) + .collect(); + } + + result +} diff --git a/easytier-core/src/config/api_input.rs b/easytier-core/src/config/api_input.rs new file mode 100644 index 00000000..281985cc --- /dev/null +++ b/easytier-core/src/config/api_input.rs @@ -0,0 +1,652 @@ +//! Conversion between the management NetworkConfig schema and shared TOML. + +use std::net::SocketAddr; + +use anyhow::Context; +use easytier_proto::api::manage; + +use crate::config::{ + MappedListenerPolicy, normalize_secure_mode_config, + toml::{ + ConfigLoader, NetworkIdentity, PeerConfig, PortForwardConfig, TomlConfigLoader, + VpnPortalConfig, gen_default_flags, + }, +}; + +fn parse_mapped_listener_urls(mapped_listeners: &[String]) -> Result, anyhow::Error> { + MappedListenerPolicy::new(["tcp", "udp", "wg", "quic", "ws", "wss", "faketcp"]) + .parse_urls(mapped_listeners) +} + +pub fn add_proxy_network_to_config( + proxy_network: &str, + cfg: &TomlConfigLoader, +) -> Result<(), anyhow::Error> { + let parts: Vec<&str> = proxy_network.split("->").collect(); + let real_cidr = parts[0] + .parse() + .with_context(|| format!("failed to parse proxy network: {}", parts[0]))?; + + if parts.len() > 2 { + return Err(anyhow::anyhow!( + "invalid proxy network format: {}, support format: or ->, example: + 10.0.0.0/24 or 10.0.0.0/24->192.168.0.0/24", + proxy_network + )); + } + + let mapped_cidr = if parts.len() == 2 { + Some( + parts[1] + .parse() + .with_context(|| format!("failed to parse mapped network: {}", parts[1]))?, + ) + } else { + None + }; + cfg.add_proxy_cidr(real_cidr, mapped_cidr)?; + Ok(()) +} + +pub type NetworkingMethod = easytier_proto::api::manage::NetworkingMethod; +pub type NetworkConfig = easytier_proto::api::manage::NetworkConfig; + +pub trait NetworkConfigExt { + fn gen_config(&self) -> Result; + fn new_from_config(config: impl ConfigLoader) -> Result; +} + +fn parse_peer(peer: &manage::NetworkPeerConfig) -> Result, anyhow::Error> { + let uri = peer.uri.trim(); + if uri.is_empty() { + return Ok(None); + } + + Ok(Some(PeerConfig { + uri: uri + .parse() + .with_context(|| format!("failed to parse peer uri: {}", uri))?, + peer_public_key: peer.peer_public_key.clone(), + })) +} + +fn parse_peers(peers: &[manage::NetworkPeerConfig]) -> Result, anyhow::Error> { + let mut ret = Vec::new(); + for peer in peers { + if let Some(peer) = parse_peer(peer)? { + ret.push(peer); + } + } + Ok(ret) +} + +fn parse_peer_urls(peer_urls: &[String]) -> Result, anyhow::Error> { + let mut peers = vec![]; + for peer_url in peer_urls.iter() { + let peer_url = peer_url.trim(); + if peer_url.is_empty() { + continue; + } + peers.push(PeerConfig { + uri: peer_url + .parse() + .with_context(|| format!("failed to parse peer uri: {}", peer_url))?, + peer_public_key: None, + }); + } + Ok(peers) +} + +impl NetworkConfigExt for NetworkConfig { + fn gen_config(&self) -> Result { + let cfg = TomlConfigLoader::default(); + cfg.set_id( + self.instance_id + .clone() + .unwrap_or(uuid::Uuid::new_v4().to_string()) + .parse() + .with_context(|| format!("failed to parse instance id: {:?}", self.instance_id))?, + ); + cfg.set_hostname(self.hostname.clone()); + cfg.set_dhcp(self.dhcp.unwrap_or_default()); + cfg.set_inst_name(self.network_name.clone().unwrap_or_default()); + + // The web UI does not expose credential inputs directly, but imported/saved + // NetworkConfig objects still need to preserve credential-mode instances via + // secure_mode.local_private_key + empty network_secret. + let credential_secret = if self.network_secret.is_some() { + None + } else { + self.secure_mode + .as_ref() + .and_then(|mode| mode.local_private_key.clone()) + .filter(|s| !s.is_empty()) + }; + + if credential_secret.is_some() { + cfg.set_network_identity(NetworkIdentity::new_credential( + self.network_name.clone().unwrap_or_default(), + )); + } else { + cfg.set_network_identity(NetworkIdentity::new( + self.network_name.clone().unwrap_or_default(), + self.network_secret.clone().unwrap_or_default(), + )); + } + + if !cfg.get_dhcp() { + let virtual_ipv4 = self.virtual_ipv4.clone().unwrap_or_default(); + if !virtual_ipv4.is_empty() { + let ip = format!("{}/{}", virtual_ipv4, self.network_length.unwrap_or(24)) + .parse() + .with_context(|| { + format!( + "failed to parse ipv4 inet address: {}, {:?}", + virtual_ipv4, self.network_length + ) + })?; + cfg.set_ipv4(Some(ip)); + } + } + + match NetworkingMethod::try_from(self.networking_method.unwrap_or_default()) + .unwrap_or_default() + { + NetworkingMethod::PublicServer => { + let peers = parse_peers(&self.peers)?; + if peers.is_empty() { + let public_server_url = self.public_server_url.clone().unwrap_or_default(); + cfg.set_peers(vec![PeerConfig { + uri: public_server_url.parse().with_context(|| { + format!("failed to parse public server uri: {}", public_server_url) + })?, + peer_public_key: None, + }]); + } else { + cfg.set_peers(peers); + } + } + NetworkingMethod::Manual => { + let mut peers = parse_peers(&self.peers)?; + if peers.is_empty() { + peers = parse_peer_urls(&self.peer_urls)?; + } + if !peers.is_empty() { + cfg.set_peers(peers); + } + } + NetworkingMethod::Standalone => {} + } + + let mut listener_urls = vec![]; + for listener_url in self.listener_urls.iter() { + if listener_url.is_empty() { + continue; + } + listener_urls.push( + listener_url + .parse() + .with_context(|| format!("failed to parse listener uri: {}", listener_url))?, + ); + } + cfg.set_listeners(listener_urls); + + for n in self.proxy_cidrs.iter() { + add_proxy_network_to_config(n, &cfg)?; + } + + if !self.port_forwards.is_empty() { + cfg.set_port_forwards( + self.port_forwards + .iter() + .filter(|pf| !pf.bind_ip.is_empty() && !pf.dst_ip.is_empty()) + .filter_map(|pf| { + let bind_addr = + format!("{}:{}", pf.bind_ip, pf.bind_port).parse::(); + let dst_addr = + format!("{}:{}", pf.dst_ip, pf.dst_port).parse::(); + + match (bind_addr, dst_addr) { + (Ok(bind_addr), Ok(dst_addr)) => Some(PortForwardConfig { + bind_addr, + dst_addr, + proto: pf.proto.clone(), + }), + _ => None, + } + }) + .collect::>(), + ); + } + + if self.enable_vpn_portal.unwrap_or_default() { + let cidr = format!( + "{}/{}", + self.vpn_portal_client_network_addr + .clone() + .unwrap_or_default(), + self.vpn_portal_client_network_len.unwrap_or(24) + ); + cfg.set_vpn_portal_config(VpnPortalConfig { + client_cidr: cidr + .parse() + .with_context(|| format!("failed to parse vpn portal client cidr: {}", cidr))?, + wireguard_listen: format!( + "0.0.0.0:{}", + self.vpn_portal_listen_port.unwrap_or_default() + ) + .parse() + .with_context(|| { + format!( + "failed to parse vpn portal wireguard listen port. {:?}", + self.vpn_portal_listen_port + ) + })?, + }); + } + + if self.enable_manual_routes.unwrap_or_default() { + let mut routes = Vec::::with_capacity(self.routes.len()); + for route in self.routes.iter() { + routes.push( + route + .parse() + .with_context(|| format!("failed to parse route: {}", route))?, + ); + } + cfg.set_routes(Some(routes)); + } + + if !self.exit_nodes.is_empty() { + let mut exit_nodes = Vec::::with_capacity(self.exit_nodes.len()); + for node in self.exit_nodes.iter() { + exit_nodes.push( + node.parse() + .with_context(|| format!("failed to parse exit node: {}", node))?, + ); + } + cfg.set_exit_nodes(exit_nodes); + } + + if self.enable_socks5.unwrap_or_default() + && let Some(socks5_port) = self.socks5_port + { + cfg.set_socks5_portal(Some( + format!("socks5://0.0.0.0:{}", socks5_port).parse().unwrap(), + )); + } + + if !self.mapped_listeners.is_empty() { + let mapped_listeners = parse_mapped_listener_urls(&self.mapped_listeners)?; + cfg.set_mapped_listeners(Some(mapped_listeners)); + } + + if let Some(credential_file) = self + .credential_file + .as_ref() + .filter(|path| !path.is_empty()) + { + cfg.set_credential_file(Some(credential_file.into())); + } + + if let Some(credential_secret) = credential_secret { + cfg.set_secure_mode(Some(normalize_secure_mode_config( + easytier_proto::common::SecureModeConfig { + enabled: true, + local_private_key: Some(credential_secret), + local_public_key: None, + }, + )?)); + } else { + cfg.set_secure_mode( + self.secure_mode + .clone() + .map(normalize_secure_mode_config) + .transpose()?, + ); + } + + let mut flags = gen_default_flags(); + if let Some(latency_first) = self.latency_first { + flags.latency_first = latency_first; + } + + if let Some(dev_name) = self.dev_name.clone() { + flags.dev_name = dev_name; + } + + if let Some(use_smoltcp) = self.use_smoltcp { + flags.use_smoltcp = use_smoltcp; + } + + if let Some(ipv6_public_addr_provider) = self.ipv6_public_addr_provider { + cfg.set_ipv6_public_addr_provider(ipv6_public_addr_provider); + } + + if let Some(ipv6_public_addr_auto) = self.ipv6_public_addr_auto { + cfg.set_ipv6_public_addr_auto(ipv6_public_addr_auto); + } + + if let Some(ipv6_public_addr_prefix) = self + .ipv6_public_addr_prefix + .as_ref() + .filter(|prefix| !prefix.is_empty()) + { + cfg.set_ipv6_public_addr_prefix(Some(ipv6_public_addr_prefix.parse().with_context( + || format!("failed to parse ipv6 public address prefix: {ipv6_public_addr_prefix}"), + )?)); + } + + if let Some(disable_ipv6) = self.disable_ipv6 { + flags.enable_ipv6 = !disable_ipv6; + } + + if let Some(enable_kcp_proxy) = self.enable_kcp_proxy { + flags.enable_kcp_proxy = enable_kcp_proxy; + } + + if let Some(disable_kcp_input) = self.disable_kcp_input { + flags.disable_kcp_input = disable_kcp_input; + } + + if let Some(enable_quic_proxy) = self.enable_quic_proxy { + flags.enable_quic_proxy = enable_quic_proxy; + } + + if let Some(disable_quic_input) = self.disable_quic_input { + flags.disable_quic_input = disable_quic_input; + } + + if let Some(disable_p2p) = self.disable_p2p { + flags.disable_p2p = disable_p2p; + } + + if let Some(p2p_only) = self.p2p_only { + flags.p2p_only = p2p_only; + } + + if let Some(lazy_p2p) = self.lazy_p2p { + flags.lazy_p2p = lazy_p2p; + } + + if let Some(bind_device) = self.bind_device { + flags.bind_device = bind_device; + } + + if self.socket_mark.is_some() { + flags.socket_mark = self.socket_mark; + } + + if let Some(no_tun) = self.no_tun { + flags.no_tun = no_tun; + } + + if let Some(enable_exit_node) = self.enable_exit_node { + flags.enable_exit_node = enable_exit_node; + } + + if let Some(relay_all_peer_rpc) = self.relay_all_peer_rpc { + flags.relay_all_peer_rpc = relay_all_peer_rpc; + } + + if let Some(need_p2p) = self.need_p2p { + flags.need_p2p = need_p2p; + } + + if let Some(multi_thread) = self.multi_thread { + flags.multi_thread = multi_thread; + } + + if let Some(proxy_forward_by_system) = self.proxy_forward_by_system { + flags.proxy_forward_by_system = proxy_forward_by_system; + } + + if let Some(disable_encryption) = self.disable_encryption { + flags.enable_encryption = !disable_encryption; + } + + if self.enable_relay_network_whitelist.unwrap_or_default() { + if !self.relay_network_whitelist.is_empty() { + flags.relay_network_whitelist = self.relay_network_whitelist.join(" "); + } else { + flags.relay_network_whitelist = "".to_string(); + } + } + + if let Some(disable_tcp_hole_punching) = self.disable_tcp_hole_punching { + flags.disable_tcp_hole_punching = disable_tcp_hole_punching; + } + + if let Some(disable_udp_hole_punching) = self.disable_udp_hole_punching { + flags.disable_udp_hole_punching = disable_udp_hole_punching; + } + + if let Some(disable_upnp) = self.disable_upnp { + flags.disable_upnp = disable_upnp; + } + + if let Some(disable_relay_data) = self.disable_relay_data { + flags.disable_relay_data = disable_relay_data; + } + + if let Some(enable_udp_broadcast_relay) = self.enable_udp_broadcast_relay { + flags.enable_udp_broadcast_relay = enable_udp_broadcast_relay; + } + + if let Some(disable_sym_hole_punching) = self.disable_sym_hole_punching { + flags.disable_sym_hole_punching = disable_sym_hole_punching; + } + + if let Some(enable_magic_dns) = self.enable_magic_dns { + flags.accept_dns = enable_magic_dns; + } + + if let Some(mtu) = self.mtu { + flags.mtu = mtu as u32; + } + + if let Some(instance_recv_bps_limit) = self.instance_recv_bps_limit { + flags.instance_recv_bps_limit = instance_recv_bps_limit; + } + + if let Some(enable_private_mode) = self.enable_private_mode { + flags.private_mode = enable_private_mode; + } + + if let Some(encryption_algorithm) = self.encryption_algorithm.clone() { + flags.encryption_algorithm = encryption_algorithm; + } + + if let Some(acl) = self.acl.as_ref() + && !acl.is_empty() + { + cfg.set_acl(Some(acl.clone())); + } + + if let Some(data_compress_algo) = self.data_compress_algo { + if data_compress_algo < 1 { + flags.data_compress_algo = 1; + } else { + flags.data_compress_algo = data_compress_algo + } + } + + cfg.set_flags(flags); + Ok(cfg) + } + + fn new_from_config(config: impl ConfigLoader) -> Result { + let default_config = TomlConfigLoader::default(); + + let mut result = Self { + ..Default::default() + }; + + result.instance_id = Some(config.get_id().to_string()); + if config.get_hostname() != default_config.get_hostname() { + result.hostname = Some(config.get_hostname()); + } + + result.dhcp = Some(config.get_dhcp()); + + let network_identity = config.get_network_identity(); + result.network_name = Some(network_identity.network_name.clone()); + result.network_secret = network_identity.network_secret; + + if let Some(ipv4) = config.get_ipv4() { + result.virtual_ipv4 = Some(ipv4.address().to_string()); + result.network_length = Some(ipv4.network_length() as i32); + } + + if config.get_ipv6_public_addr_provider() != default_config.get_ipv6_public_addr_provider() + { + result.ipv6_public_addr_provider = Some(config.get_ipv6_public_addr_provider()); + } + if config.get_ipv6_public_addr_auto() != default_config.get_ipv6_public_addr_auto() { + result.ipv6_public_addr_auto = Some(config.get_ipv6_public_addr_auto()); + } + result.ipv6_public_addr_prefix = config + .get_ipv6_public_addr_prefix() + .map(|prefix| prefix.to_string()); + + let peers = config.get_peers(); + result.networking_method = Some(NetworkingMethod::Manual as i32); + if !peers.is_empty() { + result.peer_urls = peers.iter().map(|p| p.uri.to_string()).collect(); + result.peers = peers + .iter() + .map(|p| manage::NetworkPeerConfig { + uri: p.uri.to_string(), + peer_public_key: p.peer_public_key.clone(), + }) + .collect(); + } + + result.listener_urls = config + .get_listeners() + .unwrap_or_default() + .iter() + .map(|l| l.to_string()) + .collect(); + + result.proxy_cidrs = config + .get_proxy_cidrs() + .iter() + .map(|c| { + if let Some(mapped) = c.mapped_cidr { + format!("{}->{}", c.cidr, mapped) + } else { + c.cidr.to_string() + } + }) + .collect(); + + let port_forwards = config.get_port_forwards(); + if !port_forwards.is_empty() { + result.port_forwards = port_forwards + .iter() + .map(|f| manage::PortForwardConfig { + proto: f.proto.clone(), + bind_ip: f.bind_addr.ip().to_string(), + bind_port: f.bind_addr.port() as u32, + dst_ip: f.dst_addr.ip().to_string(), + dst_port: f.dst_addr.port() as u32, + }) + .collect(); + } + + if let Some(vpn_config) = config.get_vpn_portal_config() { + result.enable_vpn_portal = Some(true); + + let cidr = vpn_config.client_cidr; + result.vpn_portal_client_network_addr = Some(cidr.first_address().to_string()); + result.vpn_portal_client_network_len = Some(cidr.network_length() as i32); + + result.vpn_portal_listen_port = Some(vpn_config.wireguard_listen.port() as i32); + } + + if let Some(routes) = config.get_routes() + && !routes.is_empty() + { + result.enable_manual_routes = Some(true); + result.routes = routes.iter().map(|r| r.to_string()).collect(); + } + + let exit_nodes = config.get_exit_nodes(); + if !exit_nodes.is_empty() { + result.exit_nodes = exit_nodes.iter().map(|n| n.to_string()).collect(); + } + + if let Some(socks5_portal) = config.get_socks5_portal() { + result.enable_socks5 = Some(true); + result.socks5_port = socks5_portal.port().map(|p| p as i32); + } + + let mapped_listeners = config.get_mapped_listeners(); + if !mapped_listeners.is_empty() { + result.mapped_listeners = mapped_listeners.iter().map(|l| l.to_string()).collect(); + } + + result.secure_mode = config.get_secure_mode(); + result.credential_file = config + .get_credential_file() + .map(|path| path.to_string_lossy().into_owned()); + let flags = config.get_flags(); + let default_flags = default_config.get_flags(); + result.latency_first = Some(flags.latency_first); + result.dev_name = Some(flags.dev_name.clone()); + result.use_smoltcp = Some(flags.use_smoltcp); + result.disable_ipv6 = Some(!flags.enable_ipv6); + result.enable_kcp_proxy = Some(flags.enable_kcp_proxy); + result.disable_kcp_input = Some(flags.disable_kcp_input); + result.enable_quic_proxy = Some(flags.enable_quic_proxy); + result.disable_quic_input = Some(flags.disable_quic_input); + result.disable_p2p = Some(flags.disable_p2p); + result.p2p_only = Some(flags.p2p_only); + result.lazy_p2p = Some(flags.lazy_p2p); + result.bind_device = Some(flags.bind_device); + result.socket_mark = flags.socket_mark; + result.no_tun = Some(flags.no_tun); + result.enable_exit_node = Some(flags.enable_exit_node); + result.relay_all_peer_rpc = Some(flags.relay_all_peer_rpc); + result.need_p2p = Some(flags.need_p2p); + result.multi_thread = Some(flags.multi_thread); + result.proxy_forward_by_system = Some(flags.proxy_forward_by_system); + result.disable_encryption = Some(!flags.enable_encryption); + result.disable_tcp_hole_punching = Some(flags.disable_tcp_hole_punching); + result.disable_udp_hole_punching = Some(flags.disable_udp_hole_punching); + result.disable_upnp = Some(flags.disable_upnp); + result.disable_relay_data = Some(flags.disable_relay_data); + result.enable_udp_broadcast_relay = Some(flags.enable_udp_broadcast_relay); + result.disable_sym_hole_punching = Some(flags.disable_sym_hole_punching); + result.enable_magic_dns = Some(flags.accept_dns); + result.mtu = Some(flags.mtu as i32); + result.data_compress_algo = (flags.data_compress_algo != default_flags.data_compress_algo) + .then_some(flags.data_compress_algo); + result.encryption_algorithm = (flags.encryption_algorithm + != default_flags.encryption_algorithm) + .then_some(flags.encryption_algorithm.clone()); + result.instance_recv_bps_limit = + (flags.instance_recv_bps_limit != u64::MAX).then_some(flags.instance_recv_bps_limit); + result.enable_private_mode = Some(flags.private_mode); + + result.acl = config.get_acl(); + + if flags.relay_network_whitelist == "*" { + result.enable_relay_network_whitelist = Some(false); + } else { + result.enable_relay_network_whitelist = Some(true); + if flags.relay_network_whitelist.is_empty() { + result.relay_network_whitelist = vec![]; + } else { + result.relay_network_whitelist = flags + .relay_network_whitelist + .split_whitespace() + .map(|s| s.to_string()) + .collect(); + } + } + + Ok(result) + } +} diff --git a/easytier-core/src/config/encryption.rs b/easytier-core/src/config/encryption.rs new file mode 100644 index 00000000..f865a03c --- /dev/null +++ b/easytier-core/src/config/encryption.rs @@ -0,0 +1,73 @@ +use std::{fmt, str::FromStr}; + +use strum::VariantArray; + +/// Stable configuration vocabulary for every known encryption algorithm. +#[derive(Debug, Default, Clone, Copy, PartialEq, Eq, VariantArray)] +pub enum EncryptionAlgorithm { + Xor, + #[default] + AesGcm, + Aes256Gcm, + ChaCha20, +} + +impl EncryptionAlgorithm { + pub const fn as_str(self) -> &'static str { + match self { + Self::Xor => "xor", + Self::AesGcm => "aes-gcm", + Self::Aes256Gcm => "aes-256-gcm", + Self::ChaCha20 => "chacha20", + } + } +} + +impl fmt::Display for EncryptionAlgorithm { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter.write_str(self.as_str()) + } +} + +impl FromStr for EncryptionAlgorithm { + type Err = (); + + fn from_str(value: &str) -> Result { + match value.to_ascii_lowercase().as_str() { + "xor" => Ok(Self::Xor), + "aes-gcm" | "openssl-aes-gcm" => Ok(Self::AesGcm), + "aes-256-gcm" | "openssl-aes-256-gcm" => Ok(Self::Aes256Gcm), + "chacha20" | "chacha20-poly1305" | "openssl-chacha20" => Ok(Self::ChaCha20), + _ => Err(()), + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn known_algorithm_names_are_stable() { + let cases = [ + ("xor", EncryptionAlgorithm::Xor), + ("aes-gcm", EncryptionAlgorithm::AesGcm), + ("aes-256-gcm", EncryptionAlgorithm::Aes256Gcm), + ("chacha20", EncryptionAlgorithm::ChaCha20), + ("chacha20-poly1305", EncryptionAlgorithm::ChaCha20), + ("openssl-aes-gcm", EncryptionAlgorithm::AesGcm), + ("openssl-aes-256-gcm", EncryptionAlgorithm::Aes256Gcm), + ("openssl-chacha20", EncryptionAlgorithm::ChaCha20), + ]; + + for (name, expected) in cases { + assert_eq!(name.parse(), Ok(expected)); + } + assert_eq!(EncryptionAlgorithm::ChaCha20.to_string(), "chacha20"); + } + + #[test] + fn aes_is_the_stable_default() { + assert_eq!(EncryptionAlgorithm::default(), EncryptionAlgorithm::AesGcm); + } +} diff --git a/easytier-core/src/config/gateway.rs b/easytier-core/src/config/gateway.rs new file mode 100644 index 00000000..abaac4d5 --- /dev/null +++ b/easytier-core/src/config/gateway.rs @@ -0,0 +1,138 @@ +use std::net::SocketAddr; + +use serde::{Deserialize, Serialize}; + +use easytier_proto::common::{PortForwardConfigPb, SocketType}; + +/// Runtime configuration for the core-owned SOCKS and port-forward gateway. +#[derive(Debug, Clone, Default, PartialEq, Eq, Serialize, Deserialize)] +pub struct GatewayRuntimeConfig { + pub socks5_bind: Option, + pub port_forwards: Vec, +} + +/// One TCP or UDP port-forward rule. +#[derive(Debug, Clone, PartialEq, Eq, Hash, Serialize, Deserialize)] +pub struct PortForwardConfig { + pub bind_addr: SocketAddr, + pub dst_addr: SocketAddr, + pub proto: String, +} + +impl From for PortForwardConfig { + fn from(config: PortForwardConfigPb) -> Self { + Self { + bind_addr: config.bind_addr.unwrap_or_default().into(), + dst_addr: config.dst_addr.unwrap_or_default().into(), + proto: match SocketType::try_from(config.socket_type) { + Ok(SocketType::Tcp) => "tcp".to_string(), + Ok(SocketType::Udp) => "udp".to_string(), + _ => "tcp".to_string(), + }, + } + } +} + +impl From for PortForwardConfigPb { + fn from(config: PortForwardConfig) -> Self { + Self { + bind_addr: Some(config.bind_addr.into()), + dst_addr: Some(config.dst_addr.into()), + socket_type: match config.proto.to_lowercase().as_str() { + "tcp" => SocketType::Tcp as i32, + "udp" => SocketType::Udp as i32, + _ => SocketType::Tcp as i32, + }, + } + } +} + +#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize, Deserialize)] +pub struct ProxyRuntimeConfig { + pub enable_exit_node: bool, + pub no_tun: bool, + pub forward_by_system: bool, + pub force_smoltcp: bool, + pub icmp_failure_is_fatal: bool, + pub udp_response_ipv4_mtu: usize, +} + +impl ProxyRuntimeConfig { + pub fn should_start(self, has_proxy_networks: bool) -> bool { + if !has_proxy_networks && !self.enable_exit_node && !self.no_tun { + return false; + } + + !self.forward_by_system || self.no_tun + } +} + +impl Default for ProxyRuntimeConfig { + fn default() -> Self { + Self { + enable_exit_node: false, + no_tun: false, + forward_by_system: false, + force_smoltcp: false, + icmp_failure_is_fatal: false, + udp_response_ipv4_mtu: 1280, + } + } +} + +#[cfg(test)] +mod tests { + use super::ProxyRuntimeConfig; + + #[test] + fn proxy_startup_policy_preserves_runtime_modes() { + assert!(!ProxyRuntimeConfig::default().should_start(false)); + assert!(ProxyRuntimeConfig::default().should_start(true)); + assert!( + ProxyRuntimeConfig { + enable_exit_node: true, + ..Default::default() + } + .should_start(false) + ); + assert!( + ProxyRuntimeConfig { + no_tun: true, + ..Default::default() + } + .should_start(false) + ); + } + + #[test] + fn proxy_startup_policy_preserves_system_forwarding_rules() { + assert!( + !ProxyRuntimeConfig { + forward_by_system: true, + ..Default::default() + } + .should_start(true) + ); + assert!( + !ProxyRuntimeConfig { + enable_exit_node: true, + forward_by_system: true, + ..Default::default() + } + .should_start(false) + ); + assert!( + ProxyRuntimeConfig { + no_tun: true, + forward_by_system: true, + ..Default::default() + } + .should_start(false) + ); + } + + #[test] + fn proxy_runtime_defaults_preserve_udp_mtu() { + assert_eq!(ProxyRuntimeConfig::default().udp_response_ipv4_mtu, 1280); + } +} diff --git a/easytier-core/src/config/mod.rs b/easytier-core/src/config/mod.rs new file mode 100644 index 00000000..fc48a73a --- /dev/null +++ b/easytier-core/src/config/mod.rs @@ -0,0 +1,804 @@ +//! Static configuration schema plus the live runtime configuration store. + +#[cfg(feature = "management")] +pub mod api; +#[cfg(feature = "management")] +pub mod api_input; +mod encryption; +pub mod gateway; +pub mod peers; +pub mod runtime; +pub mod toml; + +pub use encryption::EncryptionAlgorithm; + +pub(crate) const DEFAULT_UDP_STUN_SERVERS: &[&str] = &[ + "txt:stun.easytier.cn", + "stun.miwifi.com", + "stun.chat.bilibili.com", + "stun.hitv.com", +]; +pub(crate) const DEFAULT_TCP_STUN_SERVERS: &[&str] = &[ + "stun.hot-chilli.net", + "stun.fitauto.ru", + "fwa.lifesizecloud.com", + "global.turn.twilio.com", + "turn.cloudflare.com", + "stun.voip.blackberry.com", + "stun.radiojar.com", +]; +pub(crate) const DEFAULT_UDP_V6_STUN_SERVERS: &[&str] = &["txt:stun-v6.easytier.cn"]; + +pub(crate) fn default_stun_servers(servers: &[&str]) -> Vec { + servers.iter().map(ToString::to_string).collect() +} + +use std::{ + collections::{BTreeSet, hash_map::DefaultHasher}, + hash::{Hash, Hasher}, + net::IpAddr, +}; + +use anyhow::Context as _; +use base64::{Engine as _, prelude::BASE64_STANDARD}; +use easytier_proto::{common as common_pb, core_config as pb}; +use serde::{Deserialize, Serialize}; +use url::Url; + +pub type PeerId = u32; + +pub type NetworkSecretDigest = [u8; 32]; + +/// Host capabilities used by the portable mapped-listener validation rule. +#[derive(Debug, Clone, Default, PartialEq, Eq)] +pub struct MappedListenerPolicy { + implicit_port_schemes: BTreeSet, +} + +impl MappedListenerPolicy { + pub fn new(implicit_port_schemes: I) -> Self + where + I: IntoIterator, + S: Into, + { + Self { + implicit_port_schemes: implicit_port_schemes + .into_iter() + .map(Into::into) + .map(|scheme: String| scheme.to_ascii_lowercase()) + .collect(), + } + } + + pub fn validate(&self, url: &Url) -> anyhow::Result<()> { + if url.port().is_none() && !self.implicit_port_schemes.contains(url.scheme()) { + anyhow::bail!("mapped listener port is missing: {}", url); + } + + Ok(()) + } + + pub fn parse_urls(&self, mapped_listeners: &[String]) -> anyhow::Result> { + mapped_listeners + .iter() + .map(|value| { + let url: Url = value + .parse() + .with_context(|| format!("mapped listener is not a valid url: {}", value))?; + self.validate(&url)?; + Ok(url) + }) + .collect() + } +} + +/// Completes and validates the portable secure-mode key configuration. +pub fn normalize_secure_mode_config( + mut config: common_pb::SecureModeConfig, +) -> anyhow::Result { + if !config.enabled { + return Ok(config); + } + + let private_key = if config.local_private_key.is_none() { + let private = x25519_dalek::StaticSecret::random_from_rng(rand::rngs::OsRng); + config.local_private_key = Some(BASE64_STANDARD.encode(private.as_bytes())); + private + } else { + config.private_key()? + }; + let generated_public_key = x25519_dalek::PublicKey::from(&private_key); + let generated_public_key = BASE64_STANDARD.encode(generated_public_key.as_bytes()); + + match config.local_public_key.as_ref() { + None => config.local_public_key = Some(generated_public_key), + Some(configured_public_key) => { + let public_key = config.public_key()?; + let canonical_public_key = BASE64_STANDARD.encode(public_key.as_bytes()); + if configured_public_key != &canonical_public_key { + anyhow::bail!( + "local public key {} does not match generated public key {}", + configured_public_key, + canonical_public_key + ); + } + } + } + + Ok(config) +} + +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct NetworkIdentity { + pub network_name: String, + pub network_secret: Option, + pub network_secret_digest: Option, +} + +impl NetworkIdentity { + pub fn new(network_name: String, network_secret: String) -> Self { + Self { + network_secret_digest: Some(network_secret_digest(&network_name, &network_secret)), + network_name, + network_secret: Some(network_secret), + } + } + + pub fn new_credential(network_name: String) -> Self { + Self { + network_name, + network_secret: None, + network_secret_digest: None, + } + } + + pub fn secret_digest(&self) -> Option { + if self.network_secret_digest.is_some() { + self.network_secret_digest + } else if let Some(network_secret) = &self.network_secret { + let mut network_secret_digest = [0u8; 32]; + generate_digest_from_str( + &self.network_name, + network_secret, + &mut network_secret_digest, + ); + Some(network_secret_digest) + } else { + None + } + } + + pub fn with_secret_digest(mut self) -> Self { + self.network_secret_digest = self.secret_digest(); + self + } +} + +#[derive(Eq, PartialEq, Hash)] +struct NetworkIdentityWithOnlyDigest { + network_name: String, + network_secret_digest: Option, +} + +fn generate_digest_from_str(str1: &str, str2: &str, digest: &mut [u8]) { + let mut hasher = DefaultHasher::new(); + hasher.write(str1.as_bytes()); + hasher.write(str2.as_bytes()); + + assert_eq!(digest.len() % 8, 0, "digest length must be multiple of 8"); + + let shard_count = digest.len() / 8; + for i in 0..shard_count { + digest[i * 8..(i + 1) * 8].copy_from_slice(&hasher.finish().to_be_bytes()); + hasher.write(&digest[..(i + 1) * 8]); + } +} + +fn network_secret_digest(network_name: &str, network_secret: &str) -> NetworkSecretDigest { + let mut digest = [0u8; 32]; + generate_digest_from_str(network_name, network_secret, &mut digest); + digest +} + +impl From for NetworkIdentityWithOnlyDigest { + fn from(identity: NetworkIdentity) -> Self { + Self { + network_secret_digest: identity.secret_digest(), + network_name: identity.network_name, + } + } +} + +impl PartialEq for NetworkIdentity { + fn eq(&self, other: &Self) -> bool { + let self_with_digest = NetworkIdentityWithOnlyDigest::from(self.clone()); + let other_with_digest = NetworkIdentityWithOnlyDigest::from(other.clone()); + self_with_digest == other_with_digest + } +} + +impl Eq for NetworkIdentity {} + +impl Hash for NetworkIdentity { + fn hash(&self, state: &mut H) { + let self_with_digest = NetworkIdentityWithOnlyDigest::from(self.clone()); + self_with_digest.hash(state); + } +} + +impl Default for NetworkIdentity { + fn default() -> Self { + Self::new("default".to_string(), "".to_string()) + } +} + +#[derive(Debug, Clone, PartialEq, Eq, Default, Serialize, Deserialize)] +pub struct CoreConfig { + pub node: NodeConfig, + pub routes: RouteConfig, + pub peer_policy: PeerPolicyConfig, + pub traffic: TrafficConfig, +} + +#[derive(Debug, Clone, PartialEq, Eq, Default, Serialize, Deserialize)] +pub struct NodeConfig { + pub peer_id: Option, + pub instance_id: Option<[u8; 16]>, + pub hostname: Option, + pub network_name: String, +} + +#[derive(Debug, Clone, PartialEq, Eq, Default, Serialize, Deserialize)] +pub struct RouteConfig { + pub ipv4: Option, + pub ipv6: Option, + pub advertised_routes: Vec, + pub proxy_networks: Vec, + pub foreign_networks: Vec, +} + +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +pub struct IpPrefix { + pub address: IpAddr, + pub prefix_len: u8, +} + +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +pub struct ProxyNetworkConfig { + pub real: IpPrefix, + pub mapped: Option, +} + +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +pub struct ForeignNetworkConfig { + pub name: String, + pub cidrs: Vec, +} + +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +pub struct PeerPolicyConfig { + pub p2p_enabled: bool, + pub relay_peer_rpc: bool, + pub relay_data: bool, + pub latency_first: bool, + pub encryption_required: bool, +} + +impl Default for PeerPolicyConfig { + fn default() -> Self { + Self { + p2p_enabled: true, + relay_peer_rpc: false, + relay_data: true, + latency_first: false, + encryption_required: true, + } + } +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)] +pub struct P2pPolicyFlags { + pub disable_udp_hole_punching: bool, + pub disable_sym_hole_punching: bool, + pub disable_upnp: bool, + pub lazy_p2p: bool, + pub disable_p2p: bool, + pub need_p2p: bool, +} + +#[derive(Debug, Clone, PartialEq, Eq, Default, Serialize, Deserialize)] +pub struct TrafficConfig { + pub mtu: Option, + pub instance_recv_bps_limit: Option, + pub foreign_relay_bps_limit: Option, +} + +#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)] +pub enum ConfigError { + #[error("missing required field: {0}")] + MissingField(&'static str), + #[error("invalid IPv4 prefix length: {0}")] + InvalidIpv4Prefix(u8), + #[error("invalid IPv6 prefix length: {0}")] + InvalidIpv6Prefix(u8), + #[error("invalid MTU: {0}")] + InvalidMtu(u32), +} + +impl IpPrefix { + pub fn new(address: IpAddr, prefix_len: u8) -> Result { + match address { + IpAddr::V4(_) if prefix_len <= 32 => Ok(Self { + address, + prefix_len, + }), + IpAddr::V4(_) => Err(ConfigError::InvalidIpv4Prefix(prefix_len)), + IpAddr::V6(_) if prefix_len <= 128 => Ok(Self { + address, + prefix_len, + }), + IpAddr::V6(_) => Err(ConfigError::InvalidIpv6Prefix(prefix_len)), + } + } +} + +impl TryFrom for CoreConfig { + type Error = ConfigError; + + fn try_from(value: pb::CoreConfig) -> Result { + Ok(Self { + node: value + .node + .map(TryInto::try_into) + .transpose()? + .unwrap_or_default(), + routes: value + .routes + .map(TryInto::try_into) + .transpose()? + .unwrap_or_default(), + peer_policy: value.peer_policy.map(Into::into).unwrap_or_default(), + traffic: value + .traffic + .map(TryInto::try_into) + .transpose()? + .unwrap_or_default(), + }) + } +} + +impl From for pb::CoreConfig { + fn from(value: CoreConfig) -> Self { + Self { + node: Some(value.node.into()), + routes: Some(value.routes.into()), + peer_policy: Some(value.peer_policy.into()), + traffic: Some(value.traffic.into()), + } + } +} + +impl TryFrom for NodeConfig { + type Error = ConfigError; + + fn try_from(value: pb::NodeConfig) -> Result { + Ok(Self { + peer_id: value.peer_id, + instance_id: value.instance_id.map(uuid_to_bytes), + hostname: value.hostname, + network_name: value.network_name, + }) + } +} + +impl From for pb::NodeConfig { + fn from(value: NodeConfig) -> Self { + Self { + peer_id: value.peer_id, + instance_id: value.instance_id.map(uuid_from_bytes), + hostname: value.hostname, + network_name: value.network_name, + } + } +} + +impl TryFrom for RouteConfig { + type Error = ConfigError; + + fn try_from(value: pb::RouteConfig) -> Result { + Ok(Self { + ipv4: value.ipv4.map(TryInto::try_into).transpose()?, + ipv6: value.ipv6.map(TryInto::try_into).transpose()?, + advertised_routes: value + .advertised_routes + .into_iter() + .map(TryInto::try_into) + .collect::>()?, + proxy_networks: value + .proxy_networks + .into_iter() + .map(TryInto::try_into) + .collect::>()?, + foreign_networks: value + .foreign_networks + .into_iter() + .map(TryInto::try_into) + .collect::>()?, + }) + } +} + +impl From for pb::RouteConfig { + fn from(value: RouteConfig) -> Self { + Self { + ipv4: value.ipv4.map(Into::into), + ipv6: value.ipv6.map(Into::into), + advertised_routes: value + .advertised_routes + .into_iter() + .map(Into::into) + .collect(), + proxy_networks: value.proxy_networks.into_iter().map(Into::into).collect(), + foreign_networks: value.foreign_networks.into_iter().map(Into::into).collect(), + } + } +} + +impl TryFrom for IpPrefix { + type Error = ConfigError; + + fn try_from(value: pb::IpPrefix) -> Result { + let address = pb_ip_addr_to_std( + value + .address + .ok_or(ConfigError::MissingField("IpPrefix.address"))?, + )?; + let prefix_len = u8::try_from(value.prefix_len) + .map_err(|_| invalid_prefix_for_address(address, value.prefix_len))?; + Self::new(address, prefix_len) + } +} + +impl From for pb::IpPrefix { + fn from(value: IpPrefix) -> Self { + Self { + address: Some(value.address.into()), + prefix_len: value.prefix_len.into(), + } + } +} + +impl TryFrom for ProxyNetworkConfig { + type Error = ConfigError; + + fn try_from(value: pb::ProxyNetworkConfig) -> Result { + Ok(Self { + real: value + .real + .ok_or(ConfigError::MissingField("ProxyNetworkConfig.real"))? + .try_into()?, + mapped: value.mapped.map(TryInto::try_into).transpose()?, + }) + } +} + +impl From for pb::ProxyNetworkConfig { + fn from(value: ProxyNetworkConfig) -> Self { + Self { + real: Some(value.real.into()), + mapped: value.mapped.map(Into::into), + } + } +} + +impl TryFrom for ForeignNetworkConfig { + type Error = ConfigError; + + fn try_from(value: pb::ForeignNetworkConfig) -> Result { + Ok(Self { + name: value.name, + cidrs: value + .cidrs + .into_iter() + .map(TryInto::try_into) + .collect::>()?, + }) + } +} + +impl From for pb::ForeignNetworkConfig { + fn from(value: ForeignNetworkConfig) -> Self { + Self { + name: value.name, + cidrs: value.cidrs.into_iter().map(Into::into).collect(), + } + } +} + +impl From for PeerPolicyConfig { + fn from(value: pb::PeerPolicyConfig) -> Self { + let default = Self::default(); + Self { + p2p_enabled: value.p2p_enabled.unwrap_or(default.p2p_enabled), + relay_peer_rpc: value.relay_peer_rpc.unwrap_or(default.relay_peer_rpc), + relay_data: value.relay_data.unwrap_or(default.relay_data), + latency_first: value.latency_first.unwrap_or(default.latency_first), + encryption_required: value + .encryption_required + .unwrap_or(default.encryption_required), + } + } +} + +impl From for pb::PeerPolicyConfig { + fn from(value: PeerPolicyConfig) -> Self { + Self { + p2p_enabled: Some(value.p2p_enabled), + relay_peer_rpc: Some(value.relay_peer_rpc), + relay_data: Some(value.relay_data), + latency_first: Some(value.latency_first), + encryption_required: Some(value.encryption_required), + } + } +} + +impl TryFrom for TrafficConfig { + type Error = ConfigError; + + fn try_from(value: pb::TrafficConfig) -> Result { + Ok(Self { + mtu: value + .mtu + .map(|mtu| u16::try_from(mtu).map_err(|_| ConfigError::InvalidMtu(mtu))) + .transpose()?, + instance_recv_bps_limit: value.instance_recv_bps_limit, + foreign_relay_bps_limit: value.foreign_relay_bps_limit, + }) + } +} + +impl From for pb::TrafficConfig { + fn from(value: TrafficConfig) -> Self { + Self { + mtu: value.mtu.map(Into::into), + instance_recv_bps_limit: value.instance_recv_bps_limit, + foreign_relay_bps_limit: value.foreign_relay_bps_limit, + } + } +} + +fn pb_ip_addr_to_std(value: common_pb::IpAddr) -> Result { + match value.ip.ok_or(ConfigError::MissingField("IpAddr.ip"))? { + common_pb::ip_addr::Ip::Ipv4(addr) => Ok(IpAddr::V4(addr.into())), + common_pb::ip_addr::Ip::Ipv6(addr) => Ok(IpAddr::V6(addr.into())), + } +} + +fn invalid_prefix_for_address(address: IpAddr, prefix_len: u32) -> ConfigError { + let prefix_len = u8::try_from(prefix_len).unwrap_or(u8::MAX); + match address { + IpAddr::V4(_) => ConfigError::InvalidIpv4Prefix(prefix_len), + IpAddr::V6(_) => ConfigError::InvalidIpv6Prefix(prefix_len), + } +} + +fn uuid_to_bytes(value: common_pb::Uuid) -> [u8; 16] { + let mut bytes = [0; 16]; + bytes[0..4].copy_from_slice(&value.part1.to_be_bytes()); + bytes[4..8].copy_from_slice(&value.part2.to_be_bytes()); + bytes[8..12].copy_from_slice(&value.part3.to_be_bytes()); + bytes[12..16].copy_from_slice(&value.part4.to_be_bytes()); + bytes +} + +fn uuid_from_bytes(bytes: [u8; 16]) -> common_pb::Uuid { + common_pb::Uuid { + part1: u32::from_be_bytes([bytes[0], bytes[1], bytes[2], bytes[3]]), + part2: u32::from_be_bytes([bytes[4], bytes[5], bytes[6], bytes[7]]), + part3: u32::from_be_bytes([bytes[8], bytes[9], bytes[10], bytes[11]]), + part4: u32::from_be_bytes([bytes[12], bytes[13], bytes[14], bytes[15]]), + } +} + +#[cfg(test)] +mod tests { + use super::*; + use base64::prelude::BASE64_STANDARD; + use x25519_dalek::{PublicKey, StaticSecret}; + + fn digest(network_name: &str, network_secret: &str) -> NetworkSecretDigest { + let mut digest = [0u8; 32]; + generate_digest_from_str(network_name, network_secret, &mut digest); + digest + } + + #[test] + fn network_identity_matches_secret_to_digest_identity() { + let local = NetworkIdentity { + network_name: "net".to_string(), + network_secret: Some("secret".to_string()), + network_secret_digest: None, + }; + let remote = NetworkIdentity { + network_name: "net".to_string(), + network_secret: None, + network_secret_digest: Some(digest("net", "secret")), + }; + + assert_eq!(local, remote); + } + + #[test] + fn network_identity_rejects_different_digest() { + let local = NetworkIdentity { + network_name: "net".to_string(), + network_secret: Some("secret".to_string()), + network_secret_digest: None, + }; + let remote = NetworkIdentity { + network_name: "net".to_string(), + network_secret: None, + network_secret_digest: Some(digest("net", "other")), + }; + + assert_ne!(local, remote); + } + + #[test] + fn network_identity_equal_values_have_equal_hash() { + let local = NetworkIdentity { + network_name: "net".to_string(), + network_secret: Some("secret".to_string()), + network_secret_digest: None, + }; + let remote = NetworkIdentity { + network_name: "net".to_string(), + network_secret: None, + network_secret_digest: Some(digest("net", "secret")), + }; + let mut local_hasher = DefaultHasher::new(); + let mut remote_hasher = DefaultHasher::new(); + + local.hash(&mut local_hasher); + remote.hash(&mut remote_hasher); + + assert_eq!(local_hasher.finish(), remote_hasher.finish()); + } + + #[test] + fn network_identity_derives_digest_from_plaintext_secret() { + let identity = NetworkIdentity { + network_name: "net".to_string(), + network_secret: Some("secret".to_string()), + network_secret_digest: None, + }; + + assert_eq!(identity.secret_digest(), Some(digest("net", "secret"))); + } + + #[test] + fn network_identity_default_matches_native_default_network() { + assert_eq!( + NetworkIdentity::default(), + NetworkIdentity::new("default".to_string(), "".to_string()) + ); + } + + #[test] + fn mapped_listener_policy_uses_explicit_host_capabilities() { + let policy = MappedListenerPolicy::new(["tcp", "ws", "wss"]); + let parsed = policy + .parse_urls(&[ + "tcp://127.0.0.1".to_string(), + "ws://example.com".to_string(), + "wss://example.com/path".to_string(), + "ring://peer-id:1000".to_string(), + ]) + .unwrap(); + + assert_eq!(parsed.len(), 4); + assert_eq!(parsed[0].scheme(), "tcp"); + assert_eq!(parsed[1].scheme(), "ws"); + assert_eq!(parsed[2].scheme(), "wss"); + assert_eq!(parsed[3].port(), Some(1000)); + + let error = policy + .parse_urls(&["ring://peer-id".to_string()]) + .unwrap_err(); + assert!( + error + .to_string() + .contains("mapped listener port is missing") + ); + } + + #[test] + fn secure_mode_normalization_generates_missing_key_pair() { + let normalized = normalize_secure_mode_config(common_pb::SecureModeConfig { + enabled: true, + local_private_key: None, + local_public_key: None, + }) + .unwrap(); + + let private_key = normalized.private_key().unwrap(); + let public_key = normalized.public_key().unwrap(); + assert_eq!(public_key, PublicKey::from(&private_key)); + } + + #[test] + fn secure_mode_normalization_preserves_existing_key_configuration() { + let private_key = StaticSecret::from([7; 32]); + let public_key = PublicKey::from(&private_key); + let config = common_pb::SecureModeConfig { + enabled: true, + local_private_key: Some(BASE64_STANDARD.encode(private_key.as_bytes())), + local_public_key: Some(BASE64_STANDARD.encode(public_key.as_bytes())), + }; + + assert_eq!( + normalize_secure_mode_config(config.clone()).unwrap(), + config + ); + } + + #[test] + fn disabled_secure_mode_does_not_validate_keys() { + let config = common_pb::SecureModeConfig { + enabled: false, + local_private_key: Some("not-base64".to_string()), + local_public_key: Some("not-base64".to_string()), + }; + + assert_eq!( + normalize_secure_mode_config(config.clone()).unwrap(), + config + ); + } + + #[test] + fn validates_ip_prefix_lengths() { + assert!(IpPrefix::new("10.0.0.1".parse().unwrap(), 24).is_ok()); + assert_eq!( + IpPrefix::new("10.0.0.1".parse().unwrap(), 33), + Err(ConfigError::InvalidIpv4Prefix(33)) + ); + assert!(IpPrefix::new("2001:db8::1".parse().unwrap(), 64).is_ok()); + assert_eq!( + IpPrefix::new("2001:db8::1".parse().unwrap(), 129), + Err(ConfigError::InvalidIpv6Prefix(129)) + ); + } + + #[test] + fn converts_core_config_from_proto_defaults() { + let config = CoreConfig::try_from(pb::CoreConfig { + node: Some(pb::NodeConfig { + peer_id: Some(7), + instance_id: None, + hostname: Some("node-a".to_string()), + network_name: "net".to_string(), + }), + routes: None, + peer_policy: None, + traffic: Some(pb::TrafficConfig { + mtu: Some(1380), + instance_recv_bps_limit: Some(100), + foreign_relay_bps_limit: None, + }), + }) + .unwrap(); + + assert_eq!(config.node.peer_id, Some(7)); + assert_eq!(config.node.hostname.as_deref(), Some("node-a")); + assert!(config.peer_policy.p2p_enabled); + assert_eq!(config.traffic.mtu, Some(1380)); + } + + #[test] + fn converts_ip_prefix_round_trip() { + let prefix = IpPrefix::new("10.1.0.1".parse().unwrap(), 16).unwrap(); + let pb: pb::IpPrefix = prefix.clone().into(); + assert_eq!(IpPrefix::try_from(pb).unwrap(), prefix); + } +} diff --git a/easytier-core/src/config/peers.rs b/easytier-core/src/config/peers.rs new file mode 100644 index 00000000..e2ff3d69 --- /dev/null +++ b/easytier-core/src/config/peers.rs @@ -0,0 +1,315 @@ +//! Peer-flavored configuration data owned by the config layer. +//! +//! These types are pure serializable configuration snapshots. Normalization +//! and derivation behavior that depends on peer-domain logic stays in +//! `crate::peers`. + +use anyhow::Context as _; +use cidr::{Ipv4Cidr, Ipv6Cidr}; +use easytier_proto::common::{FlagsInConfig, PeerFeatureFlag, SecureModeConfig, StunInfo}; +use serde::{Deserialize, Serialize}; + +use crate::proto::acl::{Acl, AclV1, Action, Chain, ChainType, GroupInfo, Protocol, Rule}; + +use super::{CoreConfig, NetworkIdentity}; + +#[derive(Debug, Clone, Default, PartialEq, Serialize, Deserialize)] +pub struct AclRuleConfig { + pub acl: Option, + pub tcp_whitelist: Vec, + pub udp_whitelist: Vec, + pub whitelist_priority: Option, +} + +#[derive(Debug, Clone, Default, PartialEq, Eq)] +pub struct AclWhitelistSnapshot { + pub tcp_ports: Vec, + pub udp_ports: Vec, +} + +impl From<&AclRuleConfig> for AclWhitelistSnapshot { + fn from(config: &AclRuleConfig) -> Self { + Self { + tcp_ports: config.tcp_whitelist.clone(), + udp_ports: config.udp_whitelist.clone(), + } + } +} + +impl AclRuleConfig { + fn parse_port_list(port_list: &[String]) -> anyhow::Result> { + let mut ports = Vec::new(); + + for port_spec in port_list { + if port_spec.contains('-') { + let parts: Vec<&str> = port_spec.split('-').collect(); + if parts.len() != 2 { + return Err(anyhow::anyhow!("Invalid port range format: {}", port_spec)); + } + + let start: u16 = parts[0] + .parse() + .with_context(|| format!("Invalid start port in range: {}", port_spec))?; + let end: u16 = parts[1] + .parse() + .with_context(|| format!("Invalid end port in range: {}", port_spec))?; + + if start > end { + return Err(anyhow::anyhow!( + "Start port must be <= end port in range: {}", + port_spec + )); + } + ports.push(port_spec.clone()); + } else { + let port: u16 = port_spec + .parse() + .with_context(|| format!("Invalid port number: {}", port_spec))?; + ports.push(port.to_string()); + } + } + + Ok(ports) + } + + fn generate_acl_from_whitelists(&mut self) -> anyhow::Result<()> { + if self.tcp_whitelist.is_empty() && self.udp_whitelist.is_empty() { + return Ok(()); + } + + let mut inbound_chain = Chain { + name: "inbound_whitelist".to_string(), + chain_type: ChainType::Inbound as i32, + description: "Auto-generated inbound whitelist from CLI".to_string(), + enabled: true, + rules: vec![], + default_action: Action::Allow as i32, + }; + + let mut rule_priority = self.whitelist_priority.unwrap_or(1000u32); + + if !self.tcp_whitelist.is_empty() { + let tcp_ports = Self::parse_port_list(&self.tcp_whitelist)?; + inbound_chain.rules.push(Rule { + name: "tcp_whitelist".to_string(), + description: "Auto-generated TCP whitelist rule".to_string(), + priority: rule_priority, + enabled: true, + protocol: Protocol::Tcp as i32, + ports: tcp_ports, + source_ips: vec![], + destination_ips: vec![], + source_ports: vec![], + action: Action::Allow as i32, + rate_limit: 0, + burst_limit: 0, + stateful: true, + source_groups: vec![], + destination_groups: vec![], + }); + inbound_chain.rules.push(Rule { + name: "tcp_whitelist_deny_other".to_string(), + description: "Auto-generated TCP whitelist rule to deny other ports".to_string(), + priority: 0, + enabled: true, + protocol: Protocol::Tcp as i32, + ports: vec!["0-65535".to_string()], + source_ips: vec![], + destination_ips: vec![], + source_ports: vec![], + action: Action::Drop as i32, + rate_limit: 0, + burst_limit: 0, + stateful: false, + source_groups: vec![], + destination_groups: vec![], + }); + rule_priority -= 1; + } + + if !self.udp_whitelist.is_empty() { + let udp_ports = Self::parse_port_list(&self.udp_whitelist)?; + inbound_chain.rules.push(Rule { + name: "udp_whitelist".to_string(), + description: "Auto-generated UDP whitelist rule".to_string(), + priority: rule_priority, + enabled: true, + protocol: Protocol::Udp as i32, + ports: udp_ports, + source_ips: vec![], + destination_ips: vec![], + source_ports: vec![], + action: Action::Allow as i32, + rate_limit: 0, + burst_limit: 0, + stateful: false, + source_groups: vec![], + destination_groups: vec![], + }); + inbound_chain.rules.push(Rule { + name: "udp_whitelist_deny_other".to_string(), + description: "Auto-generated UDP whitelist rule to deny other ports".to_string(), + priority: 0, + enabled: true, + protocol: Protocol::Udp as i32, + ports: vec!["0-65535".to_string()], + source_ips: vec![], + destination_ips: vec![], + source_ports: vec![], + action: Action::Drop as i32, + rate_limit: 0, + burst_limit: 0, + stateful: false, + source_groups: vec![], + destination_groups: vec![], + }); + } + + if self.acl.is_none() { + self.acl = Some(Acl::default()); + } + + let acl = self.acl.as_mut().expect("ACL was initialized above"); + if let Some(acl_v1) = acl.acl_v1.as_mut() { + acl_v1.chains.push(inbound_chain); + } else { + acl.acl_v1 = Some(AclV1 { + chains: vec![inbound_chain], + group: Some(GroupInfo { + declares: vec![], + members: vec![], + }), + }); + } + + Ok(()) + } + + pub fn build(&self) -> anyhow::Result> { + let mut config = self.clone(); + config.generate_acl_from_whitelists()?; + Ok(config.acl) + } +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] +pub struct PublicIpv6ProviderConfig { + pub provider_enabled: bool, + pub configured_prefix: Option, + pub provider_supported: bool, +} + +impl PublicIpv6ProviderConfig { + pub fn should_run_reconcile(self) -> bool { + self.provider_enabled + } +} + +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct PeerRuntimeConfig { + pub core: CoreConfig, + pub network_identity: NetworkIdentity, + pub stun_info: StunInfo, + pub feature_flags: PeerFeatureFlag, + pub secure_mode: Option, + pub host_routing: HostRoutingPolicy, +} + +#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Serialize, Deserialize)] +pub struct HostRoutingPolicy { + /// Route otherwise-unreachable external IPv4 traffic through this node and + /// keep self-delivered packets eligible for the host TUN/proxy path. + pub local_exit_node_fallback: bool, +} + +/// One normalized peer configuration version submitted by a host. +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct PeerRuntimeSnapshot { + pub runtime: PeerRuntimeConfig, + pub easytier_version: String, + pub avoid_relay_data_preference: bool, + pub flags: FlagsInConfig, + pub vpn_portal_cidr: Option, + pub pinned_peers: Vec<(url::Url, Option)>, + pub peer_group_memberships: Vec, + pub acl_group_declarations: Vec, + pub ospf_update_my_foreign_network_interval_sec: u64, + pub max_direct_conns_per_peer_in_foreign_network: usize, + pub hmac_secret_digest: bool, +} + +impl PeerRuntimeSnapshot { + pub fn new(runtime: PeerRuntimeConfig, flags: FlagsInConfig) -> Self { + let avoid_relay_data_preference = runtime.feature_flags.avoid_relay_data; + Self { + runtime, + easytier_version: env!("CARGO_PKG_VERSION").to_owned(), + avoid_relay_data_preference, + flags, + vpn_portal_cidr: None, + pinned_peers: Vec::new(), + peer_group_memberships: Vec::new(), + acl_group_declarations: Vec::new(), + ospf_update_my_foreign_network_interval_sec: 10, + max_direct_conns_per_peer_in_foreign_network: 3, + hmac_secret_digest: false, + } + } +} + +impl Default for PeerRuntimeSnapshot { + fn default() -> Self { + Self::new( + PeerRuntimeConfig { + core: CoreConfig::default(), + network_identity: NetworkIdentity::default(), + stun_info: StunInfo::default(), + feature_flags: PeerFeatureFlag::default(), + secure_mode: None, + host_routing: HostRoutingPolicy::default(), + }, + FlagsInConfig::default(), + ) + } +} + +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +pub struct PeerGroupIdentity { + pub group_name: String, + pub group_secret: String, +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn whitelist_rules_are_built_in_core() { + let acl = AclRuleConfig { + tcp_whitelist: vec!["80".to_string(), "8000-9000".to_string()], + udp_whitelist: vec!["53".to_string()], + ..Default::default() + } + .build() + .unwrap() + .unwrap(); + + let chain = &acl.acl_v1.unwrap().chains[0]; + assert_eq!(chain.name, "inbound_whitelist"); + assert_eq!(chain.rules.len(), 4); + assert_eq!(chain.rules[0].ports, ["80", "8000-9000"]); + assert_eq!(chain.rules[2].ports, ["53"]); + } + + #[test] + fn invalid_whitelist_range_is_rejected() { + let error = AclRuleConfig { + tcp_whitelist: vec!["9000-8000".to_string()], + ..Default::default() + } + .build() + .unwrap_err(); + + assert!(error.to_string().contains("Start port must be <= end port")); + } +} diff --git a/easytier-core/src/config/runtime.rs b/easytier-core/src/config/runtime.rs new file mode 100644 index 00000000..d90c9026 --- /dev/null +++ b/easytier-core/src/config/runtime.rs @@ -0,0 +1,253 @@ +//! Atomic runtime configuration owned by one core instance. + +use std::{collections::BTreeSet, sync::Arc}; + +use arc_swap::ArcSwap; +use cidr::Ipv4Cidr; +use parking_lot::Mutex; +use serde::{Deserialize, Serialize}; + +use super::{ + gateway::{GatewayRuntimeConfig, ProxyRuntimeConfig}, + peers::{AclRuleConfig, PeerRuntimeSnapshot, PublicIpv6ProviderConfig}, +}; + +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct CoreRuntimeConfig { + pub acl: AclRuleConfig, + pub dhcp_ipv4: bool, + pub gateway: GatewayRuntimeConfig, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub manual_routes: Option>, + pub proxy: ProxyRuntimeConfig, + #[serde(default)] + pub public_ipv6_auto: bool, + pub public_ipv6_provider: PublicIpv6ProviderConfig, +} + +impl Default for CoreRuntimeConfig { + fn default() -> Self { + Self { + acl: AclRuleConfig::default(), + dhcp_ipv4: false, + gateway: GatewayRuntimeConfig::default(), + manual_routes: None, + proxy: ProxyRuntimeConfig::default(), + public_ipv6_auto: false, + public_ipv6_provider: PublicIpv6ProviderConfig { + provider_enabled: false, + configured_prefix: None, + provider_supported: false, + }, + } + } +} + +#[derive(Debug, Clone)] +pub struct CoreInstanceRuntimeConfig { + pub services: CoreRuntimeConfig, + pub peer: Arc, +} + +struct CoreRuntimeConfigStoreInner { + snapshot: ArcSwap, + update: Mutex<()>, + peer_changes: tokio::sync::watch::Sender, + service_changes: tokio::sync::watch::Sender, +} + +/// Atomic configuration authority shared by one core instance and its peer +/// context. Readers always observe a complete submitted version. +#[derive(Clone)] +pub struct CoreRuntimeConfigStore { + inner: Arc, +} + +impl CoreRuntimeConfigStore { + pub fn new(services: CoreRuntimeConfig, peer: Arc) -> Self { + let (peer_changes, _) = tokio::sync::watch::channel(0); + let (service_changes, _) = tokio::sync::watch::channel(0); + Self { + inner: Arc::new(CoreRuntimeConfigStoreInner { + snapshot: ArcSwap::from_pointee(CoreInstanceRuntimeConfig { services, peer }), + update: Mutex::new(()), + peer_changes, + service_changes, + }), + } + } + + pub fn snapshot(&self) -> Arc { + self.inner.snapshot.load_full() + } + + pub fn replace(&self, config: CoreInstanceRuntimeConfig) { + let _update = self.inner.update.lock(); + self.inner.snapshot.store(Arc::new(config)); + self.inner.peer_changes.send_modify(|version| *version += 1); + self.inner + .service_changes + .send_modify(|version| *version += 1); + } + + pub fn update_services(&self, update: impl FnOnce(&mut CoreRuntimeConfig)) { + let _update = self.inner.update.lock(); + let mut config = self.inner.snapshot.load_full().as_ref().clone(); + update(&mut config.services); + self.inner.snapshot.store(Arc::new(config)); + self.inner + .service_changes + .send_modify(|version| *version += 1); + } + + pub fn update_peer(&self, peer: Arc) { + let _update = self.inner.update.lock(); + let mut config = self.inner.snapshot.load_full().as_ref().clone(); + config.peer = peer; + self.inner.snapshot.store(Arc::new(config)); + self.inner.peer_changes.send_modify(|version| *version += 1); + } + + pub(crate) fn update_peer_with(&self, update: impl FnOnce(&mut PeerRuntimeSnapshot)) { + let _update = self.inner.update.lock(); + let mut config = self.inner.snapshot.load_full().as_ref().clone(); + update(Arc::make_mut(&mut config.peer)); + self.inner.snapshot.store(Arc::new(config)); + self.inner.peer_changes.send_modify(|version| *version += 1); + } + + pub fn subscribe_peer_runtime_changes(&self) -> tokio::sync::watch::Receiver { + self.inner.peer_changes.subscribe() + } + + pub fn subscribe_service_runtime_changes(&self) -> tokio::sync::watch::Receiver { + self.inner.service_changes.subscribe() + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn replaces_service_and_peer_as_one_version() { + let mut before_peer = PeerRuntimeSnapshot::default(); + before_peer.runtime.core.node.hostname = Some("before".to_owned()); + let store = + CoreRuntimeConfigStore::new(CoreRuntimeConfig::default(), Arc::new(before_peer)); + let before = store.snapshot(); + + let after_services = CoreRuntimeConfig { + dhcp_ipv4: true, + ..Default::default() + }; + let mut after_peer = PeerRuntimeSnapshot::default(); + after_peer.runtime.core.node.hostname = Some("after".to_owned()); + store.replace(CoreInstanceRuntimeConfig { + services: after_services, + peer: Arc::new(after_peer), + }); + + assert!(!before.services.dhcp_ipv4); + assert_eq!( + before.peer.runtime.core.node.hostname.as_deref(), + Some("before") + ); + let after = store.snapshot(); + assert!(after.services.dhcp_ipv4); + assert_eq!( + after.peer.runtime.core.node.hostname.as_deref(), + Some("after") + ); + } + + #[tokio::test] + async fn notifies_peer_snapshot_changes() { + let store = CoreRuntimeConfigStore::new( + CoreRuntimeConfig::default(), + Arc::new(PeerRuntimeSnapshot::default()), + ); + let mut changes = store.subscribe_peer_runtime_changes(); + let mut peer = PeerRuntimeSnapshot::default(); + peer.runtime.core.node.hostname = Some("updated".to_owned()); + + store.update_peer(Arc::new(peer)); + + assert!(changes.changed().await.is_ok()); + } + + #[tokio::test] + async fn notifies_service_snapshot_changes() { + let store = CoreRuntimeConfigStore::new( + CoreRuntimeConfig::default(), + Arc::new(PeerRuntimeSnapshot::default()), + ); + let mut changes = store.subscribe_service_runtime_changes(); + + store.update_services(|services| services.dhcp_ipv4 = true); + + assert!(changes.changed().await.is_ok()); + assert!(store.snapshot().services.dhcp_ipv4); + } + + #[tokio::test] + async fn peer_update_does_not_notify_service_watchers() { + let store = CoreRuntimeConfigStore::new( + CoreRuntimeConfig::default(), + Arc::new(PeerRuntimeSnapshot::default()), + ); + let changes = store.subscribe_service_runtime_changes(); + let mut peer = PeerRuntimeSnapshot::default(); + peer.runtime.core.node.hostname = Some("updated".to_owned()); + + store.update_peer(Arc::new(peer)); + + assert!(!changes.has_changed().unwrap()); + } + + #[test] + fn peer_in_place_update_preserves_the_rest_of_the_atomic_snapshot() { + let services = CoreRuntimeConfig { + dhcp_ipv4: true, + ..Default::default() + }; + let mut peer = PeerRuntimeSnapshot::default(); + peer.runtime.core.node.hostname = Some("preserved".to_owned()); + let store = CoreRuntimeConfigStore::new(services, Arc::new(peer)); + + store.update_peer_with(|peer| { + peer.runtime.core.routes.ipv4 = Some(crate::config::IpPrefix { + address: "10.20.30.7".parse().unwrap(), + prefix_len: 24, + }); + }); + + let snapshot = store.snapshot(); + assert!(snapshot.services.dhcp_ipv4); + assert_eq!( + snapshot.peer.runtime.core.node.hostname.as_deref(), + Some("preserved") + ); + assert_eq!( + snapshot + .peer + .runtime + .core + .routes + .ipv4 + .as_ref() + .unwrap() + .address, + "10.20.30.7".parse::().unwrap() + ); + } + + #[test] + fn missing_manual_routes_preserves_portable_config_compatibility() { + let encoded = serde_json::to_value(CoreRuntimeConfig::default()).unwrap(); + assert!(encoded.get("manual_routes").is_none()); + + let decoded: CoreRuntimeConfig = serde_json::from_value(encoded).unwrap(); + assert_eq!(decoded.manual_routes, None); + } +} diff --git a/easytier-core/src/config/toml.rs b/easytier-core/src/config/toml.rs new file mode 100644 index 00000000..6ca0eeeb --- /dev/null +++ b/easytier-core/src/config/toml.rs @@ -0,0 +1,1632 @@ +//! Complete EasyTier TOML configuration model. + +use std::{ + net::{IpAddr, SocketAddr}, + path::PathBuf, + sync::{Arc, Mutex}, +}; + +pub use super::{EncryptionAlgorithm, gateway::PortForwardConfig}; +use anyhow::Context; +#[cfg(feature = "rich-config-errors")] +use ariadne::{CharSet, Config as AriadneConfig, IndexType, Label, Report, ReportKind, Source}; +use serde::{Deserialize, Serialize}; + +#[cfg(feature = "config-write")] +use crate::config::{DEFAULT_UDP_STUN_SERVERS, DEFAULT_UDP_V6_STUN_SERVERS, default_stun_servers}; +use crate::proto::{ + acl::Acl, + common::{CompressionAlgoPb, SecureModeConfig}, +}; + +pub const DEFAULT_ET_DNS_ZONE: &str = "et.net."; + +pub type Flags = crate::proto::common::FlagsInConfig; + +pub(crate) fn default_instance_name() -> String { + "default".to_owned() +} + +#[cfg(feature = "config-write")] +fn default_udp_stun_servers() -> Vec { + default_stun_servers(DEFAULT_UDP_STUN_SERVERS) +} + +#[cfg(feature = "config-write")] +fn default_udp_v6_stun_servers() -> Vec { + default_stun_servers(DEFAULT_UDP_V6_STUN_SERVERS) +} + +pub fn gen_default_flags() -> Flags { + #[allow(deprecated)] + Flags { + default_protocol: "tcp".to_string(), + dev_name: "".to_string(), + enable_encryption: true, + enable_ipv6: true, + mtu: 1380, + latency_first: false, + enable_exit_node: false, + proxy_forward_by_system: false, + no_tun: false, + use_smoltcp: false, + relay_network_whitelist: "*".to_string(), + disable_p2p: false, + p2p_only: false, + lazy_p2p: false, + relay_all_peer_rpc: false, + disable_tcp_hole_punching: false, + disable_udp_hole_punching: false, + multi_thread: true, + data_compress_algo: CompressionAlgoPb::None.into(), + bind_device: true, + enable_kcp_proxy: false, + disable_kcp_input: false, + disable_relay_kcp: false, + enable_relay_foreign_network_kcp: false, + accept_dns: false, + private_mode: false, + enable_quic_proxy: false, + disable_quic_input: false, + disable_relay_quic: false, + enable_relay_foreign_network_quic: false, + foreign_relay_bps_limit: u64::MAX, + multi_thread_count: 2, + encryption_algorithm: EncryptionAlgorithm::default().to_string(), + disable_sym_hole_punching: false, + tld_dns_zone: DEFAULT_ET_DNS_ZONE.to_string(), + + quic_listen_port: u32::MAX, + need_p2p: false, + instance_recv_bps_limit: u64::MAX, + disable_upnp: false, + disable_relay_data: false, + enable_udp_broadcast_relay: false, + socket_mark: None, + } +} + +#[cfg(feature = "config-write")] +macro_rules! define_flags_diff { + ( + fields: [$($field:ident),* $(,)?], + u64s: [$($u64_field:ident),* $(,)?], + enums: [$($enum_field:ident),* $(,)?] + ) => { + #[allow(deprecated)] + fn flags_diff_from_default(flags: &Flags) -> serde_json::Map { + let defaults = gen_default_flags(); + let mut changed = serde_json::Map::new(); + $( + if flags.$field != defaults.$field { + changed.insert( + stringify!($field).to_owned(), + serde_json::to_value(&flags.$field) + .expect("FlagsInConfig field should serialize to JSON"), + ); + } + )* + $( + if flags.$u64_field != defaults.$u64_field { + changed.insert( + stringify!($u64_field).to_owned(), + serde_json::json!(flags.$u64_field.to_string()), + ); + } + )* + $( + if flags.$enum_field != defaults.$enum_field { + let value = CompressionAlgoPb::try_from(flags.$enum_field) + .map(|value| serde_json::to_value(value).expect("enum should serialize")) + .unwrap_or_else(|_| serde_json::json!(flags.$enum_field)); + changed.insert(stringify!($enum_field).to_owned(), value); + } + )* + changed + } + + #[cfg(all(test, feature = "config-write"))] + const FLAGS_DIFF_FIELDS: &[&str] = &[ + $(stringify!($field),)* + $(stringify!($u64_field),)* + $(stringify!($enum_field),)* + ]; + }; +} + +#[cfg(feature = "config-write")] +define_flags_diff! { + fields: [ + default_protocol, + dev_name, + enable_encryption, + enable_ipv6, + mtu, + latency_first, + enable_exit_node, + no_tun, + use_smoltcp, + relay_network_whitelist, + disable_p2p, + relay_all_peer_rpc, + disable_udp_hole_punching, + multi_thread, + bind_device, + enable_kcp_proxy, + disable_kcp_input, + disable_relay_kcp, + proxy_forward_by_system, + accept_dns, + private_mode, + enable_quic_proxy, + disable_quic_input, + disable_relay_quic, + quic_listen_port, + multi_thread_count, + enable_relay_foreign_network_kcp, + enable_relay_foreign_network_quic, + encryption_algorithm, + disable_sym_hole_punching, + tld_dns_zone, + p2p_only, + disable_tcp_hole_punching, + lazy_p2p, + need_p2p, + disable_upnp, + disable_relay_data, + enable_udp_broadcast_relay, + socket_mark, + ], + u64s: [foreign_relay_bps_limit, instance_recv_bps_limit], + enums: [data_compress_algo] +} + +#[auto_impl::auto_impl(Box, &)] +pub trait ConfigLoader: Send + Sync { + fn get_id(&self) -> uuid::Uuid; + fn set_id(&self, id: uuid::Uuid); + + fn get_hostname(&self) -> String; + fn set_hostname(&self, name: Option); + + fn get_inst_name(&self) -> String; + fn set_inst_name(&self, name: String); + + fn get_netns(&self) -> Option; + fn set_netns(&self, ns: Option); + + fn get_ipv4(&self) -> Option; + fn set_ipv4(&self, addr: Option); + + fn get_ipv6(&self) -> Option; + fn set_ipv6(&self, addr: Option); + + fn get_ipv6_public_addr_provider(&self) -> bool; + fn set_ipv6_public_addr_provider(&self, enabled: bool); + + fn get_ipv6_public_addr_auto(&self) -> bool; + fn set_ipv6_public_addr_auto(&self, enabled: bool); + + fn get_ipv6_public_addr_prefix(&self) -> Option; + fn set_ipv6_public_addr_prefix(&self, prefix: Option); + + fn get_dhcp(&self) -> bool; + fn set_dhcp(&self, dhcp: bool); + + fn add_proxy_cidr( + &self, + cidr: cidr::Ipv4Cidr, + mapped_cidr: Option, + ) -> Result<(), anyhow::Error>; + fn remove_proxy_cidr(&self, cidr: cidr::Ipv4Cidr); + fn clear_proxy_cidrs(&self); + fn get_proxy_cidrs(&self) -> Vec; + + fn get_network_identity(&self) -> NetworkIdentity; + fn set_network_identity(&self, identity: NetworkIdentity); + + fn get_listener_uris(&self) -> Vec; + + fn get_peers(&self) -> Vec; + fn set_peers(&self, peers: Vec); + + fn get_listeners(&self) -> Option>; + fn set_listeners(&self, listeners: Vec); + + fn get_mapped_listeners(&self) -> Vec; + fn set_mapped_listeners(&self, listeners: Option>); + + fn get_vpn_portal_config(&self) -> Option; + fn set_vpn_portal_config(&self, config: VpnPortalConfig); + + fn get_flags(&self) -> Flags; + fn set_flags(&self, flags: Flags); + + fn get_exit_nodes(&self) -> Vec; + fn set_exit_nodes(&self, nodes: Vec); + + fn get_routes(&self) -> Option>; + fn set_routes(&self, routes: Option>); + + fn get_socks5_portal(&self) -> Option; + fn set_socks5_portal(&self, addr: Option); + + fn get_port_forwards(&self) -> Vec; + fn set_port_forwards(&self, forwards: Vec); + + fn get_acl(&self) -> Option; + fn set_acl(&self, acl: Option); + + fn get_tcp_whitelist(&self) -> Vec; + fn set_tcp_whitelist(&self, whitelist: Vec); + + fn get_udp_whitelist(&self) -> Vec; + fn set_udp_whitelist(&self, whitelist: Vec); + + fn get_stun_servers(&self) -> Option>; + fn set_stun_servers(&self, servers: Option>); + + fn get_stun_servers_v6(&self) -> Option>; + fn set_stun_servers_v6(&self, servers: Option>); + + fn get_secure_mode(&self) -> Option; + fn set_secure_mode(&self, secure_mode: Option); + + fn get_credential_file(&self) -> Option { + None + } + fn set_credential_file(&self, _path: Option) {} + + fn get_network_config_source(&self) -> ConfigSource { + ConfigSource::User + } + fn set_network_config_source(&self, _source: Option) {} + + fn dump(&self) -> String; +} + +pub trait LoggingConfigLoader { + fn get_file_logger_config(&self) -> FileLoggerConfig; + + fn get_console_logger_config(&self) -> ConsoleLoggerConfig; +} + +use super::NetworkSecretDigest; + +#[derive(Debug, Clone, Deserialize, Serialize)] +pub struct NetworkIdentity { + pub network_name: String, + pub network_secret: Option, + #[serde(skip)] + pub network_secret_digest: Option, +} + +impl From for NetworkIdentity { + fn from(value: super::NetworkIdentity) -> Self { + Self { + network_name: value.network_name, + network_secret: value.network_secret, + network_secret_digest: value.network_secret_digest, + } + } +} + +impl From<&NetworkIdentity> for super::NetworkIdentity { + fn from(value: &NetworkIdentity) -> Self { + Self { + network_name: value.network_name.clone(), + network_secret: value.network_secret.clone(), + network_secret_digest: value.network_secret_digest, + } + } +} + +impl From for super::NetworkIdentity { + fn from(value: NetworkIdentity) -> Self { + Self { + network_name: value.network_name, + network_secret: value.network_secret, + network_secret_digest: value.network_secret_digest, + } + } +} + +#[derive(Debug, Clone, Copy, Deserialize, Serialize, PartialEq, Eq, Default)] +#[serde(rename_all = "snake_case")] +pub enum ConfigSource { + #[default] + User, + Web, +} + +impl ConfigSource { + pub fn as_str(self) -> &'static str { + match self { + Self::User => "user", + Self::Web => "web", + } + } +} + +impl std::str::FromStr for ConfigSource { + type Err = String; + + fn from_str(s: &str) -> Result { + match s { + "user" => Ok(Self::User), + "web" => Ok(Self::Web), + other => Err(format!("unknown network config source: {other}")), + } + } +} + +#[derive(Debug, Clone, Deserialize, Serialize, PartialEq)] +struct ConfigSourceConfig { + source: ConfigSource, +} + +impl PartialEq for NetworkIdentity { + fn eq(&self, other: &Self) -> bool { + super::NetworkIdentity::from(self) == super::NetworkIdentity::from(other) + } +} + +impl Eq for NetworkIdentity {} + +impl std::hash::Hash for NetworkIdentity { + fn hash(&self, state: &mut H) { + std::hash::Hash::hash(&super::NetworkIdentity::from(self), state); + } +} + +impl NetworkIdentity { + pub fn new(network_name: String, network_secret: String) -> Self { + super::NetworkIdentity::new(network_name, network_secret).into() + } + + /// Create a NetworkIdentity for a credential node (no network_secret). + /// The node identifies by network_name only and authenticates via credential keypair. + pub fn new_credential(network_name: String) -> Self { + super::NetworkIdentity::new_credential(network_name).into() + } +} + +impl Default for NetworkIdentity { + fn default() -> Self { + super::NetworkIdentity::default().into() + } +} + +#[derive(Debug, Clone, Deserialize, Serialize, PartialEq)] +pub struct PeerConfig { + pub uri: url::Url, + pub peer_public_key: Option, +} + +#[derive(Debug, Clone, Deserialize, Serialize, PartialEq)] +pub struct ProxyNetworkConfig { + pub cidr: cidr::Ipv4Cidr, // the CIDR of the proxy network + pub mapped_cidr: Option, // allow remap the proxy CIDR to another CIDR + pub allow: Option>, +} + +#[derive(Debug, Clone, Deserialize, Serialize, PartialEq, Default)] +pub struct FileLoggerConfig { + pub level: Option, + pub file: Option, + pub dir: Option, + pub size_mb: Option, + pub count: Option, +} + +#[derive(Debug, Clone, Deserialize, Serialize, PartialEq, Default)] +pub struct ConsoleLoggerConfig { + pub level: Option, +} + +#[derive(Debug, Clone, Deserialize, Serialize, PartialEq, derive_builder::Builder)] +pub struct LoggingConfig { + #[builder(setter(into, strip_option), default = None)] + pub file_logger: Option, + #[builder(setter(into, strip_option), default = None)] + pub console_logger: Option, +} + +impl LoggingConfigLoader for &LoggingConfig { + fn get_file_logger_config(&self) -> FileLoggerConfig { + self.file_logger.clone().unwrap_or_default() + } + + fn get_console_logger_config(&self) -> ConsoleLoggerConfig { + self.console_logger.clone().unwrap_or_default() + } +} + +#[derive(Debug, Clone, Deserialize, Serialize, PartialEq)] +pub struct VpnPortalConfig { + pub client_cidr: cidr::Ipv4Cidr, + pub wireguard_listen: SocketAddr, +} + +#[derive(Debug, Clone, PartialEq, Deserialize)] +#[cfg_attr(feature = "config-write", derive(Serialize))] +struct Config { + netns: Option, + hostname: Option, + instance_name: Option, + instance_id: Option, + ipv4: Option, + ipv6: Option, + ipv6_public_addr_provider: Option, + ipv6_public_addr_auto: Option, + ipv6_public_addr_prefix: Option, + dhcp: Option, + network_identity: Option, + listeners: Option>, + mapped_listeners: Option>, + exit_nodes: Option>, + + peer: Option>, + proxy_network: Option>, + + vpn_portal_config: Option, + + routes: Option>, + + socks5_proxy: Option, + + port_forward: Option>, + + secure_mode: Option, + + flags: Option>, + + #[serde(skip)] + flags_struct: Option, + + acl: Option, + + tcp_whitelist: Option>, + udp_whitelist: Option>, + stun_servers: Option>, + stun_servers_v6: Option>, + + credential_file: Option, + source: Option, +} + +#[cfg(feature = "rich-config-errors")] +fn format_toml_parse_error(source_name: &str, config_str: &str, error: &toml::de::Error) -> String { + let message = format!("failed to parse config TOML from {source_name}"); + + let Some(span) = error.span() else { + return format!("{message}\ndetail: {error}"); + }; + + let mut output = Vec::new(); + let report = Report::build(ReportKind::Error, (source_name, span.clone())) + .with_config( + AriadneConfig::default() + .with_color(false) + .with_char_set(CharSet::Ascii) + .with_index_type(IndexType::Byte), + ) + .with_message(&message) + .with_label(Label::new((source_name, span)).with_message(error.message())) + .finish(); + + if report + .write((source_name, Source::from(config_str)), &mut output) + .is_ok() + { + String::from_utf8_lossy(&output).into_owned() + } else { + format!("{message}\ndetail: {error}") + } +} + +#[cfg(not(feature = "rich-config-errors"))] +fn format_toml_parse_error( + source_name: &str, + _config_str: &str, + error: &toml::de::Error, +) -> String { + format!("failed to parse config TOML from {source_name}: {error}") +} + +#[derive(Debug, Clone)] +pub struct TomlConfig { + config: Arc>, +} + +impl Default for TomlConfig { + fn default() -> Self { + TomlConfig::new_from_str("").unwrap() + } +} + +impl TomlConfig { + fn normalize_config_source(config: &mut Config) { + if matches!( + config.source.as_ref().map(|source| source.source), + Some(ConfigSource::User) + ) { + config.source = None; + } + } + + pub fn new_from_str(config_str: &str) -> Result { + Self::new_from_str_with_source("inline config", config_str) + } + + pub fn new_from_str_with_source( + source_name: &str, + config_str: &str, + ) -> Result { + let mut config = toml::de::from_str::(config_str).map_err(|err| { + let message = format_toml_parse_error(source_name, config_str, &err); + anyhow::Error::new(err).context(message) + })?; + + Self::normalize_config_source(&mut config); + + Self::new_from_config(config).map_err(|err| { + let message = format!("failed to load config from {source_name}: {err}"); + err.context(message) + }) + } + + fn new_from_config(mut config: Config) -> Result { + config.flags_struct = Some( + Self::gen_flags(config.flags.clone().unwrap_or_default()) + .context("failed to parse flags")?, + ); + let has_network_identity = config.network_identity.is_some(); + + let config = TomlConfig { + config: Arc::new(Mutex::new(config)), + }; + + let old_ns = config.get_network_identity(); + + // Detect credential mode: secure_mode enabled + no network_secret in TOML + let is_credential = has_network_identity + && config + .get_secure_mode() + .map(|sm| sm.enabled) + .unwrap_or(false) + && old_ns + .network_secret + .as_deref() + .is_none_or(|s| s.is_empty()); + + if is_credential { + config.set_network_identity(NetworkIdentity::new_credential(old_ns.network_name)); + } else { + config.set_network_identity(NetworkIdentity::new( + old_ns.network_name, + old_ns.network_secret.unwrap_or_default(), + )); + } + + Ok(config) + } + + fn gen_flags( + flags_hashmap: serde_json::Map, + ) -> serde_json::Result { + let mut merged_hashmap = match serde_json::to_value(gen_default_flags()) { + Ok(serde_json::Value::Object(map)) => map, + _ => serde_json::Map::new(), + }; + merged_hashmap.extend(flags_hashmap); + serde_json::from_value(serde_json::Value::Object(merged_hashmap)) + } +} + +#[cfg(feature = "management")] +mod snapshot; + +impl ConfigLoader for TomlConfig { + fn get_inst_name(&self) -> String { + self.config + .lock() + .unwrap() + .instance_name + .clone() + .unwrap_or_else(default_instance_name) + } + + fn set_inst_name(&self, name: String) { + self.config.lock().unwrap().instance_name = Some(name); + } + + fn get_hostname(&self) -> String { + let hostname = self.config.lock().unwrap().hostname.clone(); + + match hostname { + Some(hostname) => { + let hostname = hostname + .chars() + .filter(|c| !c.is_control()) + .take(32) + .collect::(); + + if !hostname.is_empty() { + self.set_hostname(Some(hostname.clone())); + hostname + } else { + self.set_hostname(None); + String::new() + } + } + None => String::new(), + } + } + + fn set_hostname(&self, name: Option) { + self.config.lock().unwrap().hostname = name; + } + + fn get_netns(&self) -> Option { + self.config.lock().unwrap().netns.clone() + } + + fn set_netns(&self, ns: Option) { + self.config.lock().unwrap().netns = ns; + } + + fn get_ipv4(&self) -> Option { + let locked_config = self.config.lock().unwrap(); + locked_config + .ipv4 + .as_ref() + .and_then(|s| s.parse().ok()) + .map(|c: cidr::Ipv4Inet| { + if c.network_length() == 32 { + cidr::Ipv4Inet::new(c.address(), 24).unwrap() + } else { + c + } + }) + } + + fn set_ipv4(&self, addr: Option) { + self.config.lock().unwrap().ipv4 = addr.map(|addr| addr.to_string()); + } + + fn get_ipv6(&self) -> Option { + let locked_config = self.config.lock().unwrap(); + locked_config.ipv6.as_ref().and_then(|s| s.parse().ok()) + } + + fn set_ipv6(&self, addr: Option) { + self.config.lock().unwrap().ipv6 = addr.map(|addr| addr.to_string()); + } + + fn get_ipv6_public_addr_provider(&self) -> bool { + self.config + .lock() + .unwrap() + .ipv6_public_addr_provider + .unwrap_or_default() + } + + fn set_ipv6_public_addr_provider(&self, enabled: bool) { + self.config.lock().unwrap().ipv6_public_addr_provider = Some(enabled); + } + + fn get_ipv6_public_addr_auto(&self) -> bool { + self.config + .lock() + .unwrap() + .ipv6_public_addr_auto + .unwrap_or_default() + } + + fn set_ipv6_public_addr_auto(&self, enabled: bool) { + self.config.lock().unwrap().ipv6_public_addr_auto = Some(enabled); + } + + fn get_ipv6_public_addr_prefix(&self) -> Option { + let locked_config = self.config.lock().unwrap(); + locked_config + .ipv6_public_addr_prefix + .as_ref() + .and_then(|s| s.parse().ok()) + } + + fn set_ipv6_public_addr_prefix(&self, prefix: Option) { + self.config.lock().unwrap().ipv6_public_addr_prefix = + prefix.map(|prefix| prefix.to_string()); + } + + fn get_dhcp(&self) -> bool { + self.config.lock().unwrap().dhcp.unwrap_or_default() + } + + fn set_dhcp(&self, dhcp: bool) { + self.config.lock().unwrap().dhcp = Some(dhcp); + } + + fn add_proxy_cidr( + &self, + cidr: cidr::Ipv4Cidr, + mapped_cidr: Option, + ) -> Result<(), anyhow::Error> { + let mut locked_config = self.config.lock().unwrap(); + if locked_config.proxy_network.is_none() { + locked_config.proxy_network = Some(vec![]); + } + if let Some(mapped_cidr) = mapped_cidr.as_ref() + && cidr.network_length() != mapped_cidr.network_length() + { + return Err(anyhow::anyhow!( + "Mapped CIDR must have the same network length as the original CIDR: {} != {}", + cidr.network_length(), + mapped_cidr.network_length() + )); + } + // insert if no duplicate + if !locked_config + .proxy_network + .as_ref() + .unwrap() + .iter() + .any(|c| c.cidr == cidr && c.mapped_cidr == mapped_cidr) + { + locked_config + .proxy_network + .as_mut() + .unwrap() + .push(ProxyNetworkConfig { + cidr, + mapped_cidr, + allow: None, + }); + } + Ok(()) + } + + fn remove_proxy_cidr(&self, cidr: cidr::Ipv4Cidr) { + let mut locked_config = self.config.lock().unwrap(); + if let Some(proxy_cidrs) = &mut locked_config.proxy_network { + proxy_cidrs.retain(|c| c.cidr != cidr); + } + } + + fn clear_proxy_cidrs(&self) { + let mut locked_config = self.config.lock().unwrap(); + locked_config.proxy_network = None; + } + + fn get_proxy_cidrs(&self) -> Vec { + self.config + .lock() + .unwrap() + .proxy_network + .as_ref() + .cloned() + .unwrap_or_default() + } + + fn get_id(&self) -> uuid::Uuid { + let mut locked_config = self.config.lock().unwrap(); + match locked_config.instance_id { + Some(id) => id, + None => { + let id = uuid::Uuid::new_v4(); + locked_config.instance_id = Some(id); + id + } + } + } + + fn set_id(&self, id: uuid::Uuid) { + self.config.lock().unwrap().instance_id = Some(id); + } + + fn get_network_identity(&self) -> NetworkIdentity { + self.config + .lock() + .unwrap() + .network_identity + .clone() + .unwrap_or_default() + } + + fn set_network_identity(&self, identity: NetworkIdentity) { + self.config.lock().unwrap().network_identity = Some(identity); + } + + fn get_listener_uris(&self) -> Vec { + self.config + .lock() + .unwrap() + .listeners + .clone() + .unwrap_or_default() + } + + fn get_peers(&self) -> Vec { + self.config.lock().unwrap().peer.clone().unwrap_or_default() + } + + fn set_peers(&self, peers: Vec) { + self.config.lock().unwrap().peer = Some(peers); + } + + fn get_listeners(&self) -> Option> { + self.config.lock().unwrap().listeners.clone() + } + + fn set_listeners(&self, listeners: Vec) { + self.config.lock().unwrap().listeners = Some(listeners); + } + + fn get_mapped_listeners(&self) -> Vec { + self.config + .lock() + .unwrap() + .mapped_listeners + .clone() + .unwrap_or_default() + } + + fn set_mapped_listeners(&self, listeners: Option>) { + self.config.lock().unwrap().mapped_listeners = listeners; + } + + fn get_vpn_portal_config(&self) -> Option { + self.config.lock().unwrap().vpn_portal_config.clone() + } + fn set_vpn_portal_config(&self, config: VpnPortalConfig) { + self.config.lock().unwrap().vpn_portal_config = Some(config); + } + + fn get_flags(&self) -> Flags { + self.config + .lock() + .unwrap() + .flags_struct + .clone() + .unwrap_or_default() + } + + fn set_flags(&self, flags: Flags) { + self.config.lock().unwrap().flags_struct = Some(flags); + } + + fn get_exit_nodes(&self) -> Vec { + self.config + .lock() + .unwrap() + .exit_nodes + .clone() + .unwrap_or_default() + } + + fn set_exit_nodes(&self, nodes: Vec) { + self.config.lock().unwrap().exit_nodes = Some(nodes); + } + + fn get_routes(&self) -> Option> { + self.config.lock().unwrap().routes.clone() + } + + fn set_routes(&self, routes: Option>) { + self.config.lock().unwrap().routes = routes; + } + + fn get_socks5_portal(&self) -> Option { + self.config.lock().unwrap().socks5_proxy.clone() + } + + fn set_socks5_portal(&self, addr: Option) { + self.config.lock().unwrap().socks5_proxy = addr; + } + + fn get_port_forwards(&self) -> Vec { + self.config + .lock() + .unwrap() + .port_forward + .clone() + .unwrap_or_default() + } + + fn set_port_forwards(&self, forwards: Vec) { + self.config.lock().unwrap().port_forward = Some(forwards); + } + + fn get_acl(&self) -> Option { + self.config.lock().unwrap().acl.clone() + } + + fn set_acl(&self, acl: Option) { + self.config.lock().unwrap().acl = acl; + } + + fn get_tcp_whitelist(&self) -> Vec { + self.config + .lock() + .unwrap() + .tcp_whitelist + .clone() + .unwrap_or_default() + } + + fn set_tcp_whitelist(&self, whitelist: Vec) { + self.config.lock().unwrap().tcp_whitelist = Some(whitelist); + } + + fn get_udp_whitelist(&self) -> Vec { + self.config + .lock() + .unwrap() + .udp_whitelist + .clone() + .unwrap_or_default() + } + + fn set_udp_whitelist(&self, whitelist: Vec) { + self.config.lock().unwrap().udp_whitelist = Some(whitelist); + } + + fn get_stun_servers(&self) -> Option> { + self.config.lock().unwrap().stun_servers.clone() + } + + fn set_stun_servers(&self, servers: Option>) { + self.config.lock().unwrap().stun_servers = servers; + } + + fn get_stun_servers_v6(&self) -> Option> { + self.config.lock().unwrap().stun_servers_v6.clone() + } + + fn set_stun_servers_v6(&self, servers: Option>) { + self.config.lock().unwrap().stun_servers_v6 = servers; + } + + fn get_secure_mode(&self) -> Option { + self.config.lock().unwrap().secure_mode.clone() + } + + fn set_secure_mode(&self, secure_mode: Option) { + self.config.lock().unwrap().secure_mode = secure_mode; + } + + fn get_credential_file(&self) -> Option { + self.config.lock().unwrap().credential_file.clone() + } + + fn set_credential_file(&self, path: Option) { + self.config.lock().unwrap().credential_file = path; + } + + fn get_network_config_source(&self) -> ConfigSource { + self.config + .lock() + .unwrap() + .source + .as_ref() + .map(|source| source.source) + .unwrap_or(ConfigSource::User) + } + + fn set_network_config_source(&self, source: Option) { + self.config.lock().unwrap().source = source.and_then(|source| match source { + ConfigSource::User => None, + other => Some(ConfigSourceConfig { source: other }), + }); + } + + fn dump(&self) -> String { + #[cfg(feature = "config-write")] + { + let mut config = self.config.lock().unwrap().clone(); + Self::normalize_config_source(&mut config); + config.flags = Some(flags_diff_from_default(&self.get_flags())); + if config.stun_servers == Some(default_udp_stun_servers()) { + config.stun_servers = None; + } + if config.stun_servers_v6 == Some(default_udp_v6_stun_servers()) { + config.stun_servers_v6 = None; + } + toml::to_string_pretty(&config).unwrap() + } + #[cfg(not(feature = "config-write"))] + { + panic!("this build does not include TOML configuration serialization") + } + } +} + +/// Transitional name retained while native consumers migrate to [`TomlConfig`]. +pub type TomlConfigLoader = TomlConfig; + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn parse_error_preserves_source_and_location() { + let error = + TomlConfig::new_from_str_with_source("fixture.toml", "dhcp = \"yes\"").unwrap_err(); + let display = error.to_string(); + + assert!(display.contains("fixture.toml")); + assert!(display.contains("dhcp = \"yes\"")); + assert!(display.contains("invalid type: string")); + assert!( + error + .chain() + .any(|cause| cause.downcast_ref::().is_some()) + ); + } + + #[cfg(feature = "config-write")] + #[test] + fn toml_round_trip_preserves_config_and_non_default_flags() { + let config = TomlConfig::new_from_str( + r#" +instance_name = "node-a" +instance_id = "018f85a8-a9d0-7d4c-b73d-4ab62c048a20" +hostname = "host-a" +listeners = ["tcp://0.0.0.0:11010"] + +[network_identity] +network_name = "network-a" +network_secret = "secret-a" + +[flags] +mtu = 1420 +socket_mark = 0 +"#, + ) + .unwrap(); + + let dumped = config.dump(); + let restored = TomlConfig::new_from_str(&dumped).unwrap(); + + assert_eq!(restored.get_id(), config.get_id()); + assert_eq!(restored.get_hostname(), "host-a"); + assert_eq!( + restored.get_network_identity(), + config.get_network_identity() + ); + assert_eq!(restored.get_listener_uris(), config.get_listener_uris()); + assert_eq!(restored.get_flags().mtu, 1420); + assert_eq!(restored.get_flags().socket_mark, Some(0)); + } + + #[test] + fn hostname_normalization_is_portable_and_has_no_host_fallback() { + let absent = TomlConfig::default(); + assert_eq!(absent.get_hostname(), ""); + + let configured = TomlConfig::new_from_str("hostname = \"node\\u0007-name\"").unwrap(); + assert_eq!(configured.get_hostname(), "node-name"); + } + + #[test] + fn credential_mode_does_not_synthesize_a_network_secret() { + let config = TomlConfig::new_from_str( + r#" +[network_identity] +network_name = "credential-network" + +[secure_mode] +enabled = true +"#, + ) + .unwrap(); + + let identity = config.get_network_identity(); + assert_eq!(identity.network_name, "credential-network"); + assert_eq!(identity.network_secret, None); + } + + #[cfg(feature = "config-write")] + #[test] + fn user_source_is_implicit_while_web_source_round_trips() { + let user = TomlConfig::new_from_str( + r#" +[source] +source = "user" +"#, + ) + .unwrap(); + assert_eq!(user.get_network_config_source(), ConfigSource::User); + assert!(!user.dump().contains("[source]")); + + let web = TomlConfig::new_from_str( + r#" +[source] +source = "web" +"#, + ) + .unwrap(); + assert_eq!(web.get_network_config_source(), ConfigSource::Web); + assert!(web.dump().contains("source = \"web\"")); + } +} + +#[cfg(test)] +mod compatibility_tests { + use super::*; + + #[cfg(feature = "config-write")] + #[test] + fn flags_diff_covers_every_protobuf_field() { + use prost::Message as _; + + let descriptor_set = + prost_types::FileDescriptorSet::decode(crate::proto::DESCRIPTOR_POOL_BYTES).unwrap(); + let proto_fields = descriptor_set + .file + .iter() + .find(|file| file.package.as_deref() == Some("common")) + .and_then(|file| { + file.message_type + .iter() + .find(|message| message.name.as_deref() == Some("FlagsInConfig")) + }) + .unwrap() + .field + .iter() + .map(|field| field.name.as_deref().unwrap()) + .collect::>(); + let diff_fields = FLAGS_DIFF_FIELDS + .iter() + .copied() + .collect::>(); + + assert_eq!(diff_fields, proto_fields); + } + + #[test] + fn socket_mark_config_file_roundtrip_none_some_and_zero() { + // Omitting the flag leaves socket_mark unset (None) -> SO_MARK untouched. + let cfg = TomlConfigLoader::new_from_str( + r#" +[network_identity] +network_name = "n" +network_secret = "s" +"#, + ) + .unwrap(); + assert_eq!(cfg.get_flags().socket_mark, None); + + // socket_mark = 0 is a legitimate value distinct from "unset". + let cfg = TomlConfigLoader::new_from_str( + r#" +[network_identity] +network_name = "n" +network_secret = "s" + +[flags] +socket_mark = 0 +"#, + ) + .unwrap(); + assert_eq!(cfg.get_flags().socket_mark, Some(0)); + + // A non-zero mark round-trips as Some(v). + let cfg = TomlConfigLoader::new_from_str( + r#" +[network_identity] +network_name = "n" +network_secret = "s" + +[flags] +socket_mark = 66 +"#, + ) + .unwrap(); + assert_eq!(cfg.get_flags().socket_mark, Some(66)); + + // set_flags(None) must serialize back through gen_config without + // resurrecting a value (guards the gen_flags merge against dropping + // the key when the serialized default is null). + cfg.set_flags(Flags { + socket_mark: None, + ..cfg.get_flags() + }); + assert_eq!(cfg.get_flags().socket_mark, None); + } + + #[cfg(feature = "config-write")] + #[test] + fn dump_preserves_flags_that_differ_from_easytier_defaults() { + let cfg = TomlConfigLoader::default(); + let mut flags = gen_default_flags(); + flags.dev_name = "et_test".to_string(); + flags.enable_quic_proxy = true; + flags.disable_tcp_hole_punching = true; + flags.disable_sym_hole_punching = true; + flags.multi_thread = false; + flags.bind_device = false; + flags.enable_ipv6 = false; + flags.relay_network_whitelist = "".to_string(); + flags.mtu = 0; + flags.foreign_relay_bps_limit = u64::MAX - 1; + flags.instance_recv_bps_limit = u64::MAX - 2; + flags.data_compress_algo = CompressionAlgoPb::Zstd.into(); + flags.socket_mark = Some(0); + cfg.set_flags(flags); + + let dumped = cfg.dump(); + + assert!(dumped.contains("dev_name = \"et_test\"")); + assert!(dumped.contains("enable_quic_proxy = true")); + assert!(dumped.contains("disable_tcp_hole_punching = true")); + assert!(dumped.contains("disable_sym_hole_punching = true")); + assert!(dumped.contains("multi_thread = false")); + assert!(dumped.contains("bind_device = false")); + assert!(dumped.contains("enable_ipv6 = false")); + assert!(dumped.contains("relay_network_whitelist = \"\"")); + assert!(dumped.contains("mtu = 0")); + assert!(dumped.contains("foreign_relay_bps_limit = \"18446744073709551614\"")); + assert!(dumped.contains("instance_recv_bps_limit = \"18446744073709551613\"")); + assert!(dumped.contains("data_compress_algo = \"Zstd\"")); + assert!(dumped.contains("socket_mark = 0")); + + let reloaded = TomlConfigLoader::new_from_str(&dumped).unwrap(); + let reloaded_flags = reloaded.get_flags(); + assert_eq!(reloaded_flags.dev_name, "et_test"); + assert!(reloaded_flags.enable_quic_proxy); + assert!(reloaded_flags.disable_tcp_hole_punching); + assert!(reloaded_flags.disable_sym_hole_punching); + assert!(!reloaded_flags.multi_thread); + assert!(!reloaded_flags.bind_device); + assert!(!reloaded_flags.enable_ipv6); + assert_eq!(reloaded_flags.relay_network_whitelist, ""); + assert_eq!(reloaded_flags.mtu, 0); + assert_eq!(reloaded_flags.foreign_relay_bps_limit, u64::MAX - 1); + assert_eq!(reloaded_flags.instance_recv_bps_limit, u64::MAX - 2); + assert_eq!( + reloaded_flags.data_compress_algo, + i32::from(CompressionAlgoPb::Zstd) + ); + assert_eq!(reloaded_flags.socket_mark, Some(0)); + } + + #[test] + fn test_stun_servers_config() { + let config = TomlConfigLoader::default(); + let stun_servers = config.get_stun_servers(); + assert!(stun_servers.is_none()); + + // Test setting custom stun servers + let custom_servers = vec!["txt:stun.easytier.cn".to_string()]; + config.set_stun_servers(Some(custom_servers.clone())); + + let retrieved_servers = config.get_stun_servers(); + assert_eq!(retrieved_servers.unwrap(), custom_servers); + } + + #[test] + fn test_stun_servers_toml_parsing() { + let config_str = r#" +instance_name = "test" +stun_servers = [ + "stun.l.google.com:19302", + "stun1.l.google.com:19302", + "txt:stun.easytier.cn" +]"#; + + let config = TomlConfigLoader::new_from_str(config_str).unwrap(); + let stun_servers = config.get_stun_servers().unwrap(); + + assert_eq!(stun_servers.len(), 3); + assert_eq!(stun_servers[0], "stun.l.google.com:19302"); + assert_eq!(stun_servers[1], "stun1.l.google.com:19302"); + assert_eq!(stun_servers[2], "txt:stun.easytier.cn"); + } + + #[cfg(feature = "config-write")] + #[test] + fn test_network_config_source_toml_roundtrip() { + let config = TomlConfigLoader::default(); + assert_eq!(config.get_network_config_source(), ConfigSource::User); + + config.set_network_config_source(Some(ConfigSource::Web)); + let dumped = config.dump(); + + assert!(dumped.contains("[source]")); + assert!(dumped.contains("source = \"web\"")); + + let loaded = TomlConfigLoader::new_from_str(&dumped).unwrap(); + assert_eq!(loaded.get_network_config_source(), ConfigSource::Web); + } + + #[cfg(feature = "config-write")] + #[test] + fn test_toml_credential_mode_omits_network_secret() { + for network_secret in ["", r#"network_secret = """#] { + let config = TomlConfigLoader::new_from_str(&format!( + r#" +[network_identity] +network_name = "credential-network" +{network_secret} + +[secure_mode] +enabled = true +"# + )) + .unwrap(); + + let identity = config.get_network_identity(); + assert_eq!(identity.network_name, "credential-network"); + assert_eq!(identity.network_secret, None); + assert_eq!(identity.network_secret_digest, None); + assert!(!config.dump().contains("network_secret")); + } + } + + #[test] + fn test_toml_secure_mode_without_network_identity_uses_default_secret() { + let config = TomlConfigLoader::new_from_str( + r#" +[secure_mode] +enabled = true +"#, + ) + .unwrap(); + + let identity = config.get_network_identity(); + assert_eq!(identity.network_name, "default"); + assert_eq!(identity.network_secret.as_deref(), Some("")); + assert!(identity.network_secret_digest.is_some()); + } + + #[test] + fn test_acl_toml_rule_uses_defaults_for_omitted_fields() { + use crate::proto::acl::{Action, ChainType, Protocol}; + + let config_str = r#" +[[acl.acl_v1.chains]] +name = "subnet_proxy_protect" +chain_type = 3 +enabled = true +default_action = 2 + +[[acl.acl_v1.chains.rules]] +name = "allow_my_devices" +priority = 1000 +action = 1 +source_ips = ["10.172.192.2/32"] +protocol = 5 +enabled = true +"#; + + let config = TomlConfigLoader::new_from_str(config_str).unwrap(); + let acl = config.get_acl().unwrap(); + let acl_v1 = acl.acl_v1.unwrap(); + let chain = &acl_v1.chains[0]; + let rule = &chain.rules[0]; + + assert_eq!(chain.chain_type, ChainType::Forward as i32); + assert_eq!(chain.default_action, Action::Drop as i32); + assert_eq!(rule.action, Action::Allow as i32); + assert_eq!(rule.protocol, Protocol::Any as i32); + assert_eq!(rule.source_ips, vec!["10.172.192.2/32"]); + assert!(rule.ports.is_empty()); + assert!(rule.source_ports.is_empty()); + assert!(rule.destination_ips.is_empty()); + assert!(rule.source_groups.is_empty()); + assert!(rule.destination_groups.is_empty()); + assert_eq!(rule.rate_limit, 0); + assert_eq!(rule.burst_limit, 0); + assert!(!rule.stateful); + } + + #[test] + fn test_acl_toml_group_can_omit_declares_or_members() { + let declares_only = r#" +[acl.acl_v1.group] + +[[acl.acl_v1.group.declares]] +group_name = "admin" +group_secret = "admin-pw" +"#; + let config = TomlConfigLoader::new_from_str(declares_only).unwrap(); + let group = config.get_acl().unwrap().acl_v1.unwrap().group.unwrap(); + assert_eq!(group.declares.len(), 1); + assert!(group.members.is_empty()); + + let members_only = r#" +[acl.acl_v1.group] +members = ["admin"] +"#; + let config = TomlConfigLoader::new_from_str(members_only).unwrap(); + let group = config.get_acl().unwrap().acl_v1.unwrap().group.unwrap(); + assert!(group.declares.is_empty()); + assert_eq!(group.members, vec!["admin"]); + } + + #[cfg(feature = "config-write")] + #[test] + fn test_network_config_source_user_is_implicit() { + let config = TomlConfigLoader::default(); + config.set_network_config_source(Some(ConfigSource::User)); + let dumped = config.dump(); + + assert!(!dumped.contains("[source]")); + + let loaded = TomlConfigLoader::new_from_str(&dumped).unwrap(); + assert_eq!(loaded.get_network_config_source(), ConfigSource::User); + + let explicit_user = TomlConfigLoader::new_from_str( + r#" +[source] +source = "user" +"#, + ) + .unwrap(); + assert_eq!( + explicit_user.get_network_config_source(), + ConfigSource::User + ); + assert!(!explicit_user.dump().contains("[source]")); + } + + #[cfg(feature = "config-write")] + #[test] + fn test_ipv6_public_addr_config_roundtrip() { + let config = TomlConfigLoader::default(); + let prefix: cidr::Ipv6Cidr = "2001:db8:100::/64".parse().unwrap(); + + config.set_ipv6_public_addr_provider(true); + config.set_ipv6_public_addr_auto(true); + config.set_ipv6_public_addr_prefix(Some(prefix)); + + assert!(config.get_ipv6_public_addr_provider()); + assert!(config.get_ipv6_public_addr_auto()); + assert_eq!(config.get_ipv6_public_addr_prefix(), Some(prefix)); + + let dumped = config.dump(); + let loaded = TomlConfigLoader::new_from_str(&dumped).unwrap(); + assert!(loaded.get_ipv6_public_addr_provider()); + assert!(loaded.get_ipv6_public_addr_auto()); + assert_eq!(loaded.get_ipv6_public_addr_prefix(), Some(prefix)); + } +} + +#[cfg(test)] +mod full_example_tests { + use super::*; + + #[cfg(feature = "config-write")] + #[test] + fn full_example_test() { + let config_str = r#" +instance_name = "default" +instance_id = "87ede5a2-9c3d-492d-9bbe-989b9d07e742" +ipv4 = "10.144.144.10" +listeners = [ "tcp://0.0.0.0:11010", "udp://0.0.0.0:11010" ] +routes = [ "192.168.0.0/16" ] + +[network_identity] +network_name = "default" +network_secret = "" + +[[peer]] +uri = "tcp://public.kkrainbow.top:11010" + +[[peer]] +uri = "udp://192.168.94.33:11010" + +[[proxy_network]] +cidr = "10.147.223.0/24" +allow = ["tcp", "udp", "icmp"] + +[[proxy_network]] +cidr = "10.1.1.0/24" +allow = ["tcp", "icmp"] + +[file_logger] +level = "info" +file = "easytier" +dir = "/tmp/easytier" + +[console_logger] +level = "warn" + +[[port_forward]] +bind_addr = "0.0.0.0:11011" +dst_addr = "192.168.94.33:11011" +proto = "tcp" +"#; + let ret = TomlConfigLoader::new_from_str(config_str); + if let Err(e) = &ret { + println!("{}", e); + } else { + println!("{:?}", ret.as_ref().unwrap()); + } + assert!(ret.is_ok()); + + let ret = ret.unwrap(); + assert_eq!("10.144.144.10/24", ret.get_ipv4().unwrap().to_string()); + + assert_eq!( + vec!["tcp://0.0.0.0:11010", "udp://0.0.0.0:11010"], + ret.get_listener_uris() + .iter() + .map(|u| u.to_string()) + .collect::>() + ); + + assert_eq!( + vec![PortForwardConfig { + bind_addr: "0.0.0.0:11011".parse().unwrap(), + dst_addr: "192.168.94.33:11011".parse().unwrap(), + proto: "tcp".to_string(), + }], + ret.get_port_forwards() + ); + println!("{}", ret.dump()); + } +} + +#[cfg(test)] +mod diagnostic_compatibility_tests { + use super::*; + + #[test] + fn stdin_source_name_and_caret_are_preserved() { + let error = TomlConfig::new_from_str_with_source("stdin", "dhcp = \"yes\"") + .unwrap_err() + .to_string(); + + assert!(error.contains("stdin")); + assert!(error.contains("dhcp = \"yes\"")); + assert!(error.contains('^')); + assert!(!error.contains("")); + } + + #[test] + fn non_ascii_before_typed_error_keeps_byte_location() { + let error = TomlConfig::new_from_str("hostname = \"节点\"\ndhcp = \"yes\"") + .unwrap_err() + .to_string(); + + assert!(error.contains("dhcp = \"yes\"")); + assert!(error.contains('^')); + assert!(error.contains("invalid type: string")); + } + + #[cfg(feature = "rich-config-errors")] + #[test] + fn non_ascii_on_syntax_error_line_keeps_source_location() { + let error = TomlConfig::new_from_str("hostname = \"节点\" dhcp = \"yes\"") + .unwrap_err() + .to_string(); + + assert!(error.contains("inline config:1:")); + assert!(error.contains("hostname = \"节点\" dhcp = \"yes\"")); + assert!(error.contains("expected newline")); + assert!(!error.contains("")); + } + + #[test] + fn flags_conversion_error_keeps_source_and_cause_chain() { + let error = TomlConfig::new_from_str_with_source( + "flags-fixture.toml", + "[flags]\nsocket_mark = \"bad\"", + ) + .unwrap_err(); + let display = error.to_string(); + + assert!(display.contains("flags-fixture.toml")); + assert!(display.contains("failed to load config")); + assert!(display.contains("failed to parse flags")); + assert!( + error + .chain() + .any(|cause| cause.to_string().contains("failed to parse flags")) + ); + } +} diff --git a/easytier-core/src/config/toml/snapshot.rs b/easytier-core/src/config/toml/snapshot.rs new file mode 100644 index 00000000..f3bb7e90 --- /dev/null +++ b/easytier-core/src/config/toml/snapshot.rs @@ -0,0 +1,15 @@ +use super::{Arc, Mutex, TomlConfig}; + +impl TomlConfig { + pub(crate) fn detached_snapshot(&self) -> Self { + let config = self.config.lock().unwrap().clone(); + Self { + config: Arc::new(Mutex::new(config)), + } + } + + pub(crate) fn replace_from_snapshot(&self, snapshot: &Self) { + let config = snapshot.config.lock().unwrap().clone(); + *self.config.lock().unwrap() = config; + } +} diff --git a/easytier-core/src/connectivity/composite.rs b/easytier-core/src/connectivity/composite.rs new file mode 100644 index 00000000..c4af43a6 --- /dev/null +++ b/easytier-core/src/connectivity/composite.rs @@ -0,0 +1,474 @@ +//! Composes process-wide socket capabilities with instance-scoped network facts. + +use std::{ + future::Future, + net::{IpAddr, Ipv6Addr, SocketAddr}, + sync::Arc, + time::{Duration, Instant}, +}; + +use async_trait::async_trait; +use url::Url; + +use crate::{ + connectivity::{ + direct::DirectConnectorHost, + manual::{ManualConnectorHost, ManualInterfaceAddrs}, + transport::ConnectedByteStream, + }, + proto::peer_rpc::GetIpListResponse, + socket::{ + SocketContext, + tcp::{ + TcpConnectOptions, TcpListenOptions, VirtualTcpListenerFactory, VirtualTcpSocketFactory, + }, + udp::{PreferredIpv6Source, UdpBindOptions, VirtualUdpSocketFactory}, + }, +}; + +const INTERFACE_ADDR_CACHE_TTL: Duration = Duration::from_secs(60); + +#[derive(Clone)] +struct CachedInterfaceAddrs { + collected_at: Instant, + response: GetIpListResponse, +} + +struct InterfaceAddrCacheEntry { + context: SocketContext, + value: Arc>>, +} + +struct InterfaceAddrCache { + entries: tokio::sync::Mutex>, + ttl: Duration, +} + +impl InterfaceAddrCache { + fn new(ttl: Duration) -> Self { + Self { + entries: tokio::sync::Mutex::new(Vec::new()), + ttl, + } + } + + async fn get_or_collect(&self, context: &SocketContext, collect: F) -> GetIpListResponse + where + F: FnOnce() -> Fut, + Fut: Future, + { + let value = { + let mut entries = self.entries.lock().await; + entries.retain(|entry| { + if Arc::strong_count(&entry.value) > 1 { + return true; + } + entry.value.try_lock().map_or(true, |cached| { + cached + .as_ref() + .is_some_and(|cached| cached.collected_at.elapsed() < self.ttl) + }) + }); + if let Some(entry) = entries.iter().find(|entry| &entry.context == context) { + entry.value.clone() + } else { + let value = Arc::new(tokio::sync::Mutex::new(None)); + entries.push(InterfaceAddrCacheEntry { + context: context.clone(), + value: value.clone(), + }); + value + } + }; + + // Only collectors for the same socket context share this lock. A slow + // namespace observation cannot block fresh hits or refreshes elsewhere. + let mut cached = value.lock().await; + if let Some(cached) = cached + .as_ref() + .filter(|cached| cached.collected_at.elapsed() < self.ttl) + { + return cached.response.clone(); + } + + let response = collect().await; + *cached = Some(CachedInterfaceAddrs { + collected_at: Instant::now(), + response: response.clone(), + }); + response + } +} + +/// Mechanical connector operations supplied by one process-wide runtime. +#[async_trait] +pub trait ConnectorRuntime: VirtualTcpSocketFactory + Send + Sync + 'static { + async fn connect_byte_stream( + &self, + url: &Url, + ) -> anyhow::Result>; + + async fn local_addr_for_remote( + &self, + remote_addr: SocketAddr, + context: SocketContext, + ) -> anyhow::Result; + + async fn collect_ip_addrs(&self, context: &SocketContext) -> GetIpListResponse; + + async fn preferred_ipv6_source( + &self, + ip: Ipv6Addr, + context: SocketContext, + ) -> Option; +} + +/// Instance facts consumed by portable connector policy. +/// +/// Socket creation, route probing and host interface I/O are deliberately +/// absent; those belong to the process-wide [`ConnectorRuntime`]. +pub trait ConnectorEnvironment: Send + Sync + 'static { + fn socket_context(&self) -> SocketContext; + + fn mapped_listeners(&self) -> Vec; + fn is_local_ip(&self, ip: &IpAddr) -> bool; +} + +/// Deep adapter that combines one socket runtime with one instance environment. +pub struct ConnectorHostAdapter { + sockets: Arc, + environment: Arc, + interface_addrs: InterfaceAddrCache, +} + +impl ConnectorHostAdapter { + pub fn new(sockets: Arc, environment: Arc) -> Self { + Self { + sockets, + environment, + interface_addrs: InterfaceAddrCache::new(INTERFACE_ADDR_CACHE_TTL), + } + } +} + +impl ConnectorHostAdapter +where + S: ConnectorRuntime, +{ + async fn cached_ip_addrs(&self, context: &SocketContext) -> GetIpListResponse { + self.interface_addrs + .get_or_collect(context, || self.sockets.collect_ip_addrs(context)) + .await + } +} + +#[async_trait] +impl VirtualTcpSocketFactory for ConnectorHostAdapter +where + S: VirtualTcpSocketFactory, + E: Send + Sync + 'static, +{ + type Socket = S::Socket; + + async fn connect_tcp(&self, options: TcpConnectOptions) -> anyhow::Result { + self.sockets.connect_tcp(options).await + } +} + +#[async_trait] +impl VirtualTcpListenerFactory for ConnectorHostAdapter +where + S: VirtualTcpListenerFactory, + E: Send + Sync + 'static, +{ + type Listener = S::Listener; + + async fn bind_tcp(&self, options: TcpListenOptions) -> anyhow::Result> { + self.sockets.bind_tcp(options).await + } +} + +#[async_trait] +impl VirtualUdpSocketFactory for ConnectorHostAdapter +where + S: VirtualUdpSocketFactory, + E: Send + Sync + 'static, +{ + type Socket = S::Socket; + + async fn bind_udp(&self, options: UdpBindOptions) -> anyhow::Result> { + self.sockets.bind_udp(options).await + } +} + +#[async_trait] +impl ManualConnectorHost for ConnectorHostAdapter +where + S: ConnectorRuntime + VirtualUdpSocketFactory, + E: ConnectorEnvironment, +{ + async fn local_addr_for_remote( + &self, + remote_addr: SocketAddr, + context: SocketContext, + ) -> anyhow::Result { + self.sockets + .local_addr_for_remote(remote_addr, context) + .await + } + + async fn interface_addrs(&self) -> anyhow::Result { + let addrs = self + .cached_ip_addrs(&self.environment.socket_context()) + .await; + Ok(ManualInterfaceAddrs { + interface_ipv4s: addrs + .interface_ipv4s + .into_iter() + .map(std::net::Ipv4Addr::from) + .collect(), + interface_ipv6s: addrs + .interface_ipv6s + .into_iter() + .map(std::net::Ipv6Addr::from) + .collect(), + public_ipv6: addrs.public_ipv6.map(std::net::Ipv6Addr::from), + }) + } + + async fn connect_byte_stream( + &self, + url: &Url, + ) -> anyhow::Result::Socket>> { + self.sockets.connect_byte_stream(url).await + } +} + +#[async_trait] +impl DirectConnectorHost for ConnectorHostAdapter +where + S: ConnectorRuntime + VirtualUdpSocketFactory, + E: ConnectorEnvironment, +{ + async fn collect_ip_addrs(&self, context: &SocketContext) -> GetIpListResponse { + self.cached_ip_addrs(context).await + } + + async fn collect_foreign_ip_addrs(&self, context: &SocketContext) -> GetIpListResponse { + self.cached_ip_addrs(context).await + } + + fn mapped_listeners(&self) -> Vec { + self.environment.mapped_listeners() + } + + fn is_local_ip(&self, ip: &IpAddr) -> bool { + self.environment.is_local_ip(ip) + } + + async fn preferred_ipv6_source( + &self, + ip: Ipv6Addr, + context: SocketContext, + ) -> Option { + if !valid_public_ipv6_candidate(ip) { + return None; + } + self.sockets.preferred_ipv6_source(ip, context).await + } + + async fn preferred_foreign_ipv6_source( + &self, + ip: Ipv6Addr, + context: SocketContext, + ) -> Option { + if !valid_public_ipv6_candidate(ip) { + return None; + } + self.sockets.preferred_ipv6_source(ip, context).await + } +} + +fn valid_public_ipv6_candidate(ip: Ipv6Addr) -> bool { + !(ip.is_loopback() + || ip.is_unspecified() + || ip.is_unique_local() + || ip.is_unicast_link_local() + || ip.is_multicast()) +} + +#[cfg(test)] +mod tests { + use std::sync::atomic::{AtomicUsize, Ordering}; + + use tokio::sync::oneshot; + + use super::*; + use crate::socket::{IpVersion, NetNamespace}; + + #[tokio::test] + async fn concurrent_cache_misses_share_one_collection() { + let cache = Arc::new(InterfaceAddrCache::new(Duration::from_secs(60))); + let context = SocketContext::default(); + let calls = Arc::new(AtomicUsize::new(0)); + let (started_tx, started_rx) = oneshot::channel(); + let (release_tx, release_rx) = oneshot::channel(); + + let first = tokio::spawn({ + let cache = cache.clone(); + let context = context.clone(); + let calls = calls.clone(); + async move { + cache + .get_or_collect(&context, || async move { + calls.fetch_add(1, Ordering::SeqCst); + let _ = started_tx.send(()); + let _ = release_rx.await; + GetIpListResponse::default() + }) + .await + } + }); + started_rx.await.unwrap(); + + let (second_started_tx, second_started_rx) = oneshot::channel(); + let second = tokio::spawn({ + let cache = cache.clone(); + let context = context.clone(); + let calls = calls.clone(); + async move { + let _ = second_started_tx.send(()); + cache + .get_or_collect(&context, || async move { + calls.fetch_add(1, Ordering::SeqCst); + GetIpListResponse::default() + }) + .await + } + }); + second_started_rx.await.unwrap(); + tokio::task::yield_now().await; + let _ = release_tx.send(()); + + first.await.unwrap(); + second.await.unwrap(); + assert_eq!(calls.load(Ordering::SeqCst), 1); + } + + #[tokio::test] + async fn cache_keys_include_the_complete_socket_context() { + let cache = InterfaceAddrCache::new(Duration::from_secs(60)); + let calls = AtomicUsize::new(0); + let contexts = [ + SocketContext::default(), + SocketContext::default().with_ip_version(IpVersion::V4), + SocketContext::default().with_socket_mark(Some(7)), + SocketContext::default().with_netns(Some(NetNamespace::new("foreign-a"))), + ]; + + for context in &contexts { + cache + .get_or_collect(context, || async { + calls.fetch_add(1, Ordering::SeqCst); + GetIpListResponse::default() + }) + .await; + } + cache + .get_or_collect(&contexts[0], || async { + calls.fetch_add(1, Ordering::SeqCst); + GetIpListResponse::default() + }) + .await; + + assert_eq!(calls.load(Ordering::SeqCst), contexts.len()); + } + + #[tokio::test] + async fn slow_collection_does_not_block_a_different_context() { + let cache = Arc::new(InterfaceAddrCache::new(Duration::from_secs(60))); + let slow_context = SocketContext::default().with_socket_mark(Some(1)); + let other_context = SocketContext::default().with_socket_mark(Some(2)); + let (started_tx, started_rx) = oneshot::channel(); + let (release_tx, release_rx) = oneshot::channel(); + let slow = tokio::spawn({ + let cache = cache.clone(); + async move { + cache + .get_or_collect(&slow_context, || async move { + let _ = started_tx.send(()); + let _ = release_rx.await; + GetIpListResponse::default() + }) + .await + } + }); + started_rx.await.unwrap(); + + tokio::time::timeout( + Duration::from_millis(100), + cache.get_or_collect(&other_context, || async { GetIpListResponse::default() }), + ) + .await + .expect("different socket contexts must collect independently"); + + let _ = release_tx.send(()); + slow.await.unwrap(); + } + + #[tokio::test] + async fn slow_collection_does_not_block_a_fresh_hit() { + let cache = Arc::new(InterfaceAddrCache::new(Duration::from_secs(60))); + let cached_context = SocketContext::default().with_socket_mark(Some(1)); + let slow_context = SocketContext::default().with_socket_mark(Some(2)); + cache + .get_or_collect(&cached_context, || async { GetIpListResponse::default() }) + .await; + + let (started_tx, started_rx) = oneshot::channel(); + let (release_tx, release_rx) = oneshot::channel(); + let slow = tokio::spawn({ + let cache = cache.clone(); + async move { + cache + .get_or_collect(&slow_context, || async move { + let _ = started_tx.send(()); + let _ = release_rx.await; + GetIpListResponse::default() + }) + .await + } + }); + started_rx.await.unwrap(); + + tokio::time::timeout( + Duration::from_millis(100), + cache.get_or_collect(&cached_context, || async { + panic!("fresh cache hit must not recollect") + }), + ) + .await + .expect("fresh hit must not wait for another socket context"); + + let _ = release_tx.send(()); + slow.await.unwrap(); + } + + #[tokio::test] + async fn expired_entries_are_recollected() { + let cache = InterfaceAddrCache::new(Duration::ZERO); + let calls = AtomicUsize::new(0); + let context = SocketContext::default(); + + for _ in 0..2 { + cache + .get_or_collect(&context, || async { + calls.fetch_add(1, Ordering::SeqCst); + GetIpListResponse::default() + }) + .await; + } + + assert_eq!(calls.load(Ordering::SeqCst), 2); + } +} diff --git a/easytier-core/src/connectivity/connector_host.rs b/easytier-core/src/connectivity/connector_host.rs new file mode 100644 index 00000000..aaa2cabe --- /dev/null +++ b/easytier-core/src/connectivity/connector_host.rs @@ -0,0 +1,708 @@ +//! Host-operation Adapter for the shared connector composition. +//! +//! This module was previously named `host`, which collided with the +//! top-level [`crate::host`] layer, and one concept was split across two +//! same-named `environment.rs` files. The ownership split is: +//! +//! - Mechanical environment queries implement +//! [`crate::host::environment::HostConnectorEnvironmentIo`]. +//! - [`HostConnectorRuntime`] adapts host socket/listener factories and an +//! injected environment snapshot to the shared connector runtime and +//! environment traits. +//! - [`HostConnectorEnvironmentSnapshot`] is connectivity's captured view of +//! the host environment. + +use std::{ + net::{IpAddr, Ipv4Addr, Ipv6Addr, SocketAddr}, + sync::Arc, +}; + +use async_trait::async_trait; +use serde::{Deserialize, Serialize}; +use url::Url; + +use crate::{ + connectivity::{ + composite::{ConnectorEnvironment, ConnectorHostAdapter, ConnectorRuntime}, + transport::ConnectedByteStream, + }, + host::environment::{HostConnectorEnvironmentIo, local_addr_for_remote}, + host::socket::{ + HostSocketRuntime, HostTcpStream, + factory::{HostSocketBackend, HostSocketFactory}, + listener::{HostTcpListener, HostTcpListenerBackend, HostTcpListenerFactory}, + udp::HostUdpSocket, + }, + proto::peer_rpc::GetIpListResponse, + socket::{ + SocketContext, + tcp::{ + TcpConnectOptions, TcpListenOptions, VirtualTcpListenerFactory, VirtualTcpSocketFactory, + }, + udp::{PreferredIpv6Source, UdpBindOptions, VirtualUdpSocketFactory}, + }, +}; + +/// Host-observed facts consumed by core connector policy. +/// +/// The host normalizes this snapshot before constructing an instance. Core +/// owns all selection and connection policy applied to these facts. +#[derive(Debug, Clone, Default, PartialEq, Eq, Serialize, Deserialize)] +pub struct HostConnectorEnvironmentSnapshot { + pub public_ipv4: Option, + pub interface_ipv4s: Vec, + pub public_ipv6: Option, + pub interface_ipv6s: Vec, + pub mapped_listeners: Vec, + pub local_ips: Vec, + pub protected_tcp_ports: Vec, + pub preferred_ipv6_sources: Vec, +} + +impl HostConnectorEnvironmentSnapshot { + fn ip_list(&self) -> GetIpListResponse { + GetIpListResponse { + public_ipv4: self.public_ipv4.map(Into::into), + interface_ipv4s: self + .interface_ipv4s + .iter() + .copied() + .map(Into::into) + .collect(), + public_ipv6: self.public_ipv6.map(Into::into), + interface_ipv6s: self + .interface_ipv6s + .iter() + .copied() + .map(Into::into) + .collect(), + listeners: Default::default(), + } + } + + fn preferred_ipv6_source(&self, ip: Ipv6Addr) -> Option { + if ip.is_loopback() + || ip.is_unspecified() + || ip.is_unique_local() + || ip.is_unicast_link_local() + || ip.is_multicast() + { + return None; + } + self.preferred_ipv6_sources + .iter() + .find(|source| source.ip == ip) + .copied() + } +} + +/// One host handle domain capable of creating and operating connector sockets. +pub trait ConnectorHostSocketBackend: HostSocketBackend + HostTcpListenerBackend {} + +impl ConnectorHostSocketBackend for T where T: HostSocketBackend + HostTcpListenerBackend {} + +/// Adapts mechanical host sockets and captured environment state to the +/// shared connector composition. +pub struct HostConnectorRuntime +where + B: ConnectorHostSocketBackend, + E: HostConnectorEnvironmentIo, +{ + socket_runtime: HostSocketRuntime, + sockets: HostSocketFactory, + listeners: HostTcpListenerFactory, + environment: Arc, + environment_io: Arc, +} + +impl HostConnectorRuntime +where + B: ConnectorHostSocketBackend, + E: HostConnectorEnvironmentIo, +{ + pub fn new( + runtime: HostSocketRuntime, + backend: Arc, + environment: HostConnectorEnvironmentSnapshot, + environment_io: Arc, + ) -> Self { + Self { + socket_runtime: runtime.clone(), + sockets: HostSocketFactory::new(runtime.clone(), backend.clone()), + listeners: HostTcpListenerFactory::new(runtime, backend), + environment: Arc::new(environment), + environment_io, + } + } +} + +#[async_trait] +impl VirtualTcpSocketFactory for HostConnectorRuntime +where + B: ConnectorHostSocketBackend, + E: HostConnectorEnvironmentIo, +{ + type Socket = HostTcpStream; + + async fn connect_tcp(&self, options: TcpConnectOptions) -> anyhow::Result { + self.sockets.connect_tcp(options).await + } +} + +#[async_trait] +impl VirtualUdpSocketFactory for HostConnectorRuntime +where + B: ConnectorHostSocketBackend, + E: HostConnectorEnvironmentIo, +{ + type Socket = HostUdpSocket; + + async fn bind_udp(&self, options: UdpBindOptions) -> anyhow::Result> { + self.sockets.bind_udp(options).await + } +} + +#[async_trait] +impl VirtualTcpListenerFactory for HostConnectorRuntime +where + B: ConnectorHostSocketBackend, + E: HostConnectorEnvironmentIo, +{ + type Listener = HostTcpListener; + + async fn bind_tcp(&self, options: TcpListenOptions) -> anyhow::Result> { + self.listeners.bind_tcp(options).await + } +} + +#[async_trait] +impl ConnectorRuntime for HostConnectorRuntime +where + B: ConnectorHostSocketBackend, + E: HostConnectorEnvironmentIo, +{ + async fn connect_byte_stream( + &self, + url: &Url, + ) -> anyhow::Result> { + anyhow::bail!("host does not support external byte stream: {url}") + } + + async fn local_addr_for_remote( + &self, + remote_addr: SocketAddr, + context: SocketContext, + ) -> anyhow::Result { + local_addr_for_remote( + &self.socket_runtime, + self.environment_io.clone(), + remote_addr, + context, + ) + .await + } + + async fn collect_ip_addrs(&self, _context: &SocketContext) -> GetIpListResponse { + self.environment.ip_list() + } + + async fn preferred_ipv6_source( + &self, + ip: Ipv6Addr, + _context: SocketContext, + ) -> Option { + self.environment.preferred_ipv6_source(ip) + } +} + +impl ConnectorEnvironment for HostConnectorRuntime +where + B: ConnectorHostSocketBackend, + E: HostConnectorEnvironmentIo, +{ + fn socket_context(&self) -> SocketContext { + SocketContext::default() + } + + fn mapped_listeners(&self) -> Vec { + self.environment.mapped_listeners.clone() + } + + fn is_local_ip(&self, ip: &IpAddr) -> bool { + self.environment.local_ips.contains(ip) + } +} + +/// Shared connector policy composed with a host-operation runtime. +pub type ConnectorHost = + ConnectorHostAdapter, HostConnectorRuntime>; + +/// Builds the shared connector host over one host-operation backend. +pub fn new_connector_host( + socket_runtime: HostSocketRuntime, + backend: Arc, + environment: HostConnectorEnvironmentSnapshot, + environment_io: Arc, +) -> ConnectorHost +where + B: ConnectorHostSocketBackend, + E: HostConnectorEnvironmentIo, +{ + let runtime = Arc::new(HostConnectorRuntime::new( + socket_runtime, + backend, + environment, + environment_io, + )); + ConnectorHostAdapter::new(runtime.clone(), runtime) +} + +#[cfg(test)] +mod tests { + use std::{ + io, + sync::{ + Mutex, + atomic::{AtomicBool, Ordering}, + }, + task::Poll, + }; + + use crate::{ + connectivity::{ + direct::{DirectConnectorHost, DirectConnectorRpcHandler}, + hole_punch::tcp::TcpHolePunchHost, + manual::ManualConnectorHost, + stun::StunInfoProvider, + }, + host::socket::{ + HostOperationId, HostSocketHandle, HostSocketIo, HostTcpIo, + factory::{HostSocketFactoryIo, HostTcpConnectResult, HostUdpBindResult}, + listener::{HostTcpBindResult, HostTcpListenerIo}, + udp::{HostUdpDatagram, HostUdpIo}, + }, + proto::{ + common::StunInfo, + peer_rpc::{DirectConnectorRpc as _, GetIpListRequest}, + rpc_types::controller::BaseController, + }, + socket::udp::UdpSocketSendMeta, + }; + + use super::*; + + #[derive(Default)] + struct UnsupportedBackend { + udp_send_attempts: Mutex, SocketAddr, UdpSocketSendMeta)>>, + reject_preferred_source: AtomicBool, + } + + struct FixedStunProvider; + + #[async_trait] + impl StunInfoProvider for FixedStunProvider { + fn get_stun_info(&self) -> StunInfo { + StunInfo { + public_ip: vec!["198.51.100.7".to_owned(), "2001:db8::1".to_owned()], + ..Default::default() + } + } + + async fn get_udp_port_mapping(&self, _local_port: u16) -> anyhow::Result { + anyhow::bail!("unused by direct RPC projection test") + } + + async fn get_tcp_port_mapping(&self, _local_port: u16) -> anyhow::Result { + anyhow::bail!("unused by direct RPC projection test") + } + + fn update_stun_info(&self) {} + } + + fn unsupported() -> io::Result { + Err(io::ErrorKind::Unsupported.into()) + } + + impl HostSocketIo for UnsupportedBackend { + fn cancel_operation(&self, _operation: HostOperationId) -> io::Result<()> { + Ok(()) + } + + fn close(&self, _handle: HostSocketHandle) -> io::Result<()> { + Ok(()) + } + } + + impl HostTcpIo for UnsupportedBackend { + fn submit_read( + &self, + _handle: HostSocketHandle, + _operation: HostOperationId, + _capacity: usize, + ) -> io::Result<()> { + unsupported() + } + + fn take_read(&self, _operation: HostOperationId) -> Poll>> { + Poll::Ready(unsupported()) + } + + fn submit_write( + &self, + _handle: HostSocketHandle, + _operation: HostOperationId, + _source: &[u8], + ) -> io::Result<()> { + unsupported() + } + + fn take_write(&self, _operation: HostOperationId) -> Poll> { + Poll::Ready(unsupported()) + } + } + + impl HostUdpIo for UnsupportedBackend { + fn submit_recv( + &self, + _handle: HostSocketHandle, + _operation: HostOperationId, + _capacity: usize, + ) -> io::Result<()> { + unsupported() + } + + fn take_recv(&self, _operation: HostOperationId) -> Poll> { + Poll::Ready(unsupported()) + } + + fn try_send( + &self, + _handle: HostSocketHandle, + source: &[u8], + peer_addr: SocketAddr, + meta: UdpSocketSendMeta, + ) -> io::Result<()> { + self.udp_send_attempts + .lock() + .unwrap() + .push((source.to_vec(), peer_addr, meta)); + if self.reject_preferred_source.load(Ordering::Relaxed) && meta.src_ip.is_some() { + return Err(io::ErrorKind::AddrNotAvailable.into()); + } + Ok(()) + } + + fn submit_send_ready( + &self, + _handle: HostSocketHandle, + _operation: HostOperationId, + ) -> io::Result<()> { + unsupported() + } + + fn take_send_ready(&self, _operation: HostOperationId) -> Poll> { + Poll::Ready(unsupported()) + } + } + + impl HostSocketFactoryIo for UnsupportedBackend { + fn submit_tcp_connect( + &self, + _operation: HostOperationId, + _options: &TcpConnectOptions, + ) -> io::Result<()> { + unsupported() + } + + fn take_tcp_connect( + &self, + _operation: HostOperationId, + ) -> Poll> { + Poll::Ready(unsupported()) + } + + fn submit_udp_bind( + &self, + _operation: HostOperationId, + _options: &UdpBindOptions, + ) -> io::Result<()> { + unsupported() + } + + fn take_udp_bind( + &self, + _operation: HostOperationId, + ) -> Poll> { + Poll::Ready(unsupported()) + } + } + + impl HostTcpListenerIo for UnsupportedBackend { + fn submit_tcp_bind( + &self, + _operation: HostOperationId, + _options: &TcpListenOptions, + ) -> io::Result<()> { + unsupported() + } + + fn take_tcp_bind( + &self, + _operation: HostOperationId, + ) -> Poll> { + Poll::Ready(unsupported()) + } + + fn submit_tcp_accept( + &self, + _handle: HostSocketHandle, + _operation: HostOperationId, + ) -> io::Result<()> { + unsupported() + } + + fn take_tcp_accept( + &self, + _operation: HostOperationId, + ) -> Poll> { + Poll::Ready(unsupported()) + } + } + + #[derive(Default)] + struct TestEnvironmentIo { + local_requests: Mutex>, + ready: Mutex>, + } + + impl HostConnectorEnvironmentIo for TestEnvironmentIo { + fn submit_local_addr_for_remote( + &self, + operation: HostOperationId, + remote_addr: SocketAddr, + context: &SocketContext, + ) -> io::Result<()> { + self.local_requests + .lock() + .unwrap() + .push((remote_addr, context.clone())); + self.ready.lock().unwrap().push(operation); + Ok(()) + } + + fn take_local_addr_for_remote( + &self, + operation: HostOperationId, + ) -> Poll> { + let mut ready = self.ready.lock().unwrap(); + let Some(index) = ready.iter().position(|candidate| *candidate == operation) else { + return Poll::Pending; + }; + ready.swap_remove(index); + Poll::Ready(Ok("192.0.2.1:40100".parse().unwrap())) + } + + fn cancel_operation(&self, operation: HostOperationId) -> io::Result<()> { + self.ready + .lock() + .unwrap() + .retain(|candidate| *candidate != operation); + Ok(()) + } + } + + fn test_environment_snapshot() -> HostConnectorEnvironmentSnapshot { + HostConnectorEnvironmentSnapshot { + interface_ipv4s: vec!["192.0.2.1".parse().unwrap()], + public_ipv6: Some("2001:db8::1".parse().unwrap()), + interface_ipv6s: vec!["2001:db8::1".parse().unwrap()], + mapped_listeners: vec!["tcp://192.0.2.1:11010".parse().unwrap()], + local_ips: vec!["192.0.2.1".parse().unwrap()], + protected_tcp_ports: vec![11010], + preferred_ipv6_sources: vec![PreferredIpv6Source { + ip: "2001:db8::1".parse().unwrap(), + ifindex: 7, + }], + ..Default::default() + } + } + + fn assert_core_host() + where + H: DirectConnectorHost + TcpHolePunchHost, + { + } + + #[tokio::test] + async fn delegates_connector_environment_without_owning_policy() { + type TestHost = ConnectorHost; + assert_core_host::(); + + let environment_io = Arc::new(TestEnvironmentIo::default()); + let host = new_connector_host( + HostSocketRuntime::new(), + Arc::new(UnsupportedBackend::default()), + test_environment_snapshot(), + environment_io.clone(), + ); + let remote = "203.0.113.1:11010".parse().unwrap(); + let context = SocketContext::default().with_socket_mark(Some(7)); + let local = ManualConnectorHost::local_addr_for_remote(&host, remote, context.clone()) + .await + .unwrap(); + assert_eq!(local, "192.0.2.1:40100".parse().unwrap()); + assert_eq!( + *environment_io.local_requests.lock().unwrap(), + vec![(remote, context)] + ); + assert_eq!( + ManualConnectorHost::interface_addrs(&host) + .await + .unwrap() + .public_ipv6, + Some("2001:db8::1".parse().unwrap()) + ); + let byte_stream_error = + match ManualConnectorHost::connect_byte_stream(&host, &"ring://42".parse().unwrap()) + .await + { + Ok(_) => panic!("test environment should reject byte streams"), + Err(error) => error, + }; + assert_eq!( + byte_stream_error.to_string(), + "host does not support external byte stream: ring://42" + ); + assert_eq!( + DirectConnectorHost::mapped_listeners(&host), + vec!["tcp://192.0.2.1:11010".parse::().unwrap()] + ); + } + + #[tokio::test] + async fn direct_rpc_projects_host_observations_without_instance_policy() { + let host = Arc::new(new_connector_host( + HostSocketRuntime::new(), + Arc::new(UnsupportedBackend::default()), + test_environment_snapshot(), + Arc::new(TestEnvironmentIo::default()), + )); + let handler = DirectConnectorRpcHandler::new_with_stun( + host, + SocketContext::default().with_socket_mark(Some(7)), + Some(Arc::new(FixedStunProvider)), + ); + + let response = handler + .get_ip_list(BaseController::default(), GetIpListRequest {}) + .await + .unwrap(); + + assert_eq!( + response.interface_ipv4s, + vec![std::net::Ipv4Addr::new(192, 0, 2, 1).into()] + ); + assert_eq!( + response.interface_ipv6s, + vec!["2001:db8::1".parse::().unwrap().into()] + ); + assert_eq!( + response.public_ipv4, + Some("198.51.100.7".parse::().unwrap().into()) + ); + assert_eq!( + response.public_ipv6, + Some("2001:db8::1".parse::().unwrap().into()) + ); + assert_eq!( + response + .listeners + .into_iter() + .map(Url::from) + .collect::>(), + vec!["tcp://192.0.2.1:11010".parse::().unwrap()] + ); + } + + #[tokio::test] + async fn foreign_direct_rpc_preserves_parent_managed_ipv6_addresses() { + let host = Arc::new(new_connector_host( + HostSocketRuntime::new(), + Arc::new(UnsupportedBackend::default()), + test_environment_snapshot(), + Arc::new(TestEnvironmentIo::default()), + )); + let handler = DirectConnectorRpcHandler::new_for_foreign_network_with_stun( + host, + SocketContext::default().with_socket_mark(Some(7)), + Some(Arc::new(FixedStunProvider)), + ); + + let response = handler + .get_ip_list(BaseController::default(), GetIpListRequest {}) + .await + .unwrap(); + + assert_eq!( + response.interface_ipv6s, + vec!["2001:db8::1".parse::().unwrap().into()] + ); + assert_eq!( + response.public_ipv6, + Some("2001:db8::1".parse::().unwrap().into()) + ); + assert_eq!( + response.public_ipv4, + Some("198.51.100.7".parse::().unwrap().into()) + ); + } + + fn snapshot() -> HostConnectorEnvironmentSnapshot { + HostConnectorEnvironmentSnapshot { + public_ipv4: Some("198.51.100.1".parse().unwrap()), + interface_ipv4s: vec!["192.0.2.1".parse().unwrap()], + public_ipv6: Some("2001:db8::1".parse().unwrap()), + interface_ipv6s: vec!["2001:db8::2".parse().unwrap()], + mapped_listeners: vec!["tcp://198.51.100.1:11010".parse().unwrap()], + local_ips: vec!["192.0.2.1".parse().unwrap()], + protected_tcp_ports: vec![11010], + preferred_ipv6_sources: vec![ + PreferredIpv6Source { + ip: "2001:db8::2".parse().unwrap(), + ifindex: 7, + }, + PreferredIpv6Source { + ip: "fd00::1".parse().unwrap(), + ifindex: 8, + }, + ], + } + } + + #[test] + fn projects_normalized_snapshot() { + let snapshot = snapshot(); + assert_eq!( + serde_json::from_slice::( + &serde_json::to_vec(&snapshot).unwrap() + ) + .unwrap(), + snapshot + ); + assert_eq!( + snapshot.ip_list().interface_ipv4s, + vec![Ipv4Addr::new(192, 0, 2, 1).into()] + ); + assert_eq!( + snapshot.preferred_ipv6_source("2001:db8::2".parse().unwrap()), + Some(PreferredIpv6Source { + ip: "2001:db8::2".parse().unwrap(), + ifindex: 7, + }) + ); + assert_eq!( + snapshot.preferred_ipv6_source("fd00::1".parse().unwrap()), + None + ); + } +} diff --git a/easytier-core/src/connectivity/direct/mod.rs b/easytier-core/src/connectivity/direct/mod.rs new file mode 100644 index 00000000..e4325630 --- /dev/null +++ b/easytier-core/src/connectivity/direct/mod.rs @@ -0,0 +1,1302 @@ +use std::{ + collections::HashSet, + hash::Hash, + net::{IpAddr, Ipv4Addr, Ipv6Addr, SocketAddr}, + sync::{Arc, Weak}, + time::{Duration, Instant as StdInstant}, +}; + +use anyhow::Context; +use async_trait::async_trait; +use dashmap::DashMap; +use quanta::Instant; +use rand::Rng; +use serde::{Deserialize, Serialize}; +use tokio::task::JoinSet; +use url::{Host, Url}; + +use crate::{ + config::PeerId, + connectivity::hole_punch::policy::{should_background_p2p_with_peer, should_try_p2p_with_peer}, + connectivity::stun::{StunInfoProvider, StunSocketMapper}, + connectivity::{ + LocalListenerUrls, NoLocalListeners, + protocol::{ + ClientProtocolUpgrader, ProtocolTransport, protocol_transport, protocol_uses_udp, + }, + transport::{self, ConnectedTransport, UdpSessionMode}, + }, + foundation::task::{PeerTaskLauncher, PeerTaskManager}, + host::dns::DnsResolver, + peers::{ + conn::peer_conn::PeerConnId, foreign_network::ForeignNetworkRpcRegistrar, + peer_manager::PeerManagerCore, peer_rpc::PeerRpcManager, + }, + process_runtime::ProtectedTcpPortRegistry, + proto::{ + common::Void, + peer_rpc::{ + DirectConnectorRpc, DirectConnectorRpcClientFactory, + DirectConnectorRpcServer as GeneratedDirectConnectorRpcServer, GetIpListRequest, + GetIpListResponse, SendUdpHolePunchPacketRequest, + }, + rpc_types::{self, controller::BaseController}, + }, + socket::{ + IpVersion, SocketContext, + tcp::{TcpBindOptions, TcpSocketPurpose}, + udp::{ + PreferredIpv6Source, UdpBindOptions, VirtualUdpSocket, VirtualUdpSocketFactory, + send_v4_hole_punch_control_packet, send_v6_hole_punch_control_packet, + }, + }, + tunnel::Tunnel, +}; + +use super::manual::{ + ManualConnectorHost, collect_bind_addrs, convert_idn_to_ascii, resolve_remote_addr, + resolve_url_addrs, +}; + +mod udp; + +const DIRECT_CONNECTOR_BLACKLIST_TIMEOUT: Duration = Duration::from_secs(300); +const INVALID_SERVICE_BLACKLIST_TIMEOUT: Duration = Duration::from_secs(3600); +const DIRECT_CONNECT_TIMEOUT: Duration = Duration::from_secs(3); +const DIRECT_TASK_LOOP_INTERVAL_MS: u64 = 5000; +const MAX_IPV6_HOLE_PUNCH_CONNECTOR_ADDRS: usize = 16; +const MAX_UDP_HOLE_PUNCH_CONNECTOR_ADDRS: usize = 16; + +#[async_trait] +pub trait DirectConnectorHost: ManualConnectorHost { + async fn collect_ip_addrs(&self, context: &SocketContext) -> GetIpListResponse; + + async fn collect_foreign_ip_addrs(&self, context: &SocketContext) -> GetIpListResponse { + self.collect_ip_addrs(context).await + } + + fn mapped_listeners(&self) -> Vec; + + fn is_local_ip(&self, ip: &IpAddr) -> bool; + + async fn preferred_ipv6_source( + &self, + ip: Ipv6Addr, + context: SocketContext, + ) -> Option; + + async fn preferred_foreign_ipv6_source( + &self, + ip: Ipv6Addr, + context: SocketContext, + ) -> Option { + self.preferred_ipv6_source(ip, context).await + } +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum DirectTransport { + Tcp(TcpSocketPurpose), + Udp(UdpSessionMode), +} + +impl DirectTransport { + fn from_url(url: &Url) -> anyhow::Result { + match protocol_transport(url.scheme()) { + Some(ProtocolTransport::Tcp) => Ok(Self::Tcp(TcpSocketPurpose::DirectConnect)), + Some(ProtocolTransport::FakeTcp) => Ok(Self::Tcp(TcpSocketPurpose::FakeTcp)), + Some(ProtocolTransport::Udp(mode)) => Ok(Self::Udp(mode)), + None => anyhow::bail!("unsupported direct transport scheme: {}", url.scheme()), + } + } + + fn is_udp(self) -> bool { + matches!(self, Self::Udp(_)) + } + + fn supports_interface_bind(self) -> bool { + !matches!(self, Self::Tcp(TcpSocketPurpose::FakeTcp)) + } +} + +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct DirectConnectorOptions { + pub default_protocol: String, + pub enable_ipv6: bool, + pub allow_public_server: bool, + pub bind_device: bool, + pub allow_interface_bind: bool, + pub tcp_bind: TcpBindOptions, + pub udp_bind: UdpBindOptions, + #[serde(skip)] + pub testing: bool, +} + +impl Default for DirectConnectorOptions { + fn default() -> Self { + Self { + default_protocol: "tcp".to_owned(), + enable_ipv6: true, + allow_public_server: false, + bind_device: false, + allow_interface_bind: true, + tcp_bind: TcpBindOptions::default(), + udp_bind: UdpBindOptions::direct_connect(), + testing: false, + } + } +} + +impl DirectConnectorOptions { + fn socket_context(&self, transport: DirectTransport, ip_version: IpVersion) -> SocketContext { + let context = match transport { + DirectTransport::Tcp(_) => self.tcp_bind.context.clone(), + DirectTransport::Udp(_) => self.udp_bind.context.clone(), + }; + context.with_ip_version(ip_version) + } +} + +#[derive(Debug)] +struct ExpiringSet +where + K: Eq + Hash, +{ + entries: DashMap, +} + +impl Default for ExpiringSet +where + K: Eq + Hash, +{ + fn default() -> Self { + Self { + entries: DashMap::new(), + } + } +} + +impl ExpiringSet +where + K: Eq + Hash + Clone, +{ + fn insert(&self, key: K, ttl: Duration) { + self.entries.insert(key, StdInstant::now() + ttl); + } + + fn contains(&self, key: &K) -> bool { + let active = self + .entries + .get(key) + .is_some_and(|expires_at| *expires_at > StdInstant::now()); + if !active { + self.entries.remove(key); + } + active + } + + fn cleanup(&self) { + let now = StdInstant::now(); + self.entries.retain(|_, expires_at| *expires_at > now); + } +} + +#[derive(Debug, Hash, Eq, PartialEq, Clone)] +struct ListenerBlacklistKey(PeerId, String); + +struct DirectConnectorData +where + H: DirectConnectorHost, +{ + peer_manager: Arc, + host: Arc, + protected_tcp_ports: Arc, + stun: Arc::Socket>>, + running_listeners: Arc, + dns: Arc, + protocol: + Arc::Socket>>, + options: DirectConnectorOptions, + listener_blacklist: ExpiringSet, + peer_blacklist: ExpiringSet, +} + +impl std::fmt::Debug for DirectConnectorData +where + H: DirectConnectorHost, +{ + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("DirectConnectorData") + .field("peer_id", &self.peer_manager.my_peer_id()) + .field("network_name", &self.peer_manager.network_name()) + .finish() + } +} + +pub(crate) struct DirectConnectorManager +where + H: DirectConnectorHost, +{ + data: Arc>, + client: PeerTaskManager>, +} + +impl DirectConnectorManager +where + H: DirectConnectorHost, +{ + #[allow(clippy::too_many_arguments)] + pub(crate) fn new_with_running_listeners( + peer_manager: Arc, + host: Arc, + protected_tcp_ports: Arc, + stun: Arc::Socket>>, + running_listeners: Arc, + dns: Arc, + protocol: Arc< + dyn ClientProtocolUpgrader<::Socket>, + >, + options: DirectConnectorOptions, + ) -> Self { + let data = Arc::new(DirectConnectorData { + peer_manager: peer_manager.clone(), + host, + protected_tcp_ports, + stun, + running_listeners, + dns, + protocol, + options, + listener_blacklist: ExpiringSet::default(), + peer_blacklist: ExpiringSet::default(), + }); + let client = PeerTaskManager::new_with_external_signal( + DirectConnectorLauncher(data.clone()), + Some(peer_manager.p2p_demand_notify()), + ); + Self { data, client } + } + + pub fn run(&self) { + self.run_as_server(); + self.run_as_client(); + } + + pub fn run_as_server(&self) { + self.data + .peer_manager + .get_peer_rpc_mgr() + .rpc_server() + .registry() + .register( + GeneratedDirectConnectorRpcServer::new( + DirectConnectorRpcHandler::new_with_running_listeners_and_stun( + self.data.host.clone(), + Some(Arc::downgrade(&self.data.peer_manager)), + self.data.running_listeners.clone(), + self.data.options.udp_bind.context.clone(), + Some(self.data.stun.clone()), + ), + ), + self.data.peer_manager.network_name(), + ); + } + + pub fn run_as_client(&self) { + self.client.start(); + } + + pub async fn stop(&self) { + self.client.stop().await; + self.data + .peer_manager + .get_peer_rpc_mgr() + .rpc_server() + .registry() + .unregister( + GeneratedDirectConnectorRpcServer::new( + DirectConnectorRpcHandler::new_with_running_listeners_and_stun( + self.data.host.clone(), + Some(Arc::downgrade(&self.data.peer_manager)), + self.data.running_listeners.clone(), + self.data.options.udp_bind.context.clone(), + Some(self.data.stun.clone()), + ), + ), + self.data.peer_manager.network_name(), + ); + } + + pub(crate) async fn local_address_observations_with_stun( + &self, + stun_info: &crate::proto::common::StunInfo, + ) -> GetIpListResponse { + collect_address_observations( + self.data.host.as_ref(), + Some(self.data.peer_manager.as_ref()), + &self.data.options.udp_bind.context, + false, + Some(stun_info), + ) + .await + } +} + +struct DirectConnectorLauncher(Arc>) +where + H: DirectConnectorHost; + +impl Clone for DirectConnectorLauncher +where + H: DirectConnectorHost, +{ + fn clone(&self) -> Self { + Self(self.0.clone()) + } +} + +#[async_trait] +impl PeerTaskLauncher for DirectConnectorLauncher +where + H: DirectConnectorHost, +{ + type CollectPeerItem = PeerId; + type TaskRet = (); + + async fn collect_peers_need_task(&self) -> Vec { + let data = &self.0; + data.peer_blacklist.cleanup(); + let my_peer_id = data.peer_manager.my_peer_id(); + let policy = data.peer_manager.p2p_policy_flags(); + let now = Instant::now(); + data.peer_manager + .get_route() + .list_routes() + .await + .into_iter() + .filter(|route| { + let static_allowed = should_background_p2p_with_peer( + route.feature_flag.as_ref(), + data.options.allow_public_server, + policy.lazy_p2p, + policy.disable_p2p, + policy.need_p2p, + ); + let dynamic_allowed = should_try_p2p_with_peer( + route.feature_flag.as_ref(), + data.options.allow_public_server, + policy.disable_p2p, + policy.need_p2p, + ) && data.peer_manager.has_recent_traffic(route.peer_id, now); + route.peer_id != my_peer_id + && (static_allowed || dynamic_allowed) + && !data.peer_manager.has_directly_connected_conn(route.peer_id) + && !data.peer_blacklist.contains(&route.peer_id) + }) + .map(|route| route.peer_id) + .collect() + } + + async fn launch_task(&self, peer_id: PeerId) -> tokio::task::JoinHandle> { + let data = self.0.clone(); + tokio::spawn(async move { data.try_direct_connect(peer_id).await }) + } + + fn loop_interval_ms(&self) -> u64 { + DIRECT_TASK_LOOP_INTERVAL_MS + } +} + +impl DirectConnectorData +where + H: DirectConnectorHost, +{ + async fn try_direct_connect(self: Arc, dst_peer_id: PeerId) -> anyhow::Result<()> { + let backoffs_ms = [1000, 2000, 2000, 5000, 5000, 10000, 30000, 60000]; + let mut backoff_index = 0usize; + let mut attempt = 0usize; + + loop { + if self.peer_blacklist.contains(&dst_peer_id) { + anyhow::bail!("peer {dst_peer_id} is blacklisted"); + } + if attempt > 0 { + crate::foundation::time::sleep(Duration::from_millis(backoffs_ms[backoff_index])) + .await; + backoff_index = (backoff_index + 1).min(backoffs_ms.len() - 1); + } + attempt += 1; + + let rpc_stub = self + .peer_manager + .get_peer_rpc_mgr() + .rpc_client() + .scoped_client::>( + self.peer_manager.my_peer_id(), + dst_peer_id, + self.peer_manager.network_name().to_owned(), + ); + let ip_list = match rpc_stub + .get_ip_list(BaseController::default(), GetIpListRequest {}) + .await + { + Ok(ip_list) => ip_list, + Err(error @ rpc_types::error::Error::InvalidServiceKey(_, _)) => { + self.peer_blacklist + .insert(dst_peer_id, INVALID_SERVICE_BLACKLIST_TIMEOUT); + return Err(error.into()); + } + Err(error) => return Err(error.into()), + }; + + tracing::info!(?ip_list, dst_peer_id, "got direct-connect IP list"); + let result = self + .try_direct_connect_with_ip_list(dst_peer_id, ip_list) + .await; + tracing::info!(?result, dst_peer_id, "direct-connect attempt returned"); + if self.peer_manager.has_directly_connected_conn(dst_peer_id) { + return Ok(()); + } + } + } + + async fn try_direct_connect_with_ip_list( + self: &Arc, + dst_peer_id: PeerId, + ip_list: GetIpListResponse, + ) -> anyhow::Result<()> { + let mut available_listeners = ip_list + .listeners + .clone() + .into_iter() + .map(Into::::into) + .filter(|listener| listener.scheme() != "ring") + .filter(|listener| { + mapped_listener_port(listener).is_some() && listener.host().is_some() + }) + .filter(|listener| { + self.options.enable_ipv6 || !matches!(listener.host(), Some(Host::Ipv6(_))) + }) + .collect::>(); + + if available_listeners.is_empty() { + anyhow::bail!("peer {dst_peer_id} has no valid listener"); + } + + available_listeners.sort_by_key(|listener| { + if listener.scheme() == self.options.default_protocol { + 3 + } else if listener.scheme() == "udp" { + 2 + } else { + 1 + } + }); + + while !available_listeners.is_empty() { + let mut tasks = JoinSet::new(); + let current_scheme = available_listeners + .last() + .expect("non-empty listener list") + .scheme() + .to_owned(); + while available_listeners + .last() + .is_some_and(|listener| listener.scheme() == current_scheme) + { + let listener = available_listeners.pop().expect("listener should exist"); + self.spawn_direct_connect_tasks(dst_peer_id, &ip_list, &listener, &mut tasks) + .await; + } + let _ = tasks.join_all().await; + if self.peer_manager.has_directly_connected_conn(dst_peer_id) { + return Ok(()); + } + } + Ok(()) + } + + async fn spawn_direct_connect_tasks( + self: &Arc, + dst_peer_id: PeerId, + ip_list: &GetIpListResponse, + listener: &Url, + tasks: &mut JoinSet>, + ) { + let Ok(mut addrs) = resolve_mapped_listener_addrs( + listener, + self.options.tcp_bind.context.clone(), + self.dns.as_ref(), + ) + .await + else { + tracing::error!(?listener, "failed to resolve direct listener"); + return; + }; + let listener_host = addrs.pop(); + let is_udp = protocol_uses_udp(listener.scheme()); + let local_listeners = self.running_listeners.local_listener_urls(); + let port_has_local_listener = |port: u16| { + local_listeners.iter().any(|local| { + local.port() == Some(port) && protocol_uses_udp(local.scheme()) == is_udp + }) + }; + let should_deny_target = |target: &SocketAddr| { + let port_is_protected = port_has_local_listener(target.port()) + || (!is_udp && self.protected_tcp_ports.contains(target.port())); + port_is_protected && self.host.is_local_ip(&target.ip()) + }; + + match listener_host { + Some(SocketAddr::V4(socket_addr)) if socket_addr.ip().is_unspecified() => { + for ip in ip_list + .interface_ipv4s + .iter() + .chain(ip_list.public_ipv4.iter()) + { + let target = SocketAddr::new(IpAddr::V4(ip.addr.into()), socket_addr.port()); + if should_deny_target(&target) { + continue; + } + let mut url = listener.clone(); + if url.set_ip_host(target.ip()).is_ok() { + tasks.spawn(Self::try_connect_to_url( + self.clone(), + dst_peer_id, + url.to_string(), + )); + } + } + } + Some(SocketAddr::V4(socket_addr)) + if (!socket_addr.ip().is_loopback() || self.options.testing) + && !should_deny_target(&SocketAddr::V4(socket_addr)) => + { + tasks.spawn(Self::try_connect_to_url( + self.clone(), + dst_peer_id, + listener.to_string(), + )); + } + Some(SocketAddr::V6(socket_addr)) if socket_addr.ip().is_unspecified() => { + let mut candidates = HashSet::new(); + for ip in ip_list + .interface_ipv6s + .iter() + .chain(ip_list.public_ipv6.iter()) + .map(|ip| Ipv6Addr::from(*ip)) + { + if self.is_usable_public_ipv6(&ip).await { + candidates.insert(ip); + } + } + for ip in candidates { + let target = SocketAddr::new(IpAddr::V6(ip), socket_addr.port()); + if should_deny_target(&target) { + continue; + } + let mut url = listener.clone(); + if url.set_ip_host(target.ip()).is_ok() { + tasks.spawn(Self::try_connect_to_url( + self.clone(), + dst_peer_id, + url.to_string(), + )); + } + } + } + Some(SocketAddr::V6(socket_addr)) + if self + .peer_manager + .is_easytier_managed_ipv6(socket_addr.ip()) + .await => + { + tracing::debug!(?listener, "skip managed IPv6 direct target"); + } + Some(SocketAddr::V6(socket_addr)) + if (!socket_addr.ip().is_loopback() || self.options.testing) + && !should_deny_target(&SocketAddr::V6(socket_addr)) => + { + tasks.spawn(Self::try_connect_to_url( + self.clone(), + dst_peer_id, + listener.to_string(), + )); + } + _ => {} + } + } + + async fn try_connect_to_url( + self: Arc, + dst_peer_id: PeerId, + url: String, + ) -> anyhow::Result<()> { + self.listener_blacklist.cleanup(); + let key = ListenerBlacklistKey(dst_peer_id, url.clone()); + if self.listener_blacklist.contains(&key) { + anyhow::bail!("direct listener URL is blacklisted"); + } + + let backoffs_ms = [1000i64, 2000, 4000]; + for attempt in 0..=backoffs_ms.len() { + if self.peer_manager.has_directly_connected_conn(dst_peer_id) { + return Ok(()); + } + let result = self.connect_to_url_once(dst_peer_id, &url).await; + if result.is_ok() || self.peer_manager.has_directly_connected_conn(dst_peer_id) { + return Ok(()); + } + if attempt == backoffs_ms.len() { + self.listener_blacklist + .insert(key, DIRECT_CONNECTOR_BLACKLIST_TIMEOUT); + return result; + } + + let base = backoffs_ms[attempt]; + let delta = base >> 1; + let delay_ms = { + let mut rng = rand::thread_rng(); + base + rng.gen_range(-delta..delta) + }; + crate::foundation::time::sleep(Duration::from_millis(delay_ms as u64)).await; + } + unreachable!("direct URL retry loop must return") + } + + async fn connect_to_url_once(&self, dst_peer_id: PeerId, raw_url: &str) -> anyhow::Result<()> { + let url = Url::parse(raw_url)?; + let (peer_id, conn_id) = if url.scheme() == "udp" { + match url.host() { + Some(Host::Ipv6(_)) => self.connect_public_ipv6(dst_peer_id, &url).await?, + Some(Host::Ipv4(ip)) if is_public_ipv4(ip) => { + match self.connect_public_ipv4(dst_peer_id, &url).await { + Ok(result) => result, + Err(error) => { + tracing::debug!(?error, %url, "public IPv4 UDP punch failed; fallback"); + self.connect_ordinary(dst_peer_id, url.clone()).await? + } + } + } + _ => self.connect_ordinary(dst_peer_id, url.clone()).await?, + } + } else { + self.connect_ordinary(dst_peer_id, url.clone()).await? + }; + + if peer_id != dst_peer_id && !self.options.testing { + self.peer_manager.close_peer_conn(peer_id, &conn_id).await?; + anyhow::bail!("direct peer mismatch for {url}: expected {dst_peer_id}, got {peer_id}"); + } + Ok(()) + } + + async fn connect_ordinary( + &self, + dst_peer_id: PeerId, + url: Url, + ) -> anyhow::Result<(PeerId, PeerConnId)> { + let transport = DirectTransport::from_url(&url)?; + let normalized = convert_idn_to_ascii(url.clone())?; + let default_port = mapped_listener_port(&normalized) + .ok_or_else(|| anyhow::anyhow!("listener has no port: {url}"))?; + let remote_addr = resolve_remote_addr( + self.peer_manager.as_ref(), + self.host.as_ref(), + self.dns.as_ref(), + &normalized, + default_port, + self.options.socket_context(transport, IpVersion::Both), + ) + .await?; + let bind_addrs = if self.options.bind_device + && self.options.allow_interface_bind + && transport.supports_interface_bind() + { + collect_bind_addrs( + self.peer_manager.as_ref(), + self.host.as_ref(), + transport.is_udp(), + remote_addr, + ) + .await? + } else { + Vec::new() + }; + crate::foundation::time::timeout(DIRECT_CONNECT_TIMEOUT, async { + let connected = match transport { + DirectTransport::Tcp(purpose) => ConnectedTransport::Tcp( + transport::connect_tcp( + self.host.clone(), + remote_addr, + bind_addrs, + self.options.tcp_bind.clone(), + purpose, + ) + .await?, + ), + DirectTransport::Udp(mode) => ConnectedTransport::Udp( + transport::connect_udp( + self.host.clone(), + remote_addr, + bind_addrs, + self.options.udp_bind.clone(), + mode, + ) + .await?, + ), + }; + let tunnel = self.protocol.upgrade_client(connected, url).await?; + self.admit(tunnel, dst_peer_id).await + }) + .await? + } + + async fn connect_public_ipv4( + &self, + dst_peer_id: PeerId, + url: &Url, + ) -> anyhow::Result<(PeerId, PeerConnId)> { + let socket = self + .host + .bind_udp( + UdpBindOptions::direct_connect() + .with_context( + self.options + .udp_bind + .context + .clone() + .with_ip_version(IpVersion::V4), + ) + .with_local_addr(Some("0.0.0.0:0".parse().unwrap())), + ) + .await?; + let connector_addr = self + .stun + .get_udp_port_mapping_with_socket(socket.clone()) + .await?; + let _ = self + .remote_send_udp_hole_punch_packet(dst_peer_id, vec![connector_addr], None, url) + .await; + let remote_addr = resolve_literal_url(url, IpVersion::V4)?; + let connected = udp::connect_with_socket(self.host.clone(), socket, remote_addr).await?; + let tunnel = self + .protocol + .upgrade_client(ConnectedTransport::Udp(connected), url.clone()) + .await?; + self.admit(tunnel, dst_peer_id).await + } + + async fn connect_public_ipv6( + &self, + dst_peer_id: PeerId, + url: &Url, + ) -> anyhow::Result<(PeerId, PeerConnId)> { + let socket = self + .host + .bind_udp( + UdpBindOptions::direct_connect() + .with_context( + self.options + .udp_bind + .context + .clone() + .with_ip_version(IpVersion::V6), + ) + .with_local_addr(Some("[::]:0".parse().unwrap())), + ) + .await?; + let connector_ips = self.collect_ipv6_hole_punch_candidates().await?; + if !connector_ips.is_empty() { + let port = socket.local_addr()?.port(); + let connector_addrs = connector_ips + .into_iter() + .map(|ip| SocketAddr::new(IpAddr::V6(ip), port)) + .collect(); + let preferred_src = match url.host() { + Some(Host::Ipv6(ip)) => Some(ip), + _ => None, + }; + let _ = self + .remote_send_udp_hole_punch_packet(dst_peer_id, connector_addrs, preferred_src, url) + .await; + } + let remote_addr = resolve_literal_url(url, IpVersion::V6)?; + let connected = udp::connect_with_socket(self.host.clone(), socket, remote_addr).await?; + let tunnel = self + .protocol + .upgrade_client(ConnectedTransport::Udp(connected), url.clone()) + .await?; + self.admit(tunnel, dst_peer_id).await + } + + async fn collect_ipv6_hole_punch_candidates(&self) -> anyhow::Result> { + let mut candidates = Vec::new(); + for ip in self + .stun + .get_stun_info() + .public_ip + .into_iter() + .filter_map(|ip| ip.parse().ok()) + { + if let IpAddr::V6(ip) = ip { + self.push_ipv6_candidate(&mut candidates, ip).await; + } + } + let interface_addrs = self.host.interface_addrs().await?; + for ip in interface_addrs + .interface_ipv6s + .into_iter() + .chain(interface_addrs.public_ipv6) + { + self.push_ipv6_candidate(&mut candidates, ip).await; + } + Ok(candidates) + } + + async fn push_ipv6_candidate(&self, candidates: &mut Vec, ip: Ipv6Addr) { + if candidates.len() < MAX_IPV6_HOLE_PUNCH_CONNECTOR_ADDRS + && self.is_usable_public_ipv6(&ip).await + && !candidates.contains(&ip) + { + candidates.push(ip); + } + } + + async fn is_usable_public_ipv6(&self, ip: &Ipv6Addr) -> bool { + !self.peer_manager.is_easytier_managed_ipv6(ip).await + && (self.options.testing + || (!ip.is_loopback() + && !ip.is_unspecified() + && !ip.is_unique_local() + && !ip.is_unicast_link_local() + && !ip.is_multicast())) + } + + async fn remote_send_udp_hole_punch_packet( + &self, + dst_peer_id: PeerId, + connector_addrs: Vec, + preferred_src_ipv6: Option, + remote_url: &Url, + ) -> anyhow::Result<()> { + if remote_url.scheme() != "udp" { + anyhow::bail!("UDP punch request requires a UDP listener: {remote_url}"); + } + let listener_port = mapped_listener_port(remote_url) + .ok_or_else(|| anyhow::anyhow!("listener has no port: {remote_url}"))?; + let rpc_stub = self + .peer_manager + .get_peer_rpc_mgr() + .rpc_client() + .scoped_client::>( + self.peer_manager.my_peer_id(), + dst_peer_id, + self.peer_manager.network_name().to_owned(), + ); + rpc_stub + .send_udp_hole_punch_packet( + BaseController::default(), + SendUdpHolePunchPacketRequest { + connector_addr: connector_addrs.first().copied().map(Into::into), + listener_port: listener_port as u32, + preferred_src_ipv6: preferred_src_ipv6.map(Into::into), + connector_addrs: connector_addrs.into_iter().map(Into::into).collect(), + }, + ) + .await + .with_context(|| format!("send UDP punch request to peer {dst_peer_id}"))?; + Ok(()) + } + + async fn admit( + &self, + tunnel: Box, + dst_peer_id: PeerId, + ) -> anyhow::Result<(PeerId, PeerConnId)> { + self.peer_manager + .add_client_tunnel_with_peer_id_hint(tunnel, true, Some(dst_peer_id)) + .await + .map_err(Into::into) + } +} + +pub(crate) struct DirectConnectorRpcHandler +where + H: DirectConnectorHost, +{ + host: Arc, + peer_manager: Option>, + running_listeners: Arc, + socket_context: SocketContext, + foreign_network: bool, + stun: Option>, +} + +impl Clone for DirectConnectorRpcHandler +where + H: DirectConnectorHost, +{ + fn clone(&self) -> Self { + Self { + host: self.host.clone(), + peer_manager: self.peer_manager.clone(), + running_listeners: self.running_listeners.clone(), + socket_context: self.socket_context.clone(), + foreign_network: self.foreign_network, + stun: self.stun.clone(), + } + } +} + +impl DirectConnectorRpcHandler +where + H: DirectConnectorHost, +{ + pub fn new_for_foreign_network_with_stun( + host: Arc, + socket_context: SocketContext, + stun: Option>, + ) -> Self { + let running_listeners = Arc::new(NoLocalListeners); + Self { + host, + peer_manager: None, + running_listeners, + socket_context, + foreign_network: true, + stun, + } + } + + fn new_with_running_listeners_and_stun( + host: Arc, + peer_manager: Option>, + running_listeners: Arc, + socket_context: SocketContext, + stun: Option>, + ) -> Self { + Self { + host, + peer_manager, + running_listeners, + socket_context, + foreign_network: false, + stun, + } + } +} + +async fn collect_address_observations( + host: &H, + peer_manager: Option<&PeerManagerCore>, + socket_context: &SocketContext, + foreign_network: bool, + stun_info: Option<&crate::proto::common::StunInfo>, +) -> GetIpListResponse +where + H: DirectConnectorHost, +{ + let mut response = if foreign_network { + host.collect_foreign_ip_addrs(socket_context).await + } else { + host.collect_ip_addrs(socket_context).await + }; + if let Some(stun_info) = stun_info { + for public_ip in &stun_info.public_ip { + match public_ip.parse::() { + Ok(IpAddr::V4(ip)) => response.public_ipv4 = Some(ip.into()), + Ok(IpAddr::V6(ip)) => response.public_ipv6 = Some(ip.into()), + Err(_) => {} + } + } + } + if let Some(peer_manager) = peer_manager.filter(|_| !foreign_network) { + let mut interface_ipv6s = Vec::with_capacity(response.interface_ipv6s.len()); + for ip in response.interface_ipv6s { + if !peer_manager + .is_easytier_managed_ipv6(&Ipv6Addr::from(ip)) + .await + { + interface_ipv6s.push(ip); + } + } + response.interface_ipv6s = interface_ipv6s; + if let Some(ip) = response.public_ipv6.map(Ipv6Addr::from) + && peer_manager.is_easytier_managed_ipv6(&ip).await + { + response.public_ipv6 = None; + } + } + response +} + +pub(crate) struct ForeignDirectConnectorRpcRegistrar +where + H: DirectConnectorHost, +{ + host: Arc, + stun: Arc, +} + +impl ForeignDirectConnectorRpcRegistrar +where + H: DirectConnectorHost, +{ + pub fn new(host: Arc, stun: Arc) -> Self { + Self { host, stun } + } +} + +impl ForeignNetworkRpcRegistrar for ForeignDirectConnectorRpcRegistrar +where + H: DirectConnectorHost, +{ + fn register_peer_rpc_services( + &self, + peer_rpc: &Arc, + network_name: &str, + socket_context: SocketContext, + ) { + peer_rpc.rpc_server().registry().register( + GeneratedDirectConnectorRpcServer::new( + DirectConnectorRpcHandler::new_for_foreign_network_with_stun( + self.host.clone(), + socket_context, + Some(self.stun.clone()), + ), + ), + network_name, + ); + } +} + +#[async_trait] +impl DirectConnectorRpc for DirectConnectorRpcHandler +where + H: DirectConnectorHost, +{ + type Controller = BaseController; + + async fn get_ip_list( + &self, + _: BaseController, + _: GetIpListRequest, + ) -> rpc_types::error::Result { + let peer_manager = self.peer_manager.as_ref().and_then(Weak::upgrade); + let mut response = collect_address_observations( + self.host.as_ref(), + peer_manager.as_deref(), + &self.socket_context, + self.foreign_network, + self.stun + .as_deref() + .map(StunInfoProvider::get_stun_info) + .as_ref(), + ) + .await; + response.listeners = self + .host + .mapped_listeners() + .into_iter() + .chain(self.running_listeners.local_listener_urls()) + .map(Into::into) + .collect(); + Ok(response) + } + + async fn send_udp_hole_punch_packet( + &self, + _: BaseController, + request: SendUdpHolePunchPacketRequest, + ) -> rpc_types::error::Result { + let (listener_port, connector_addrs, preferred_src_ipv6) = + connector_addrs_from_request(request)?; + let peer_manager = self.peer_manager.as_ref().and_then(Weak::upgrade); + let preferred_source = match preferred_src_ipv6.map(Ipv6Addr::from) { + Some(ip) => { + if self.foreign_network { + self.host + .preferred_foreign_ipv6_source(ip, self.socket_context.clone()) + .await + } else if let Some(peer_manager) = peer_manager.as_deref() + && peer_manager.is_easytier_managed_ipv6(&ip).await + { + None + } else { + self.host + .preferred_ipv6_source(ip, self.socket_context.clone()) + .await + } + } + None => None, + }; + + for _ in 0..3 { + for connector_addr in &connector_addrs { + let result = match connector_addr { + SocketAddr::V4(addr) => { + send_v4_hole_punch_control_packet( + self.host.as_ref(), + self.socket_context.clone(), + listener_port, + *addr, + ) + .await + } + SocketAddr::V6(addr) => { + send_v6_hole_punch_control_packet( + self.host.as_ref(), + self.socket_context.clone(), + listener_port, + *addr, + preferred_source, + ) + .await + } + }; + if let Err(error) = result { + tracing::debug!(?error, ?connector_addr, "send UDP punch packet failed"); + } + } + crate::foundation::time::sleep(Duration::from_millis(30)).await; + } + Ok(Void::default()) + } +} + +fn connector_addrs_from_request( + request: SendUdpHolePunchPacketRequest, +) -> rpc_types::error::Result<(u16, Vec, Option)> { + let listener_port = u16::try_from(request.listener_port) + .map_err(|_| anyhow::anyhow!("listener_port out of range: {}", request.listener_port))?; + let mut connector_addrs = request + .connector_addrs + .into_iter() + .map(SocketAddr::from) + .collect::>(); + if connector_addrs.is_empty() { + connector_addrs.push( + request + .connector_addr + .ok_or_else(|| anyhow::anyhow!("connector_addr is required"))? + .into(), + ); + } + let mut deduped = Vec::with_capacity(connector_addrs.len()); + for addr in connector_addrs { + if !deduped.contains(&addr) { + deduped.push(addr); + } + if deduped.len() >= MAX_UDP_HOLE_PUNCH_CONNECTOR_ADDRS { + break; + } + } + Ok((listener_port, deduped, request.preferred_src_ipv6)) +} + +fn mapped_listener_port(url: &Url) -> Option { + url.port() + .or_else(|| crate::connectivity::protocol::protocol_default_port(url.scheme())) +} + +async fn resolve_mapped_listener_addrs( + listener: &Url, + context: SocketContext, + dns: &dyn DnsResolver, +) -> anyhow::Result> { + let port = mapped_listener_port(listener) + .ok_or_else(|| anyhow::anyhow!("listener has no default port: {listener}"))?; + resolve_url_addrs( + listener, + port, + context.with_ip_version(IpVersion::Both), + dns, + ) + .await +} + +fn resolve_literal_url(url: &Url, ip_version: IpVersion) -> anyhow::Result { + let port = + mapped_listener_port(url).ok_or_else(|| anyhow::anyhow!("listener has no port: {url}"))?; + match (url.host(), ip_version) { + (Some(Host::Ipv4(ip)), IpVersion::V4 | IpVersion::Both) => { + Ok(SocketAddr::new(IpAddr::V4(ip), port)) + } + (Some(Host::Ipv6(ip)), IpVersion::V6 | IpVersion::Both) => { + Ok(SocketAddr::new(IpAddr::V6(ip), port)) + } + _ => anyhow::bail!("URL host does not match {ip_version:?}: {url}"), + } +} + +fn is_public_ipv4(ip: Ipv4Addr) -> bool { + !ip.is_private() + && !ip.is_loopback() + && !ip.is_link_local() + && !ip.is_broadcast() + && !ip.is_unspecified() +} + +#[cfg(test)] +mod tests { + use super::*; + + impl DirectConnectorRpcHandler + where + H: DirectConnectorHost, + { + pub(crate) fn new_with_stun( + host: Arc, + socket_context: SocketContext, + stun: Option>, + ) -> Self { + let running_listeners = Arc::new(NoLocalListeners); + Self { + host, + peer_manager: None, + running_listeners, + socket_context, + foreign_network: false, + stun, + } + } + } + + #[test] + fn faketcp_uses_specialized_tcp_socket_without_interface_binding() { + let transport = + DirectTransport::from_url(&"faketcp://127.0.0.1:11013".parse().unwrap()).unwrap(); + + assert_eq!(transport, DirectTransport::Tcp(TcpSocketPurpose::FakeTcp)); + assert!(!transport.supports_interface_bind()); + assert!(!transport.is_udp()); + } + + #[test] + fn connector_address_request_deduplicates_and_caps() { + let mut request = SendUdpHolePunchPacketRequest { + listener_port: 11010, + ..Default::default() + }; + for port in 1..=20 { + request + .connector_addrs + .push(SocketAddr::from(([127, 0, 0, 1], port)).into()); + } + request + .connector_addrs + .push(SocketAddr::from(([127, 0, 0, 1], 1)).into()); + + let (_, addrs, _) = connector_addrs_from_request(request).unwrap(); + assert_eq!(addrs.len(), MAX_UDP_HOLE_PUNCH_CONNECTOR_ADDRS); + assert_eq!(addrs[0], SocketAddr::from(([127, 0, 0, 1], 1))); + } + + #[test] + fn expiring_set_removes_expired_entries() { + let set = ExpiringSet::default(); + set.insert(7u32, Duration::ZERO); + assert!(!set.contains(&7)); + } +} diff --git a/easytier-core/src/connectivity/direct/udp.rs b/easytier-core/src/connectivity/direct/udp.rs new file mode 100644 index 00000000..272101ec --- /dev/null +++ b/easytier-core/src/connectivity/direct/udp.rs @@ -0,0 +1,21 @@ +use std::{net::SocketAddr, sync::Arc}; + +use crate::{ + connectivity::transport::ConnectedUdpSession, + socket::udp::{UdpSessionLayer, VirtualUdpSocketFactory}, +}; + +use super::DirectConnectorHost; + +pub(super) async fn connect_with_socket( + host: Arc, + socket: Arc<::Socket>, + remote_addr: SocketAddr, +) -> anyhow::Result +where + H: DirectConnectorHost, +{ + let layer = Arc::new(UdpSessionLayer::new_with_stun_responder(socket, host)); + let session = layer.connect(remote_addr).await?; + Ok(ConnectedUdpSession::new(session, layer)) +} diff --git a/easytier-core/src/connectivity/hole_punch/mod.rs b/easytier-core/src/connectivity/hole_punch/mod.rs new file mode 100644 index 00000000..ec1e951d --- /dev/null +++ b/easytier-core/src/connectivity/hole_punch/mod.rs @@ -0,0 +1,33 @@ +use async_trait::async_trait; + +use crate::proto::rpc_types::{controller::BaseController, handler::Handler}; +use crate::tunnel::Tunnel; + +mod peer_adapters; +pub(crate) mod policy; +pub mod port_mapping; +pub(crate) mod tcp; +pub(crate) mod udp; + +/// Registration seam for hole-punch RPC services. +/// +/// The engines build the proto-generated server wrapper around their RPC +/// endpoint; the implementation owns the peer RPC registry and the network +/// domain the service is registered under. Implemented only by the sealed +/// peer adapter in `peer_adapters.rs`. +pub(crate) trait HolePunchRpcRegistry: Send + Sync + 'static { + fn register_rpc_service(&self, service: H) + where + H: Handler; + + fn unregister_rpc_service(&self, service: H) + where + H: Handler; +} + +#[async_trait] +pub(crate) trait HolePunchTunnelSink: Send + Sync + 'static { + async fn add_client_tunnel(&self, tunnel: Box) -> anyhow::Result<()>; + + async fn add_server_tunnel(&self, tunnel: Box) -> anyhow::Result<()>; +} diff --git a/easytier-core/src/connectivity/hole_punch/peer_adapters.rs b/easytier-core/src/connectivity/hole_punch/peer_adapters.rs new file mode 100644 index 00000000..07d0a7c1 --- /dev/null +++ b/easytier-core/src/connectivity/hole_punch/peer_adapters.rs @@ -0,0 +1,188 @@ +use std::{ + net::{IpAddr, Ipv6Addr}, + sync::Arc, +}; + +use async_trait::async_trait; +use quanta::Instant; + +use crate::{ + config::{P2pPolicyFlags, PeerId}, + foundation::task::ExternalTaskSignal, + peers::peer_manager::PeerManagerCore, + proto::{ + common::NatType, + peer_rpc::{ + TcpHolePunchRpc, TcpHolePunchRpcClientFactory, UdpHolePunchRpc, + UdpHolePunchRpcClientFactory, + }, + rpc_types::{controller::BaseController, handler::Handler}, + }, + tunnel::Tunnel, +}; + +use super::{ + HolePunchRpcRegistry, HolePunchTunnelSink, + tcp::{TcpHolePunchPeerSource, TcpPunchCandidate}, + udp::{UdpHolePunchPeerSource, UdpHolePunchRpcSource, UdpPunchCandidate}, +}; + +#[async_trait] +impl UdpHolePunchPeerSource for PeerManagerCore { + fn local_peer_id(&self) -> PeerId { + PeerManagerCore::my_peer_id(self) + } + + fn p2p_policy_flags(&self) -> P2pPolicyFlags { + PeerManagerCore::p2p_policy_flags(self) + } + + async fn candidates(&self) -> Vec { + let now = Instant::now(); + let peer_map = self.get_peer_map(); + self.list_route_snapshots() + .await + .into_iter() + .filter_map(|route| { + let udp_nat_type = route + .stun_info + .as_ref() + .map(|info| info.udp_nat_type) + .unwrap_or_default(); + let Ok(udp_nat_type) = crate::proto::common::NatType::try_from(udp_nat_type) else { + return None; + }; + Some(UdpPunchCandidate { + peer_id: route.peer_id, + udp_nat_type, + feature_flag: route.feature_flag, + has_direct_connection: peer_map.has_peer(route.peer_id), + has_recent_traffic: self.has_recent_traffic(route.peer_id, now), + }) + }) + .collect() + } + + fn p2p_demand_notify(&self) -> Arc { + PeerManagerCore::p2p_demand_notify(self) + } + + fn is_local_virtual_ip(&self, ip: &IpAddr) -> bool { + PeerManagerCore::is_local_virtual_ip(self, ip) + } + + async fn is_easytier_managed_ipv6(&self, ip: &Ipv6Addr) -> bool { + PeerManagerCore::is_easytier_managed_ipv6(self, ip).await + } +} + +impl UdpHolePunchRpcSource for PeerManagerCore { + fn local_peer_id(&self) -> PeerId { + PeerManagerCore::my_peer_id(self) + } + + fn rpc_stub( + &self, + dst_peer_id: PeerId, + ) -> Box + Send + Sync + 'static> { + PeerManagerCore::get_peer_rpc_mgr(self) + .rpc_client() + .scoped_client::>( + PeerManagerCore::my_peer_id(self), + dst_peer_id, + PeerManagerCore::network_name(self).to_owned(), + ) + } +} + +#[async_trait] +impl HolePunchTunnelSink for PeerManagerCore { + async fn add_client_tunnel(&self, tunnel: Box) -> anyhow::Result<()> { + PeerManagerCore::add_client_tunnel(self, tunnel, false) + .await + .map(|_| ()) + .map_err(anyhow::Error::from) + } + + async fn add_server_tunnel(&self, tunnel: Box) -> anyhow::Result<()> { + PeerManagerCore::add_tunnel_as_server(self, tunnel, false) + .await + .map_err(anyhow::Error::from) + } +} + +#[async_trait] +impl TcpHolePunchPeerSource for PeerManagerCore { + fn local_peer_id(&self) -> PeerId { + PeerManagerCore::my_peer_id(self) + } + + fn p2p_policy_flags(&self) -> P2pPolicyFlags { + PeerManagerCore::p2p_policy_flags(self) + } + + fn tcp_hole_punching_disabled(&self) -> bool { + PeerManagerCore::tcp_hole_punching_disabled(self) + } + + fn p2p_demand_notify(&self) -> Arc { + PeerManagerCore::p2p_demand_notify(self) + } + + async fn candidates(&self) -> Vec { + let now = Instant::now(); + let peer_map = self.get_peer_map(); + self.list_route_snapshots() + .await + .into_iter() + .map(|route| TcpPunchCandidate { + peer_id: route.peer_id, + tcp_nat_type: route + .stun_info + .as_ref() + .map(|info| info.tcp_nat_type) + .and_then(|nat_type| NatType::try_from(nat_type).ok()) + .unwrap_or(NatType::Unknown), + feature_flag: route.feature_flag, + has_direct_connection: peer_map.has_peer(route.peer_id), + has_recent_traffic: self.has_recent_traffic(route.peer_id, now), + }) + .collect() + } + + fn rpc_stub( + &self, + dst_peer_id: PeerId, + ) -> Box + Send + Sync + 'static> { + PeerManagerCore::get_peer_rpc_mgr(self) + .rpc_client() + .scoped_client::>( + PeerManagerCore::my_peer_id(self), + dst_peer_id, + PeerManagerCore::network_name(self).to_owned(), + ) + } +} + +#[async_trait] +impl HolePunchRpcRegistry for PeerManagerCore { + fn register_rpc_service(&self, service: H) + where + H: Handler, + { + PeerManagerCore::get_peer_rpc_mgr(self) + .rpc_server() + .registry() + .register(service, PeerManagerCore::network_name(self)); + } + + fn unregister_rpc_service(&self, service: H) + where + H: Handler, + { + PeerManagerCore::get_peer_rpc_mgr(self) + .rpc_server() + .registry() + .unregister(service, PeerManagerCore::network_name(self)); + } +} diff --git a/easytier-core/src/connectivity/hole_punch/policy.rs b/easytier-core/src/connectivity/hole_punch/policy.rs new file mode 100644 index 00000000..68ffd4a4 --- /dev/null +++ b/easytier-core/src/connectivity/hole_punch/policy.rs @@ -0,0 +1,212 @@ +use crate::proto::common::PeerFeatureFlag; + +#[derive(Debug)] +pub struct BackOff { + backoffs_ms: Vec, + current_idx: usize, +} + +impl BackOff { + pub fn new(backoffs_ms: Vec) -> Self { + Self { + backoffs_ms, + current_idx: 0, + } + } + + pub fn next_backoff(&mut self) -> u64 { + let backoff = self.backoffs_ms[self.current_idx]; + self.current_idx = (self.current_idx + 1).min(self.backoffs_ms.len() - 1); + backoff + } + + pub fn rollback(&mut self) { + self.current_idx = self.current_idx.saturating_sub(1); + } + + pub async fn sleep_for_next_backoff(&mut self) { + let backoff = self.next_backoff(); + if backoff > 0 { + crate::foundation::time::sleep(crate::foundation::time::Duration::from_millis(backoff)) + .await; + } + } +} + +pub fn should_try_p2p_with_peer( + feature_flag: Option<&PeerFeatureFlag>, + allow_public_server: bool, + local_disable_p2p: bool, + local_need_p2p: bool, +) -> bool { + feature_flag + .map(|flag| { + (allow_public_server || !flag.is_public_server) + && (!local_disable_p2p || flag.need_p2p) + && (!flag.disable_p2p || local_need_p2p) + }) + .unwrap_or(!local_disable_p2p) +} + +pub fn should_background_p2p_with_peer( + feature_flag: Option<&PeerFeatureFlag>, + allow_public_server: bool, + lazy_p2p: bool, + local_disable_p2p: bool, + local_need_p2p: bool, +) -> bool { + should_try_p2p_with_peer( + feature_flag, + allow_public_server, + local_disable_p2p, + local_need_p2p, + ) && (!lazy_p2p || feature_flag.map(|flag| flag.need_p2p).unwrap_or(false)) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn backoff_saturates_and_can_rollback() { + let mut backoff = BackOff::new(vec![10, 20]); + + assert_eq!(backoff.next_backoff(), 10); + assert_eq!(backoff.next_backoff(), 20); + assert_eq!(backoff.next_backoff(), 20); + backoff.rollback(); + assert_eq!(backoff.next_backoff(), 10); + } + + #[test] + fn lazy_background_p2p_requires_need_p2p() { + let no_need_p2p = PeerFeatureFlag { + need_p2p: false, + ..Default::default() + }; + let need_p2p = PeerFeatureFlag { + need_p2p: true, + ..Default::default() + }; + + assert!(should_background_p2p_with_peer( + Some(&no_need_p2p), + false, + false, + false, + false + )); + assert!(!should_background_p2p_with_peer( + Some(&no_need_p2p), + false, + true, + false, + false + )); + assert!(should_background_p2p_with_peer( + Some(&need_p2p), + false, + true, + false, + false + )); + } + + #[test] + fn p2p_policy_respects_public_server_setting() { + let public_server = PeerFeatureFlag { + is_public_server: true, + ..Default::default() + }; + + assert!(!should_try_p2p_with_peer( + Some(&public_server), + false, + false, + false + )); + assert!(should_try_p2p_with_peer( + Some(&public_server), + true, + false, + false + )); + assert!(!should_background_p2p_with_peer( + Some(&public_server), + false, + false, + false, + false + )); + assert!(should_background_p2p_with_peer( + Some(&public_server), + true, + false, + false, + false + )); + } + + #[test] + fn disable_p2p_only_allows_need_p2p_exceptions() { + let normal_peer = PeerFeatureFlag::default(); + let need_peer = PeerFeatureFlag { + need_p2p: true, + ..Default::default() + }; + let disable_peer = PeerFeatureFlag { + disable_p2p: true, + ..Default::default() + }; + let disable_need_peer = PeerFeatureFlag { + disable_p2p: true, + need_p2p: true, + ..Default::default() + }; + + assert!(should_try_p2p_with_peer( + Some(&normal_peer), + false, + false, + false + )); + assert!(should_try_p2p_with_peer(None, false, false, false)); + assert!(!should_try_p2p_with_peer(None, false, true, false)); + assert!(!should_try_p2p_with_peer( + Some(&normal_peer), + false, + true, + false + )); + assert!(should_try_p2p_with_peer( + Some(&need_peer), + false, + true, + false + )); + assert!(!should_try_p2p_with_peer( + Some(&disable_peer), + false, + false, + false + )); + assert!(should_try_p2p_with_peer( + Some(&disable_peer), + false, + false, + true + )); + assert!(should_try_p2p_with_peer( + Some(&disable_need_peer), + false, + true, + true + )); + assert!(!should_try_p2p_with_peer( + Some(&disable_need_peer), + false, + true, + false + )); + } +} diff --git a/easytier-core/src/connectivity/hole_punch/port_mapping.rs b/easytier-core/src/connectivity/hole_punch/port_mapping.rs new file mode 100644 index 00000000..5d78290d --- /dev/null +++ b/easytier-core/src/connectivity/hole_punch/port_mapping.rs @@ -0,0 +1,439 @@ +use std::{ + fmt, + future::Future, + net::{Ipv4Addr, SocketAddr}, + pin::Pin, + sync::Arc, + time::Duration, +}; + +use async_trait::async_trait; +use tokio::sync::oneshot; + +use crate::events::{CoreEvent, CoreEventSink}; + +const UPNP_RENEW_INTERVAL: Duration = Duration::from_secs(240); + +pub(crate) trait UdpPortMappingLease: Send + Sync + fmt::Debug { + fn public_addr_resolved(&self, _mapped_addr: SocketAddr) {} +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum UdpPortMappingBackend { + Igd, + NatPmp, +} + +impl UdpPortMappingBackend { + pub fn name(self) -> &'static str { + match self { + Self::Igd => "igd", + Self::NatPmp => "nat-pmp", + } + } +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum UdpPortMappingAttemptPhase { + Discovery, + Establishment, +} + +#[derive(Debug)] +pub struct UdpPortMappingAttemptError { + phase: UdpPortMappingAttemptPhase, + source: anyhow::Error, +} + +impl UdpPortMappingAttemptError { + pub fn discovery(source: impl Into) -> Self { + Self { + phase: UdpPortMappingAttemptPhase::Discovery, + source: source.into(), + } + } + + pub fn establishment(source: impl Into) -> Self { + Self { + phase: UdpPortMappingAttemptPhase::Establishment, + source: source.into(), + } + } + + pub(crate) fn phase(&self) -> UdpPortMappingAttemptPhase { + self.phase + } + + pub(crate) fn source(&self) -> &anyhow::Error { + &self.source + } +} + +impl fmt::Display for UdpPortMappingAttemptError { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + self.source.fmt(f) + } +} + +impl std::error::Error for UdpPortMappingAttemptError { + fn source(&self) -> Option<&(dyn std::error::Error + 'static)> { + Some(self.source.as_ref()) + } +} + +#[async_trait] +pub trait ActiveUdpPortMapping: Send + Sync + fmt::Debug { + fn backend(&self) -> UdpPortMappingBackend; + + fn local_addr(&self) -> SocketAddr; + + fn gateway_external_port(&self) -> u16; + + async fn renew(&self) -> anyhow::Result<()>; + + async fn remove(&self) -> anyhow::Result<()>; +} + +pub type UdpPortMappingLifecycle = Pin + Send + 'static>>; + +#[async_trait] +pub trait UdpPortMappingPlatform: Send + Sync + 'static { + async fn establish_udp_port_mapping( + &self, + backend: UdpPortMappingBackend, + local_listener: &url::Url, + ) -> Result, UdpPortMappingAttemptError>; + + fn spawn_udp_port_mapping_lifecycle( + &self, + _local_listener: url::Url, + lifecycle: UdpPortMappingLifecycle, + ) { + tokio::spawn(lifecycle); + } +} + +struct ManagedUdpPortMappingLease { + events: Arc, + local_listener: url::Url, + backend: UdpPortMappingBackend, + gateway_external_port: u16, + stop_tx: Option>, +} + +impl fmt::Debug for ManagedUdpPortMappingLease { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.debug_struct("UdpPortMappingLease") + .field("backend", &self.backend.name()) + .field("gateway_external_port", &self.gateway_external_port) + .finish() + } +} + +impl Drop for ManagedUdpPortMappingLease { + fn drop(&mut self) { + if let Some(stop_tx) = self.stop_tx.take() { + let _ = stop_tx.send(()); + } + } +} + +impl UdpPortMappingLease for ManagedUdpPortMappingLease { + fn public_addr_resolved(&self, mapped_addr: SocketAddr) { + self.events.emit(CoreEvent::UdpPortMappingEstablished { + local_listener: self.local_listener.clone(), + mapped_listener: udp_url(mapped_addr), + backend: self.backend.name().to_owned(), + }); + tracing::info!( + local_listener = %self.local_listener, + backend = self.backend.name(), + gateway_external_port = self.gateway_external_port, + stun_mapped_addr = %mapped_addr, + "udp public addr resolved after port mapping" + ); + } +} + +pub(crate) async fn start_udp_port_mapping( + platform: Arc, + events: Arc, + local_listener: &url::Url, +) -> anyhow::Result>> { + if !should_map_udp_listener(local_listener) { + return Ok(None); + } + + let mapping = discover_udp_port_mapping(platform.as_ref(), local_listener).await?; + let backend = mapping.backend(); + let gateway_external_port = mapping.gateway_external_port(); + tracing::info!( + %local_listener, + backend = backend.name(), + local_addr = %mapping.local_addr(), + gateway_external_port, + "udp port mapping established" + ); + + let (stop_tx, stop_rx) = oneshot::channel(); + platform.spawn_udp_port_mapping_lifecycle( + local_listener.clone(), + Box::pin(run_udp_port_mapping_lifecycle( + local_listener.clone(), + mapping, + stop_rx, + )), + ); + + Ok(Some(Box::new(ManagedUdpPortMappingLease { + events, + local_listener: local_listener.clone(), + backend, + gateway_external_port, + stop_tx: Some(stop_tx), + }))) +} + +async fn discover_udp_port_mapping( + platform: &dyn UdpPortMappingPlatform, + local_listener: &url::Url, +) -> anyhow::Result> { + let igd_error = match platform + .establish_udp_port_mapping(UdpPortMappingBackend::Igd, local_listener) + .await + { + Ok(mapping) => return Ok(mapping), + Err(error) => error, + }; + match igd_error.phase() { + UdpPortMappingAttemptPhase::Discovery => tracing::debug!( + igd_err = ?igd_error.source(), + %local_listener, + "igd gateway discovery failed, retry with nat-pmp" + ), + UdpPortMappingAttemptPhase::Establishment => tracing::debug!( + igd_err = ?igd_error.source(), + %local_listener, + "igd udp port mapping failed, retry with nat-pmp" + ), + } + + match platform + .establish_udp_port_mapping(UdpPortMappingBackend::NatPmp, local_listener) + .await + { + Ok(mapping) => Ok(mapping), + Err(nat_pmp_error) => Err(combined_mapping_error( + local_listener, + igd_error, + nat_pmp_error, + )), + } +} + +fn combined_mapping_error( + local_listener: &url::Url, + igd_error: UdpPortMappingAttemptError, + nat_pmp_error: UdpPortMappingAttemptError, +) -> anyhow::Error { + let igd_label = match igd_error.phase() { + UdpPortMappingAttemptPhase::Discovery => "igd discovery error", + UdpPortMappingAttemptPhase::Establishment => "igd error", + }; + let nat_pmp_label = match nat_pmp_error.phase() { + UdpPortMappingAttemptPhase::Discovery => "nat-pmp discovery error", + UdpPortMappingAttemptPhase::Establishment => "nat-pmp error", + }; + anyhow::anyhow!( + "udp port mapping failed for {local_listener}: {igd_label}: {}; {nat_pmp_label}: {}", + igd_error.source(), + nat_pmp_error.source(), + ) +} + +async fn run_udp_port_mapping_lifecycle( + local_listener: url::Url, + mapping: Box, + mut stop_rx: oneshot::Receiver<()>, +) { + loop { + tokio::select! { + _ = tokio::time::sleep(UPNP_RENEW_INTERVAL) => { + if let Err(error) = mapping.renew().await { + tracing::warn!( + err = ?error, + %local_listener, + backend = mapping.backend().name(), + gateway_external_port = mapping.gateway_external_port(), + "failed to renew udp port mapping" + ); + } + } + _ = &mut stop_rx => break, + } + } + + if let Err(error) = mapping.remove().await { + tracing::debug!( + err = ?error, + %local_listener, + backend = mapping.backend().name(), + gateway_external_port = mapping.gateway_external_port(), + "failed to remove udp port mapping" + ); + } +} + +pub(crate) fn should_map_udp_listener(local_listener: &url::Url) -> bool { + if local_listener.scheme() != "udp" { + return false; + } + + let Some(host) = listener_ipv4_host(local_listener) else { + return false; + }; + + if host.is_loopback() || host.is_broadcast() { + return false; + } + + host.is_unspecified() || host.is_private() || host.is_link_local() +} + +fn listener_ipv4_host(local_listener: &url::Url) -> Option { + local_listener.host_str()?.parse().ok() +} + +fn udp_url(addr: SocketAddr) -> url::Url { + let mut url = url::Url::parse("udp://0.0.0.0").expect("static UDP URL should be valid"); + url.set_ip_host(addr.ip()) + .expect("socket IP should be a valid URL host"); + url.set_port(Some(addr.port())) + .expect("UDP URL should accept a port"); + url +} + +#[cfg(test)] +mod tests { + use std::sync::{ + Mutex, + atomic::{AtomicUsize, Ordering}, + }; + + use super::*; + + #[derive(Debug)] + struct MockMapping { + backend: UdpPortMappingBackend, + removals: Arc, + } + + #[async_trait] + impl ActiveUdpPortMapping for MockMapping { + fn backend(&self) -> UdpPortMappingBackend { + self.backend + } + + fn local_addr(&self) -> SocketAddr { + "192.168.1.5:11010".parse().unwrap() + } + + fn gateway_external_port(&self) -> u16 { + 41010 + } + + async fn renew(&self) -> anyhow::Result<()> { + Ok(()) + } + + async fn remove(&self) -> anyhow::Result<()> { + self.removals.fetch_add(1, Ordering::SeqCst); + Ok(()) + } + } + + struct MockPlatform { + attempts: Mutex>, + igd_phase: Option, + removals: Arc, + } + + #[async_trait] + impl UdpPortMappingPlatform for MockPlatform { + async fn establish_udp_port_mapping( + &self, + backend: UdpPortMappingBackend, + _local_listener: &url::Url, + ) -> Result, UdpPortMappingAttemptError> { + self.attempts.lock().unwrap().push(backend); + if backend == UdpPortMappingBackend::Igd + && let Some(phase) = self.igd_phase + { + return Err(match phase { + UdpPortMappingAttemptPhase::Discovery => { + UdpPortMappingAttemptError::discovery(anyhow::anyhow!("no igd")) + } + UdpPortMappingAttemptPhase::Establishment => { + UdpPortMappingAttemptError::establishment(anyhow::anyhow!("igd denied")) + } + }); + } + Ok(Box::new(MockMapping { + backend, + removals: self.removals.clone(), + })) + } + } + + #[test] + fn mapping_requires_private_or_unspecified_ipv4_listener() { + assert!(should_map_udp_listener( + &"udp://0.0.0.0:11010".parse().unwrap() + )); + assert!(should_map_udp_listener( + &"udp://192.168.1.10:11010".parse().unwrap() + )); + assert!(!should_map_udp_listener( + &"udp://127.0.0.1:11010".parse().unwrap() + )); + assert!(!should_map_udp_listener( + &"udp://8.8.8.8:11010".parse().unwrap() + )); + assert!(!should_map_udp_listener( + &"tcp://0.0.0.0:11010".parse().unwrap() + )); + } + + #[tokio::test] + async fn falls_back_from_igd_to_nat_pmp_and_removes_on_drop() { + let removals = Arc::new(AtomicUsize::new(0)); + let platform = Arc::new(MockPlatform { + attempts: Mutex::new(Vec::new()), + igd_phase: Some(UdpPortMappingAttemptPhase::Discovery), + removals: removals.clone(), + }); + + let lease = start_udp_port_mapping( + platform.clone(), + Arc::new(()), + &"udp://0.0.0.0:11010".parse().unwrap(), + ) + .await + .unwrap() + .unwrap(); + assert_eq!( + *platform.attempts.lock().unwrap(), + vec![UdpPortMappingBackend::Igd, UdpPortMappingBackend::NatPmp] + ); + + drop(lease); + tokio::time::timeout(Duration::from_secs(1), async { + while removals.load(Ordering::SeqCst) == 0 { + tokio::task::yield_now().await; + } + }) + .await + .unwrap(); + assert_eq!(removals.load(Ordering::SeqCst), 1); + } +} diff --git a/easytier-core/src/connectivity/hole_punch/tcp.rs b/easytier-core/src/connectivity/hole_punch/tcp.rs new file mode 100644 index 00000000..dd3c6d5b --- /dev/null +++ b/easytier-core/src/connectivity/hole_punch/tcp.rs @@ -0,0 +1,1104 @@ +#![cfg_attr(not(feature = "tcp-hole-punch"), allow(dead_code))] + +use std::{ + net::{IpAddr, Ipv4Addr, Ipv6Addr, SocketAddr}, + sync::{ + Arc, Mutex, + atomic::{AtomicBool, Ordering}, + }, + time::Duration, +}; + +use anyhow::Context as _; +use async_trait::async_trait; +use dashmap::DashMap; +use quanta::Instant; +use rand::Rng as _; +use tokio::task::JoinSet; +use tokio_util::task::AbortOnDropHandle; + +use crate::{ + config::{P2pPolicyFlags, PeerId}, + connectivity::{ + hole_punch::{ + HolePunchRpcRegistry, HolePunchTunnelSink, + policy::{BackOff, should_background_p2p_with_peer, should_try_p2p_with_peer}, + }, + protocol::{ClientProtocolUpgrader, ServerProtocolUpgrade, ServerProtocolUpgrader}, + stun::StunInfoProvider, + transport::ConnectedTransport, + }, + foundation::task::{ + ExternalTaskSignal, PeerTaskLauncher, PeerTaskManager, reap_joinset_background, + }, + proto::{ + common::{NatType, PeerFeatureFlag}, + peer_rpc::{ + TcpHolePunchRequest, TcpHolePunchResponse, TcpHolePunchRpc, TcpHolePunchRpcServer, + }, + rpc_types::{self, controller::BaseController}, + }, + socket::{ + IpVersion, SocketContext, + tcp::{ + TcpBindOptions, TcpConnectOptions, TcpListenOptions, VirtualTcpListener, + VirtualTcpListenerFactory, VirtualTcpSocketFactory, + }, + }, +}; + +pub trait TcpHolePunchHost: VirtualTcpListenerFactory + VirtualTcpSocketFactory {} + +impl TcpHolePunchHost for T where T: VirtualTcpListenerFactory + VirtualTcpSocketFactory {} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum TcpHolePunchAdmission { + Client, + Server, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct TcpPunchCandidate { + pub peer_id: PeerId, + pub tcp_nat_type: NatType, + pub feature_flag: Option, + pub has_direct_connection: bool, + pub has_recent_traffic: bool, +} + +/// Narrow peer-graph view required by the TCP hole-punch engine. +/// +/// Implemented only by the sealed peer adapter in `super::peer_adapters`. +#[async_trait] +pub trait TcpHolePunchPeerSource: Send + Sync + 'static { + fn local_peer_id(&self) -> PeerId; + + fn p2p_policy_flags(&self) -> P2pPolicyFlags; + + fn tcp_hole_punching_disabled(&self) -> bool; + + fn p2p_demand_notify(&self) -> Arc; + + async fn candidates(&self) -> Vec; + + fn rpc_stub( + &self, + dst_peer_id: PeerId, + ) -> Box + Send + Sync + 'static>; +} + +#[derive(Debug, thiserror::Error)] +pub enum TcpHolePunchTransportError { + #[error("TCP hole-punch protocol upgrade failed")] + Upgrade(#[source] anyhow::Error), + #[error("TCP hole-punch tunnel admission failed")] + Admission(#[source] anyhow::Error), +} + +#[async_trait] +pub trait TcpHolePunchTransportSink: Send + Sync + 'static { + type ConnectedSocket; + type AcceptedSocket; + + async fn add_connected_transport( + &self, + socket: Self::ConnectedSocket, + requested_url: url::Url, + admission: TcpHolePunchAdmission, + ) -> Result<(), TcpHolePunchTransportError>; + + async fn add_accepted_transport( + &self, + socket: Self::AcceptedSocket, + local_url: url::Url, + ) -> Result<(), TcpHolePunchTransportError>; +} + +pub struct ProtocolTcpHolePunchTransportSink { + client_protocol: Arc>, + server_protocol: Arc>, + tunnel_sink: Arc, +} + +impl + ProtocolTcpHolePunchTransportSink +{ + pub fn new( + client_protocol: Arc>, + server_protocol: Arc>, + tunnel_sink: Arc, + ) -> Self { + Self { + client_protocol, + server_protocol, + tunnel_sink, + } + } +} + +#[async_trait] +impl TcpHolePunchTransportSink + for ProtocolTcpHolePunchTransportSink +where + ConnectedSocket: Send + 'static, + AcceptedSocket: Send + 'static, + T: HolePunchTunnelSink, +{ + type ConnectedSocket = ConnectedSocket; + type AcceptedSocket = AcceptedSocket; + + async fn add_connected_transport( + &self, + socket: ConnectedSocket, + requested_url: url::Url, + admission: TcpHolePunchAdmission, + ) -> Result<(), TcpHolePunchTransportError> { + let tunnel = self + .client_protocol + .upgrade_client(ConnectedTransport::Tcp(socket), requested_url) + .await + .map_err(TcpHolePunchTransportError::Upgrade)?; + match admission { + TcpHolePunchAdmission::Client => self.tunnel_sink.add_client_tunnel(tunnel).await, + TcpHolePunchAdmission::Server => self.tunnel_sink.add_server_tunnel(tunnel).await, + } + .map_err(TcpHolePunchTransportError::Admission) + } + + async fn add_accepted_transport( + &self, + socket: AcceptedSocket, + local_url: url::Url, + ) -> Result<(), TcpHolePunchTransportError> { + let upgrade = self + .server_protocol + .upgrade_tcp(socket, local_url) + .await + .map_err(TcpHolePunchTransportError::Upgrade)?; + let ServerProtocolUpgrade::Tunnel(tunnel) = upgrade else { + return Err(TcpHolePunchTransportError::Upgrade(anyhow::anyhow!( + "TCP hole-punch protocol returned a tunnel acceptor" + ))); + }; + self.tunnel_sink + .add_server_tunnel(tunnel) + .await + .map_err(TcpHolePunchTransportError::Admission) + } +} + +type ConnectedTcpSocket = ::Socket; +type AcceptedTcpSocket = + <::Listener as VirtualTcpListener>::Socket; + +pub(super) type TcpHolePunchTransportSinkFor = dyn TcpHolePunchTransportSink< + ConnectedSocket = ConnectedTcpSocket, + AcceptedSocket = AcceptedTcpSocket, + >; + +fn bind_addr_for_port(port: u16, is_v6: bool) -> SocketAddr { + if is_v6 { + SocketAddr::new(IpAddr::V6(Ipv6Addr::UNSPECIFIED), port) + } else { + SocketAddr::new(IpAddr::V4(Ipv4Addr::UNSPECIFIED), port) + } +} + +pub async fn select_local_port( + host: &H, + context: SocketContext, + is_v6: bool, +) -> anyhow::Result +where + H: VirtualTcpListenerFactory, +{ + let bind_addr = bind_addr_for_port(0, is_v6); + tracing::trace!(?bind_addr, is_v6, "tcp hole punch select local port"); + let context = context.with_ip_version(if is_v6 { IpVersion::V6 } else { IpVersion::V4 }); + let listener = host + .bind_tcp( + TcpListenOptions::hole_punch(bind_addr).with_bind( + TcpBindOptions::default() + .with_context(context) + .with_local_addr(Some(bind_addr)), + ), + ) + .await?; + let port = listener.local_addr()?.port(); + tracing::debug!(?bind_addr, port, "tcp hole punch selected local port"); + Ok(port) +} + +// TCP supports simultaneous connect, so both peers may dial from the mapped port. +pub async fn try_connect_to_remote( + host: Arc, + transport_sink: Arc< + dyn TcpHolePunchTransportSink< + ConnectedSocket = ::Socket, + AcceptedSocket = AcceptedSocket, + >, + >, + remote_mapped_addr: SocketAddr, + local_port: u16, + context: SocketContext, + admission: TcpHolePunchAdmission, + max_attempts: u32, +) -> anyhow::Result<()> +where + H: VirtualTcpSocketFactory, + AcceptedSocket: 'static, +{ + tracing::info!( + ?remote_mapped_addr, + local_port, + "tcp hole punch server start connect loop" + ); + + let bind_addr = bind_addr_for_port(local_port, remote_mapped_addr.is_ipv6()); + let context = context.with_ip_version(if remote_mapped_addr.is_ipv6() { + IpVersion::V6 + } else { + IpVersion::V4 + }); + let requested_url: url::Url = format!("tcp://{remote_mapped_addr}").parse().unwrap(); + + let start = crate::foundation::time::Instant::now(); + let mut attempts = 0_u32; + while start.elapsed() < Duration::from_secs(10) && attempts < max_attempts { + attempts = attempts.wrapping_add(1); + let bind = TcpBindOptions::default() + .with_context(context.clone()) + .with_local_addr(Some(bind_addr)) + .with_only_v6(true); + let options = + TcpConnectOptions::hole_punch(remote_mapped_addr, Some(bind_addr)).with_bind(bind); + if let Ok(Ok(socket)) = + crate::foundation::time::timeout(Duration::from_secs(3), host.connect_tcp(options)) + .await + { + let admission_result = transport_sink + .add_connected_transport(socket, requested_url.clone(), admission) + .await; + match admission_result { + Ok(()) => {} + Err(TcpHolePunchTransportError::Upgrade(error)) => return Err(error), + Err(TcpHolePunchTransportError::Admission(error)) => { + tracing::error!( + ?remote_mapped_addr, + local_port, + attempts, + ?error, + "tcp hole punch server connected and added client tunnel failed" + ); + continue; + } + } + + tracing::info!( + ?remote_mapped_addr, + local_port, + attempts, + ?admission, + "tcp hole punch server connected and added tunnel" + ); + return Ok(()); + } + tracing::trace!( + ?remote_mapped_addr, + local_port, + attempts, + "tcp hole punch server connect attempt failed" + ); + let sleep_ms = rand::thread_rng().gen_range(10..100); + crate::foundation::time::sleep(Duration::from_millis(sleep_ms)).await; + } + + tracing::warn!( + ?remote_mapped_addr, + local_port, + attempts, + "tcp hole punch server connect loop timeout" + ); + + Err(anyhow::anyhow!( + "tcp hole punch server connect loop timeout" + )) +} + +pub async fn accept_connections( + listener: Arc, + transport_sink: Arc< + dyn TcpHolePunchTransportSink, + >, + dst_peer_id: PeerId, +) -> anyhow::Result<()> +where + L: VirtualTcpListener, + ConnectedSocket: 'static, +{ + loop { + match listener.accept().await { + Ok((socket, _)) => { + let local_url = format!("tcp://0.0.0.0:{}", listener.local_addr()?.port()) + .parse() + .unwrap(); + if let Err(error) = transport_sink + .add_accepted_transport(socket, local_url) + .await + { + tracing::error!(?error, "tcp hole punch transport admission error"); + continue; + } + + tracing::info!( + dst_peer_id, + "tcp hole punch initiator accepted and added server tunnel" + ); + } + Err(error) => { + tracing::error!(?error, "tcp hole punch accept error"); + } + } + } +} + +const BLACKLIST_TIMEOUT: Duration = Duration::from_secs(3600); + +fn fallback_listener_options( + socket_context: SocketContext, + bind_addr: std::net::SocketAddr, +) -> TcpListenOptions { + let bind = TcpBindOptions::default() + .with_context(socket_context.with_ip_version(IpVersion::V4)) + .with_local_addr(Some(bind_addr)) + .with_only_v6(true); + TcpListenOptions::hole_punch(bind_addr).with_bind(bind) +} + +struct TcpHolePunchBlacklist { + entries: DashMap, +} + +impl TcpHolePunchBlacklist { + fn new() -> Self { + Self { + entries: DashMap::new(), + } + } + + fn insert(&self, peer_id: PeerId) { + self.entries.insert(peer_id, Instant::now()); + } + + fn contains(&self, peer_id: PeerId) -> bool { + let active = self + .entries + .get(&peer_id) + .is_some_and(|inserted_at| inserted_at.elapsed() < BLACKLIST_TIMEOUT); + if !active { + self.entries.remove(&peer_id); + } + active + } + + fn cleanup(&self) { + self.entries + .retain(|_, inserted_at| inserted_at.elapsed() < BLACKLIST_TIMEOUT); + } +} + +fn handle_rpc_result( + result: Result, + dst_peer_id: PeerId, + blacklist: &TcpHolePunchBlacklist, +) -> Result { + match result { + Ok(result) => Ok(result), + Err(error) => { + if matches!(error, rpc_types::error::Error::InvalidServiceKey(_, _)) { + blacklist.insert(dst_peer_id); + } + Err(error) + } + } +} + +fn is_symmetric_tcp_nat(nat_type: NatType) -> bool { + matches!( + nat_type, + NatType::Symmetric | NatType::SymmetricEasyInc | NatType::SymmetricEasyDec + ) +} + +struct TcpHolePunchServer +where + H: TcpHolePunchHost, +{ + host: Arc, + stun: Arc, + socket_context: SocketContext, + transport_sink: Arc>, + tasks: Arc>>, + reaper: Mutex>>, + stopping: AtomicBool, +} + +impl TcpHolePunchServer +where + H: TcpHolePunchHost, +{ + fn new( + host: Arc, + stun: Arc, + socket_context: SocketContext, + transport_sink: Arc>, + ) -> Arc { + Arc::new(Self { + host, + stun, + socket_context, + transport_sink, + tasks: Arc::new(Mutex::new(JoinSet::new())), + reaper: Mutex::new(None), + stopping: AtomicBool::new(true), + }) + } + + fn start(&self) { + let mut reaper = self.reaper.lock().unwrap(); + if reaper.as_ref().is_some_and(|task| !task.is_finished()) { + return; + } + { + let _tasks = self.tasks.lock().unwrap(); + self.stopping.store(false, Ordering::Release); + } + reaper.replace(AbortOnDropHandle::new(tokio::spawn( + reap_joinset_background(self.tasks.clone(), "tcp hole punch"), + ))); + } + + fn begin_stop(&self) { + self.stopping.store(true, Ordering::Release); + } + + async fn stop(&self) { + let reaper = self.reaper.lock().unwrap().take(); + if let Some(reaper) = reaper { + reaper.abort(); + let _ = reaper.await; + } + let mut tasks = { + let mut task_slot = self.tasks.lock().unwrap(); + std::mem::replace(&mut *task_slot, JoinSet::new()) + }; + tasks.abort_all(); + while tasks.join_next().await.is_some() {} + } +} + +#[async_trait] +impl TcpHolePunchRpc for TcpHolePunchServer +where + H: TcpHolePunchHost, +{ + type Controller = BaseController; + + #[tracing::instrument(skip(self), fields(a_mapped_addr = ?input.connector_mapped_addr), err)] + async fn exchange_mapped_addr( + &self, + _controller: Self::Controller, + input: TcpHolePunchRequest, + ) -> rpc_types::error::Result { + let local_nat_type = + NatType::try_from(self.stun.get_stun_info().tcp_nat_type).unwrap_or(NatType::Unknown); + tracing::debug!(?local_nat_type, "tcp hole punch rpc received"); + if local_nat_type == NatType::Unknown { + tracing::warn!(?local_nat_type, "tcp hole punch rpc rejected (unknown)"); + return Err(anyhow::anyhow!("tcp nat type unknown not supported").into()); + } + + let remote_mapped_addr = input + .connector_mapped_addr + .ok_or_else(|| anyhow::anyhow!("connector_mapped_addr is required"))?; + let remote_mapped_addr: std::net::SocketAddr = remote_mapped_addr.into(); + let remote_ip = remote_mapped_addr.ip(); + if remote_ip.is_unspecified() || remote_ip.is_multicast() { + tracing::warn!( + ?remote_mapped_addr, + "tcp hole punch rpc invalid connector addr" + ); + return Err(anyhow::anyhow!("connector_mapped_addr is malformed").into()); + } + + let local_port = select_local_port( + self.host.as_ref(), + self.socket_context.clone(), + remote_mapped_addr.is_ipv6(), + ) + .await?; + let local_mapped_addr = self + .stun + .get_tcp_port_mapping(local_port) + .await + .context("failed to get tcp port mapping")?; + + tracing::info!( + ?remote_mapped_addr, + local_port, + ?local_mapped_addr, + "tcp hole punch rpc responding with listener mapped addr and start connecting" + ); + + let host = self.host.clone(); + let socket_context = self.socket_context.clone(); + let transport_sink = self.transport_sink.clone(); + let mut tasks = self.tasks.lock().unwrap(); + if self.stopping.load(Ordering::Acquire) { + return Err(rpc_types::error::Error::Shutdown); + } + tasks.spawn(async move { + let _ = try_connect_to_remote( + host, + transport_sink, + remote_mapped_addr, + local_port, + socket_context, + TcpHolePunchAdmission::Client, + 5, + ) + .await; + }); + + Ok(TcpHolePunchResponse { + listener_mapped_addr: Some(local_mapped_addr.into()), + }) + } +} + +struct TcpHolePunchConnectorData +where + H: TcpHolePunchHost, + P: TcpHolePunchPeerSource, +{ + host: Arc, + stun: Arc, + socket_context: SocketContext, + peer_source: Arc

, + transport_sink: Arc>, + blacklist: TcpHolePunchBlacklist, +} + +impl TcpHolePunchConnectorData +where + H: TcpHolePunchHost, + P: TcpHolePunchPeerSource, +{ + async fn punch_as_initiator(self: Arc, dst_peer_id: PeerId) -> anyhow::Result<()> { + let mut backoff = BackOff::new(vec![1000, 1000, 4000, 8000]); + + loop { + backoff.sleep_for_next_backoff().await; + if self.do_punch_as_initiator(dst_peer_id).await.is_ok() { + break; + } + + if self.blacklist.contains(dst_peer_id) { + tracing::warn!( + dst_peer_id, + "tcp hole punch initiator skipped (blacklisted)" + ); + break; + } + } + + Ok(()) + } + + #[tracing::instrument(skip(self), fields(dst_peer_id), err)] + async fn do_punch_as_initiator(&self, dst_peer_id: PeerId) -> anyhow::Result<()> { + let local_nat_type = + NatType::try_from(self.stun.get_stun_info().tcp_nat_type).unwrap_or(NatType::Unknown); + tracing::debug!(?local_nat_type, "tcp hole punch initiator start"); + if is_symmetric_tcp_nat(local_nat_type) || local_nat_type == NatType::Unknown { + tracing::debug!("tcp hole punch initiator skipped (symmetric)"); + return Ok(()); + } + + let local_port = + select_local_port(self.host.as_ref(), self.socket_context.clone(), false).await?; + let local_mapped_addr = self + .stun + .get_tcp_port_mapping(local_port) + .await + .context("failed to get tcp port mapping")?; + + tracing::info!( + dst_peer_id, + local_port, + ?local_mapped_addr, + "tcp hole punch initiator got mapped addr, start rpc exchange" + ); + + let rpc_stub = self.peer_source.rpc_stub(dst_peer_id); + let response = rpc_stub + .exchange_mapped_addr( + BaseController { + timeout_ms: 6000, + ..Default::default() + }, + TcpHolePunchRequest { + connector_mapped_addr: Some(local_mapped_addr.into()), + }, + ) + .await; + let response = handle_rpc_result(response, dst_peer_id, &self.blacklist)?; + let remote_mapped_addr = response + .listener_mapped_addr + .ok_or_else(|| anyhow::anyhow!("listener_mapped_addr is required"))?; + let remote_mapped_addr = remote_mapped_addr.into(); + tracing::info!( + dst_peer_id, + ?remote_mapped_addr, + "tcp hole punch initiator rpc returned" + ); + + if try_connect_to_remote( + self.host.clone(), + self.transport_sink.clone(), + remote_mapped_addr, + local_port, + self.socket_context.clone(), + TcpHolePunchAdmission::Server, + 1, + ) + .await + .is_ok() + { + tracing::info!( + dst_peer_id, + local_port, + ?remote_mapped_addr, + "tcp hole punch initiator connected to remote mapped addr with simultaneous connection" + ); + return Ok(()); + } + + tracing::debug!( + dst_peer_id, + local_port, + ?remote_mapped_addr, + "tcp hole punch initiator sent syn to remote mapped addr" + ); + + let bind_addr = + std::net::SocketAddr::new(std::net::Ipv4Addr::UNSPECIFIED.into(), local_port); + let listener = self + .host + .bind_tcp(fallback_listener_options( + self.socket_context.clone(), + bind_addr, + )) + .await?; + tracing::info!( + dst_peer_id, + local_port, + local_addr = ?listener.local_addr()?, + "tcp hole punch initiator listening" + ); + + crate::foundation::time::timeout( + Duration::from_secs(10), + accept_connections(listener, self.transport_sink.clone(), dst_peer_id), + ) + .await??; + + tracing::info!( + dst_peer_id, + "tcp hole punch initiator accepted and added server tunnel" + ); + Ok(()) + } + + async fn collect_peers_need_task(&self) -> Vec { + let local_nat_type = + NatType::try_from(self.stun.get_stun_info().tcp_nat_type).unwrap_or(NatType::Unknown); + if is_symmetric_tcp_nat(local_nat_type) || local_nat_type == NatType::Unknown { + tracing::trace!( + ?local_nat_type, + "tcp hole punch task collect skipped (symmetric)" + ); + return Vec::new(); + } + + self.blacklist.cleanup(); + let policy = self.peer_source.p2p_policy_flags(); + let local_peer_id = self.peer_source.local_peer_id(); + let mut peers_to_connect = Vec::new(); + for candidate in self.peer_source.candidates().await { + let static_allowed = should_background_p2p_with_peer( + candidate.feature_flag.as_ref(), + false, + policy.lazy_p2p, + policy.disable_p2p, + policy.need_p2p, + ); + let dynamic_allowed = should_try_p2p_with_peer( + candidate.feature_flag.as_ref(), + false, + policy.disable_p2p, + policy.need_p2p, + ) && candidate.has_recent_traffic; + if !static_allowed && !dynamic_allowed { + continue; + } + + let peer_id = candidate.peer_id; + if peer_id == local_peer_id { + tracing::trace!(peer_id, "tcp hole punch task collect skip self"); + continue; + } + if self.blacklist.contains(peer_id) { + tracing::debug!(peer_id, "tcp hole punch task collect skip blacklisted"); + continue; + } + if candidate.has_direct_connection { + tracing::trace!(peer_id, "tcp hole punch task collect skip already has peer"); + continue; + } + + let peer_nat_type = candidate.tcp_nat_type; + if peer_nat_type == NatType::Unknown { + tracing::debug!( + peer_id, + ?peer_nat_type, + "tcp hole punch task collect skip peer unknown" + ); + continue; + } + + tracing::info!( + peer_id, + local_peer_id, + ?local_nat_type, + ?peer_nat_type, + "tcp hole punch task collect add peer" + ); + peers_to_connect.push(peer_id); + } + peers_to_connect + } +} + +struct TcpHolePunchPeerTaskLauncher(Arc>) +where + H: TcpHolePunchHost, + P: TcpHolePunchPeerSource; + +impl Clone for TcpHolePunchPeerTaskLauncher +where + H: TcpHolePunchHost, + P: TcpHolePunchPeerSource, +{ + fn clone(&self) -> Self { + Self(self.0.clone()) + } +} + +#[async_trait] +impl PeerTaskLauncher for TcpHolePunchPeerTaskLauncher +where + H: TcpHolePunchHost, + P: TcpHolePunchPeerSource, +{ + type CollectPeerItem = PeerId; + type TaskRet = (); + + async fn collect_peers_need_task(&self) -> Vec { + self.0.collect_peers_need_task().await + } + + async fn launch_task( + &self, + dst_peer_id: PeerId, + ) -> tokio::task::JoinHandle> { + let data = self.0.clone(); + tokio::spawn(async move { data.punch_as_initiator(dst_peer_id).await }) + } + + fn loop_interval_ms(&self) -> u64 { + 5000 + } +} + +pub struct TcpHolePunchConnector +where + H: TcpHolePunchHost, + P: TcpHolePunchPeerSource + HolePunchTunnelSink + HolePunchRpcRegistry, +{ + server: Arc>, + client: PeerTaskManager>, + peer_source: Arc

, +} + +impl TcpHolePunchConnector +where + H: TcpHolePunchHost, + P: TcpHolePunchPeerSource + HolePunchTunnelSink + HolePunchRpcRegistry, +{ + pub fn new( + peer_source: Arc

, + host: Arc, + stun: Arc, + socket_context: SocketContext, + client_protocol: Arc>>, + server_protocol: Arc>>, + ) -> Self { + let transport_sink: Arc> = + Arc::new(ProtocolTcpHolePunchTransportSink::new( + client_protocol, + server_protocol, + peer_source.clone(), + )); + let data = Arc::new(TcpHolePunchConnectorData { + host: host.clone(), + stun: stun.clone(), + socket_context: socket_context.clone(), + peer_source: peer_source.clone(), + transport_sink: transport_sink.clone(), + blacklist: TcpHolePunchBlacklist::new(), + }); + Self { + server: TcpHolePunchServer::new(host, stun, socket_context, transport_sink), + client: PeerTaskManager::new_with_external_signal( + TcpHolePunchPeerTaskLauncher(data), + Some(peer_source.p2p_demand_notify()), + ), + peer_source, + } + } + + pub fn run(&self) { + if self.peer_source.tcp_hole_punching_disabled() { + tracing::debug!("tcp hole punch disabled by runtime configuration"); + return; + } + self.server.start(); + self.peer_source + .register_rpc_service(TcpHolePunchRpcServer::new_arc(self.server.clone())); + self.client.start(); + } + + pub async fn stop(&self) { + self.client.stop().await; + self.server.begin_stop(); + self.peer_source + .unregister_rpc_service(TcpHolePunchRpcServer::new_arc(self.server.clone())); + self.server.stop().await; + } +} + +#[cfg(test)] +mod tests { + use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering}; + + use super::*; + use crate::tunnel::Tunnel; + + #[derive(Default)] + struct MockProtocols { + client_upgrades: AtomicUsize, + server_upgrades: AtomicUsize, + fail_client_upgrade: AtomicBool, + } + + #[async_trait] + impl ClientProtocolUpgrader<()> for MockProtocols { + fn supports_scheme(&self, scheme: &str) -> bool { + scheme == "tcp" + } + + async fn upgrade_client( + &self, + connected: ConnectedTransport<()>, + _requested_url: url::Url, + ) -> anyhow::Result> { + let ConnectedTransport::Tcp(()) = connected else { + anyhow::bail!("expected TCP transport"); + }; + if self.fail_client_upgrade.load(Ordering::Relaxed) { + anyhow::bail!("mock client upgrade failure"); + } + self.client_upgrades.fetch_add(1, Ordering::Relaxed); + Ok(crate::tunnel::ring::create_ring_tunnel_pair().0) + } + } + + #[async_trait] + impl ServerProtocolUpgrader<()> for MockProtocols { + fn supports_scheme(&self, scheme: &str) -> bool { + scheme == "tcp" + } + + async fn upgrade_tcp( + &self, + _socket: (), + _local_url: url::Url, + ) -> anyhow::Result { + self.server_upgrades.fetch_add(1, Ordering::Relaxed); + Ok(ServerProtocolUpgrade::Tunnel( + crate::tunnel::ring::create_ring_tunnel_pair().0, + )) + } + + async fn upgrade_udp( + &self, + _session: crate::socket::udp::UdpSession, + _local_url: url::Url, + _admission: Option, + ) -> anyhow::Result { + anyhow::bail!("unexpected UDP transport") + } + + async fn upgrade_byte_stream( + &self, + _socket: (), + _local_url: url::Url, + _remote_url: Option, + ) -> anyhow::Result { + anyhow::bail!("unexpected byte stream") + } + } + + #[derive(Default)] + struct MockTunnelSink { + clients: AtomicUsize, + servers: AtomicUsize, + fail_client_admission: AtomicBool, + } + + #[async_trait] + impl HolePunchTunnelSink for MockTunnelSink { + async fn add_client_tunnel(&self, _tunnel: Box) -> anyhow::Result<()> { + if self.fail_client_admission.load(Ordering::Relaxed) { + anyhow::bail!("mock client admission failure"); + } + self.clients.fetch_add(1, Ordering::Relaxed); + Ok(()) + } + + async fn add_server_tunnel(&self, _tunnel: Box) -> anyhow::Result<()> { + self.servers.fetch_add(1, Ordering::Relaxed); + Ok(()) + } + } + + #[tokio::test] + async fn protocol_sink_upgrades_before_tcp_hole_punch_admission() { + let protocols = Arc::new(MockProtocols::default()); + let tunnel_sink = Arc::new(MockTunnelSink::default()); + let sink = ProtocolTcpHolePunchTransportSink::new( + protocols.clone(), + protocols.clone(), + tunnel_sink.clone(), + ); + let url = url::Url::parse("tcp://198.51.100.1:11010").unwrap(); + + sink.add_connected_transport((), url.clone(), TcpHolePunchAdmission::Client) + .await + .unwrap(); + sink.add_connected_transport((), url.clone(), TcpHolePunchAdmission::Server) + .await + .unwrap(); + sink.add_accepted_transport((), url).await.unwrap(); + + assert_eq!(protocols.client_upgrades.load(Ordering::Relaxed), 2); + assert_eq!(protocols.server_upgrades.load(Ordering::Relaxed), 1); + assert_eq!(tunnel_sink.clients.load(Ordering::Relaxed), 1); + assert_eq!(tunnel_sink.servers.load(Ordering::Relaxed), 2); + + protocols.fail_client_upgrade.store(true, Ordering::Relaxed); + assert!(matches!( + sink.add_connected_transport( + (), + url::Url::parse("tcp://198.51.100.1:11010").unwrap(), + TcpHolePunchAdmission::Client, + ) + .await, + Err(TcpHolePunchTransportError::Upgrade(_)) + )); + protocols + .fail_client_upgrade + .store(false, Ordering::Relaxed); + tunnel_sink + .fail_client_admission + .store(true, Ordering::Relaxed); + assert!(matches!( + sink.add_connected_transport( + (), + url::Url::parse("tcp://198.51.100.1:11010").unwrap(), + TcpHolePunchAdmission::Client, + ) + .await, + Err(TcpHolePunchTransportError::Admission(_)) + )); + } + + #[test] + fn bind_address_tracks_requested_family_and_port() { + assert_eq!( + bind_addr_for_port(1234, false), + "0.0.0.0:1234".parse().unwrap() + ); + assert_eq!(bind_addr_for_port(4321, true), "[::]:4321".parse().unwrap()); + } + + #[test] + fn symmetric_tcp_nat_variants_are_ineligible_initiators() { + assert!(is_symmetric_tcp_nat(NatType::Symmetric)); + assert!(is_symmetric_tcp_nat(NatType::SymmetricEasyInc)); + assert!(is_symmetric_tcp_nat(NatType::SymmetricEasyDec)); + assert!(!is_symmetric_tcp_nat(NatType::PortRestricted)); + assert!(!is_symmetric_tcp_nat(NatType::Unknown)); + } + + #[test] + fn blacklist_tracks_and_cleans_entries() { + let blacklist = TcpHolePunchBlacklist::new(); + assert!(!blacklist.contains(7)); + blacklist.insert(7); + assert!(blacklist.contains(7)); + blacklist.cleanup(); + assert!(blacklist.contains(7)); + } + + #[test] + fn fallback_listener_normalizes_context_to_ipv4() { + let bind_addr = "0.0.0.0:23333".parse().unwrap(); + let context = SocketContext::default() + .with_socket_mark(Some(0)) + .with_netns(Some(crate::socket::NetNamespace::new("test-netns"))); + + let options = fallback_listener_options(context, bind_addr); + + assert_eq!( + options.purpose, + crate::socket::tcp::TcpListenPurpose::HolePunch + ); + assert_eq!(options.bind.local_addr, Some(bind_addr)); + assert_eq!(options.bind.context.ip_version, IpVersion::V4); + assert_eq!(options.bind.context.socket_mark, Some(0)); + assert_eq!( + options + .bind + .context + .netns + .as_ref() + .map(|netns| netns.token()), + Some("test-netns") + ); + assert!(options.bind.only_v6); + assert_eq!(options.bind.reuse_addr, None); + assert!(!options.bind.reuse_port); + } +} diff --git a/easytier-core/src/connectivity/hole_punch/udp/binding.rs b/easytier-core/src/connectivity/hole_punch/udp/binding.rs new file mode 100644 index 00000000..72691a32 --- /dev/null +++ b/easytier-core/src/connectivity/hole_punch/udp/binding.rs @@ -0,0 +1,802 @@ +//! Core-owned UDP hole-punch socket/session runtime. + +use std::{ + net::{IpAddr, Ipv4Addr, SocketAddr, SocketAddrV4}, + sync::{Arc, Weak}, +}; + +use anyhow::Context as _; +use async_trait::async_trait; + +use crate::{ + connectivity::{ + direct::DirectConnectorHost, + hole_punch::port_mapping::{UdpPortMappingPlatform, start_udp_port_mapping}, + hole_punch::{HolePunchRpcRegistry, HolePunchTunnelSink}, + protocol::ClientProtocolUpgrader, + stun::{StunInfoProvider, StunSocketMapper}, + }, + proto::peer_rpc::UdpHolePunchRpcServer, + socket::{ + IpVersion, ListenerConnectionCounter, SocketContext, + tcp::VirtualTcpSocketFactory, + udp::{ + UdpBindOptions, UdpSessionLayer, UdpSessionSocket, UdpSessionStunResponder, + VirtualUdpSocket, VirtualUdpSocketFactory, + }, + }, +}; + +use super::{ + ProtocolUdpHolePunchTransportSink, UdpHolePunchConnector, UdpHolePunchPeerSource, + UdpHolePunchRuntime, UdpPunchAcceptor, UdpPunchListener, UdpPunchSocket, UdpResolvedPublicAddr, + UdpSymPunchLock, + rpc::{PeerRpcUdpHolePunchSignaling, UdpHolePunchRpcEndpoint, UdpHolePunchRpcSource}, +}; + +async fn resolve_public_addr_with_policy( + stun: &dyn StunSocketMapper, + platform: Option>, + events: Arc, + socket: Arc, + local_listener: &url::Url, + disable_upnp: bool, +) -> anyhow::Result +where + S: VirtualUdpSocket + 'static, +{ + let port_mapping_lease = if disable_upnp { + None + } else if let Some(platform) = platform { + match start_udp_port_mapping(platform, events, local_listener).await { + Ok(lease) => lease, + Err(error) => { + tracing::warn!( + ?error, + %local_listener, + "failed to establish udp port mapping, fallback to stun-only public addr resolution" + ); + None + } + } + } else { + None + }; + + let mapped_addr = stun + .get_udp_port_mapping_with_socket(socket) + .await + .with_context(|| format!("resolve udp public addr for {local_listener}"))?; + if let Some(lease) = &port_mapping_lease { + lease.public_addr_resolved(mapped_addr); + } else { + tracing::debug!( + %local_listener, + stun_mapped_addr = %mapped_addr, + "udp public addr resolved without port mapping" + ); + } + + Ok(UdpResolvedPublicAddr { + mapped_addr, + port_mapping_lease, + }) +} + +fn managed_local_addr_error( + local_addr: SocketAddr, + is_local_virtual_ipv4: bool, + is_easytier_managed_ipv6: bool, +) -> Option<&'static str> { + match local_addr.ip() { + IpAddr::V4(_) if is_local_virtual_ipv4 => Some("local address is virtual ipv4"), + IpAddr::V6(_) if is_easytier_managed_ipv6 => Some("local address is easytier-managed ipv6"), + _ => None, + } +} + +type HostUdpSocket = ::Socket; +type HostTcpSocket = ::Socket; +type CoreUdpSessionLayer = UdpSessionLayer, H>; +type CoreUdpHolePunchTransportSink = ProtocolUdpHolePunchTransportSink, P>; +type CoreUdpHolePunchConnector = UdpHolePunchConnector< + P, + PeerRpcUdpHolePunchSignaling

, + CoreUdpHolePunchTransportSink, + CoreUdpHolePunchRuntime, +>; +type CoreUdpHolePunchEndpoint = + UdpHolePunchRpcEndpoint, CoreUdpHolePunchTransportSink>; + +pub(crate) struct CoreUdpHolePunchService +where + H: DirectConnectorHost, + P: UdpHolePunchPeerSource + + HolePunchTunnelSink + + UdpHolePunchRpcSource + + HolePunchRpcRegistry + + 'static, +{ + server: Arc>, + client: CoreUdpHolePunchConnector, + peer_source: Arc

, +} + +impl CoreUdpHolePunchService +where + H: DirectConnectorHost + Send + Sync + 'static, + HostUdpSocket: VirtualUdpSocket + 'static, + P: UdpHolePunchPeerSource + + HolePunchTunnelSink + + UdpHolePunchRpcSource + + HolePunchRpcRegistry + + 'static, +{ + pub(crate) fn new( + peer_source: Arc

, + host: Arc, + stun: Arc>>, + platform: Option>, + events: Arc, + socket_context: SocketContext, + protocol: Arc>>, + ) -> Self { + let stun_mapper = stun.clone(); + let stun_info: Arc = stun; + let transport_sink = Arc::new(ProtocolUdpHolePunchTransportSink::new( + protocol, + peer_source.clone(), + )); + let runtime = Arc::new(CoreUdpHolePunchRuntime::new( + host, + peer_source.clone(), + stun_mapper, + platform, + events, + socket_context, + )); + let sym_punch_lock = UdpSymPunchLock::default(); + let client = UdpHolePunchConnector::new( + peer_source.clone(), + Arc::new(PeerRpcUdpHolePunchSignaling::new(peer_source.clone())), + transport_sink.clone(), + runtime.clone(), + stun_info.clone(), + sym_punch_lock.clone(), + Some(peer_source.p2p_demand_notify()), + ); + + Self { + server: UdpHolePunchRpcEndpoint::new( + stun_info, + transport_sink, + sym_punch_lock, + runtime, + ), + client, + peer_source, + } + } + + pub(crate) async fn start(&self) -> anyhow::Result<()> { + if self + .peer_source + .p2p_policy_flags() + .disable_udp_hole_punching + { + return Ok(()); + } + + self.server.start().await; + self.peer_source + .register_rpc_service(UdpHolePunchRpcServer::new(Arc::downgrade(&self.server))); + self.client.run_as_client(); + Ok(()) + } + + pub(crate) async fn stop(&self) { + self.client.stop().await; + self.server.begin_stop(); + self.peer_source + .unregister_rpc_service(UdpHolePunchRpcServer::new(Arc::downgrade(&self.server))); + self.server.stop().await; + } +} + +struct CoreUdpPunchAcceptor +where + H: VirtualUdpSocketFactory, + HostUdpSocket: VirtualUdpSocket, +{ + layer: Arc>, +} + +#[async_trait] +impl UdpPunchAcceptor for CoreUdpPunchAcceptor +where + H: VirtualUdpSocketFactory + UdpSessionStunResponder> + Send + Sync + 'static, + HostUdpSocket: VirtualUdpSocket + 'static, +{ + async fn accept(&mut self) -> anyhow::Result { + let session = self.layer.accept().await?; + let remote_addr = session.peer_addr()?; + Ok(UdpPunchSocket::new( + session, + remote_addr, + self.layer.clone(), + )) + } +} + +struct CoreUdpPunchConnCounter +where + H: VirtualUdpSocketFactory, + HostUdpSocket: VirtualUdpSocket, +{ + layer: Weak>, +} + +impl std::fmt::Debug for CoreUdpPunchConnCounter +where + H: VirtualUdpSocketFactory, + HostUdpSocket: VirtualUdpSocket, +{ + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("CoreUdpPunchConnCounter") + .finish_non_exhaustive() + } +} + +impl ListenerConnectionCounter for CoreUdpPunchConnCounter +where + H: VirtualUdpSocketFactory + UdpSessionStunResponder> + Send + Sync + 'static, + HostUdpSocket: VirtualUdpSocket + 'static, +{ + fn get(&self) -> Option { + Some( + self.layer + .upgrade() + .map(|layer| layer.active_session_count() as u32) + .unwrap_or(0), + ) + } +} + +pub struct CoreUdpHolePunchRuntime +where + H: DirectConnectorHost, + P: UdpHolePunchPeerSource + 'static, +{ + host: Arc, + peer_source: Arc

, + stun: Arc>>, + platform: Option>, + events: Arc, + socket_context: SocketContext, +} + +impl CoreUdpHolePunchRuntime +where + H: DirectConnectorHost + Send + Sync + 'static, + HostUdpSocket: VirtualUdpSocket + 'static, + P: UdpHolePunchPeerSource + 'static, +{ + pub fn new( + host: Arc, + peer_source: Arc

, + stun: Arc>>, + platform: Option>, + events: Arc, + socket_context: SocketContext, + ) -> Self { + Self { + host, + peer_source, + stun, + platform, + events, + socket_context, + } + } + + fn session_layer(&self, socket: Arc>) -> Arc> { + Arc::new(UdpSessionLayer::new_with_stun_responder( + socket, + self.host.clone(), + )) + } + + async fn create_listener_with_mapping( + &self, + resolve_public_addr: bool, + port: Option, + ) -> anyhow::Result>> { + let bind = match port { + Some(port) => UdpBindOptions::hole_punch_candidate().with_local_addr(Some( + SocketAddr::V4(SocketAddrV4::new(Ipv4Addr::UNSPECIFIED, port)), + )), + None => UdpBindOptions::hole_punch_control(), + } + .with_context(self.socket_context.clone().with_ip_version(IpVersion::V4)); + let socket = self.host.bind_udp(bind).await?; + let local_port = socket.local_addr()?.port(); + let resolved = if resolve_public_addr { + self.resolve_public_addr(socket.clone()).await? + } else { + UdpResolvedPublicAddr { + mapped_addr: SocketAddr::V4(SocketAddrV4::new(Ipv4Addr::UNSPECIFIED, local_port)), + port_mapping_lease: None, + } + }; + + let layer = self.session_layer(socket.clone()); + let conn_counter = Arc::new(CoreUdpPunchConnCounter { + layer: Arc::downgrade(&layer), + }); + let acceptor = Box::new(CoreUdpPunchAcceptor { layer }); + + Ok(UdpPunchListener { + socket, + mapped_addr: resolved.mapped_addr, + conn_counter, + acceptor, + port_mapping_lease: resolved.port_mapping_lease, + }) + } + + async fn resolve_public_addr( + &self, + socket: Arc>, + ) -> anyhow::Result { + let local_port = socket.local_addr()?.port(); + let local_listener: url::Url = format!("udp://0.0.0.0:{local_port}").parse()?; + resolve_public_addr_with_policy( + self.stun.as_ref(), + self.platform.clone(), + self.events.clone(), + socket, + &local_listener, + self.peer_source.p2p_policy_flags().disable_upnp, + ) + .await + } + + async fn validate_socket_route( + &self, + context: SocketContext, + remote_addr: SocketAddr, + ) -> anyhow::Result<()> { + let local_addr = self + .host + .local_addr_for_remote(remote_addr, context) + .await?; + let is_local_virtual_ipv4 = match local_addr.ip() { + IpAddr::V4(ip) => self.peer_source.is_local_virtual_ip(&IpAddr::V4(ip)), + IpAddr::V6(_) => false, + }; + let is_easytier_managed_ipv6 = match local_addr.ip() { + IpAddr::V4(_) => false, + IpAddr::V6(ip) => self.peer_source.is_easytier_managed_ipv6(&ip).await, + }; + if let Some(error) = + managed_local_addr_error(local_addr, is_local_virtual_ipv4, is_easytier_managed_ipv6) + { + anyhow::bail!(error); + } + Ok(()) + } +} + +#[async_trait] +impl UdpHolePunchRuntime for CoreUdpHolePunchRuntime +where + H: DirectConnectorHost + Send + Sync + 'static, + HostUdpSocket: VirtualUdpSocket + 'static, + P: UdpHolePunchPeerSource + 'static, +{ + type Socket = HostUdpSocket; + + fn socket_context(&self) -> SocketContext { + self.socket_context.clone() + } + + async fn bind_udp(&self, options: UdpBindOptions) -> anyhow::Result> { + self.host.bind_udp(options).await + } + + async fn bind_direct_connect_udp(&self) -> anyhow::Result> { + self.host + .bind_udp( + UdpBindOptions::hole_punch_candidate() + .with_context(self.socket_context.clone().with_ip_version(IpVersion::V4)), + ) + .await + } + + async fn resolve_udp_public_addr( + &self, + socket: Arc, + ) -> anyhow::Result { + self.resolve_public_addr(socket).await + } + + async fn create_listener( + &self, + _prefer_port_mapping: bool, + ) -> anyhow::Result> { + self.create_listener_with_mapping(true, None).await + } + + async fn create_port_bound_listener( + &self, + port: u16, + ) -> anyhow::Result> { + self.create_listener_with_mapping(false, Some(port)).await + } + + async fn connect_with_socket( + &self, + socket: Arc, + remote: SocketAddr, + ) -> anyhow::Result { + self.validate_socket_route(socket.socket_context(), remote) + .await?; + let layer = self.session_layer(socket); + let session = layer.connect(remote).await?; + if session.peer_addr()? != remote { + tracing::debug!( + recv_addr = ?session.peer_addr()?, + ?remote, + "udp connect addr not match" + ); + } + Ok(UdpPunchSocket::new(session, remote, layer)) + } +} + +#[cfg(test)] +mod tests { + use std::{ + io, + net::{Ipv6Addr, SocketAddrV6}, + sync::{ + Mutex, + atomic::{AtomicUsize, Ordering}, + }, + }; + + use crate::{connectivity::stun::StunInfoProvider, proto::common::StunInfo}; + + use super::*; + + #[derive(Debug)] + struct MockSocket { + local_addr: SocketAddr, + } + + #[async_trait] + impl VirtualUdpSocket for MockSocket { + fn local_addr(&self) -> io::Result { + Ok(self.local_addr) + } + + async fn send_to(&self, data: &[u8], _addr: SocketAddr) -> io::Result { + Ok(data.len()) + } + + async fn recv_from(&self, _buf: &mut [u8]) -> io::Result<(usize, SocketAddr)> { + std::future::pending().await + } + } + + struct MockStun { + mapped_addr: SocketAddr, + fail: bool, + calls: AtomicUsize, + } + + impl MockStun { + fn succeeds_with(mapped_addr: SocketAddr) -> Self { + Self { + mapped_addr, + fail: false, + calls: AtomicUsize::new(0), + } + } + + fn failing() -> Self { + Self { + mapped_addr: "0.0.0.0:0".parse().unwrap(), + fail: true, + calls: AtomicUsize::new(0), + } + } + } + + #[async_trait] + impl StunInfoProvider for MockStun { + fn get_stun_info(&self) -> StunInfo { + StunInfo::default() + } + + async fn get_udp_port_mapping(&self, _local_port: u16) -> anyhow::Result { + Ok(self.mapped_addr) + } + + async fn get_tcp_port_mapping(&self, _local_port: u16) -> anyhow::Result { + Ok(self.mapped_addr) + } + + fn update_stun_info(&self) {} + } + + #[async_trait] + impl StunSocketMapper for MockStun { + async fn get_udp_port_mapping_with_socket( + &self, + _socket: Arc, + ) -> anyhow::Result { + self.calls.fetch_add(1, Ordering::SeqCst); + if self.fail { + anyhow::bail!("mock STUN failure"); + } + Ok(self.mapped_addr) + } + } + + #[derive(Debug, Default)] + struct LeaseState { + drops: AtomicUsize, + } + + #[derive(Default)] + struct MockEvents { + established: Mutex>, + } + + impl crate::events::CoreEventSink for MockEvents { + fn emit(&self, event: crate::events::CoreEvent) { + self.established.lock().unwrap().push(event); + } + } + + #[derive(Debug)] + struct MockMapping { + state: Arc, + backend: crate::connectivity::hole_punch::port_mapping::UdpPortMappingBackend, + } + + #[async_trait] + impl crate::connectivity::hole_punch::port_mapping::ActiveUdpPortMapping for MockMapping { + fn backend(&self) -> crate::connectivity::hole_punch::port_mapping::UdpPortMappingBackend { + self.backend + } + + fn local_addr(&self) -> SocketAddr { + "192.168.1.5:30123".parse().unwrap() + } + + fn gateway_external_port(&self) -> u16 { + 40123 + } + + async fn renew(&self) -> anyhow::Result<()> { + Ok(()) + } + + async fn remove(&self) -> anyhow::Result<()> { + self.state.drops.fetch_add(1, Ordering::SeqCst); + Ok(()) + } + } + + struct MockPlatform { + calls: AtomicUsize, + fail: bool, + lease_state: Arc, + } + + impl MockPlatform { + fn failing() -> Self { + Self { + calls: AtomicUsize::new(0), + fail: true, + lease_state: Arc::default(), + } + } + + fn with_lease(lease_state: Arc) -> Self { + Self { + calls: AtomicUsize::new(0), + fail: false, + lease_state, + } + } + } + + #[async_trait] + impl UdpPortMappingPlatform for MockPlatform { + async fn establish_udp_port_mapping( + &self, + backend: crate::connectivity::hole_punch::port_mapping::UdpPortMappingBackend, + _local_listener: &url::Url, + ) -> Result< + Box, + crate::connectivity::hole_punch::port_mapping::UdpPortMappingAttemptError, + > { + self.calls.fetch_add(1, Ordering::SeqCst); + if self.fail { + return Err( + crate::connectivity::hole_punch::port_mapping::UdpPortMappingAttemptError::establishment( + anyhow::anyhow!("mock port-mapping failure"), + ), + ); + } + Ok(Box::new(MockMapping { + state: self.lease_state.clone(), + backend, + })) + } + } + + fn socket() -> Arc { + Arc::new(MockSocket { + local_addr: "0.0.0.0:30123".parse().unwrap(), + }) + } + + fn listener_url() -> url::Url { + "udp://0.0.0.0:30123".parse().unwrap() + } + + #[tokio::test] + async fn port_mapping_failure_falls_back_to_stun() { + let mapped_addr = "198.51.100.8:40123".parse().unwrap(); + let stun = MockStun::succeeds_with(mapped_addr); + let platform = Arc::new(MockPlatform::failing()); + + let resolved = resolve_public_addr_with_policy( + &stun, + Some(platform.clone()), + Arc::new(()), + socket(), + &listener_url(), + false, + ) + .await + .unwrap(); + + assert_eq!(resolved.mapped_addr, mapped_addr); + assert!(resolved.port_mapping_lease.is_none()); + assert_eq!(platform.calls.load(Ordering::SeqCst), 2); + assert_eq!(stun.calls.load(Ordering::SeqCst), 1); + } + + #[tokio::test] + async fn stun_failure_releases_mapping_without_notification() { + let lease_state = Arc::new(LeaseState::default()); + let platform = Arc::new(MockPlatform::with_lease(lease_state.clone())); + let events = Arc::new(MockEvents::default()); + + let result = resolve_public_addr_with_policy( + &MockStun::failing(), + Some(platform), + events.clone(), + socket(), + &listener_url(), + false, + ) + .await; + + assert!(result.is_err()); + tokio::time::timeout(std::time::Duration::from_secs(1), async { + while lease_state.drops.load(Ordering::SeqCst) == 0 { + tokio::task::yield_now().await; + } + }) + .await + .unwrap(); + assert_eq!(lease_state.drops.load(Ordering::SeqCst), 1); + assert!(events.established.lock().unwrap().is_empty()); + } + + #[tokio::test] + async fn disable_upnp_is_applied_per_resolution() { + let mapped_addr = "198.51.100.9:40124".parse().unwrap(); + let stun = MockStun::succeeds_with(mapped_addr); + let lease_state = Arc::new(LeaseState::default()); + let platform = Arc::new(MockPlatform::with_lease(lease_state)); + + let disabled = resolve_public_addr_with_policy( + &stun, + Some(platform.clone()), + Arc::new(()), + socket(), + &listener_url(), + true, + ) + .await + .unwrap(); + assert!(disabled.port_mapping_lease.is_none()); + assert_eq!(platform.calls.load(Ordering::SeqCst), 0); + + let enabled = resolve_public_addr_with_policy( + &stun, + Some(platform.clone()), + Arc::new(()), + socket(), + &listener_url(), + false, + ) + .await + .unwrap(); + assert!(enabled.port_mapping_lease.is_some()); + assert_eq!(platform.calls.load(Ordering::SeqCst), 1); + } + + #[tokio::test] + async fn successful_mapping_is_notified_and_held_with_result() { + let mapped_addr = "198.51.100.10:40125".parse().unwrap(); + let lease_state = Arc::new(LeaseState::default()); + let platform = Arc::new(MockPlatform::with_lease(lease_state.clone())); + let events = Arc::new(MockEvents::default()); + + let resolved = resolve_public_addr_with_policy( + &MockStun::succeeds_with(mapped_addr), + Some(platform), + events.clone(), + socket(), + &listener_url(), + false, + ) + .await + .unwrap(); + + { + let established = events.established.lock().unwrap(); + assert!(matches!( + established.as_slice(), + [crate::events::CoreEvent::UdpPortMappingEstablished { + local_listener, + mapped_listener, + backend, + }] if local_listener == &listener_url() + && mapped_listener.as_str() == "udp://198.51.100.10:40125" + && backend == "igd" + )); + } + assert_eq!(lease_state.drops.load(Ordering::SeqCst), 0); + drop(resolved); + tokio::time::timeout(std::time::Duration::from_secs(1), async { + while lease_state.drops.load(Ordering::SeqCst) == 0 { + tokio::task::yield_now().await; + } + }) + .await + .unwrap(); + assert_eq!(lease_state.drops.load(Ordering::SeqCst), 1); + } + + #[test] + fn managed_local_addresses_are_rejected() { + let virtual_ipv4 = "10.144.0.2:1234".parse().unwrap(); + assert_eq!( + managed_local_addr_error(virtual_ipv4, true, false), + Some("local address is virtual ipv4") + ); + assert_eq!(managed_local_addr_error(virtual_ipv4, false, false), None); + + let managed_ipv6 = SocketAddr::V6(SocketAddrV6::new( + "fd00::1".parse::().unwrap(), + 1234, + 0, + 0, + )); + assert_eq!( + managed_local_addr_error(managed_ipv6, false, true), + Some("local address is easytier-managed ipv6") + ); + assert_eq!(managed_local_addr_error(managed_ipv6, false, false), None); + } +} diff --git a/easytier-core/src/connectivity/hole_punch/udp/client.rs b/easytier-core/src/connectivity/hole_punch/udp/client.rs new file mode 100644 index 00000000..a40fce5b --- /dev/null +++ b/easytier-core/src/connectivity/hole_punch/udp/client.rs @@ -0,0 +1,995 @@ +use std::{ + net::{IpAddr, Ipv4Addr, SocketAddr}, + sync::{ + Arc, + atomic::{AtomicBool, Ordering}, + }, +}; + +use guarden::defer; +use quanta::Instant; +use rand::Rng; +use tokio::sync::RwLock; +use tokio_util::task::AbortOnDropHandle; + +use crate::{ + config::PeerId, + connectivity::stun::StunInfoProvider, + packet::{HOLE_PUNCH_PACKET_BODY_LEN, new_hole_punch_packet}, + socket::udp::{UdpBindOptions, VirtualUdpSocketFactory}, +}; + +use super::{ + SelectPunchListener, SendPunchPacketBothEasySym, SendPunchPacketCone, SendPunchPacketEasySym, + SendPunchPacketHardSym, UdpHolePunchRuntime, UdpHolePunchSignalError, UdpHolePunchSignaling, + UdpNatType, UdpPunchSocket, UdpSocketArray, +}; + +#[derive(Debug, thiserror::Error)] +pub enum UdpHolePunchClientError { + #[error("signaling: {0}")] + Signaling(#[from] UdpHolePunchSignalError), + #[error(transparent)] + Other(#[from] anyhow::Error), +} + +pub type UdpHolePunchClientResult = Result; + +const UDP_ARRAY_SIZE_FOR_HARD_SYM: usize = 84; +const UDP_ARRAY_SIZE_FOR_BOTH_EASY_SYM: usize = 25; +const DST_PORT_OFFSET: u16 = 20; +const REMOTE_WAIT_TIME_MS: u64 = 5000; + +pub fn apply_peer_easy_sym_port_offset(base_port: u16, peer_is_incremental: bool) -> u16 { + let port = if peer_is_incremental { + (base_port as u32).saturating_add(DST_PORT_OFFSET as u32) + } else { + (base_port as u32).saturating_sub(DST_PORT_OFFSET as u32) + }; + port as u16 +} + +#[tracing::instrument(skip(runtime, signaling), fields(dst_peer_id), err)] +pub async fn punch_cone_to_cone( + runtime: Arc, + signaling: Arc, + dst_peer_id: PeerId, +) -> UdpHolePunchClientResult> +where + R: UdpHolePunchRuntime, + S: UdpHolePunchSignaling + 'static, +{ + tracing::info!(?dst_peer_id, "start hole punching"); + let tid = rand::random(); + + let udp_array = UdpSocketArray::new_with_context(1, runtime.clone(), runtime.socket_context()); + + let resp = signaling + .select_punch_listener( + dst_peer_id, + SelectPunchListener { + force_new: false, + prefer_port_mapping: true, + }, + ) + .await?; + let remote_mapped_addr = resp.listener_mapped_addr; + + let local_socket = UdpHolePunchRuntime::bind_udp( + runtime.as_ref(), + UdpBindOptions::hole_punch_control().with_context( + runtime + .socket_context() + .with_ip_version(crate::socket::IpVersion::V4), + ), + ) + .await?; + let resolved = runtime + .resolve_udp_public_addr(local_socket.clone()) + .await?; + let local_mapped_addr = resolved.mapped_addr; + let _local_port_mapping_lease = resolved.port_mapping_lease; + + tracing::debug!( + ?local_mapped_addr, + ?remote_mapped_addr, + "hole punch got remote listener" + ); + + udp_array.add_new_socket(local_socket).await?; + udp_array.add_intreast_tid(tid); + let punch_packet = new_hole_punch_packet(tid, HOLE_PUNCH_PACKET_BODY_LEN).into_bytes(); + + send_from_local(&udp_array, &punch_packet, remote_mapped_addr).await?; + + let signaling_for_task = signaling.clone(); + let punch_task = AbortOnDropHandle::new(tokio::spawn(async move { + if let Err(e) = signaling_for_task + .send_punch_packet_cone( + dst_peer_id, + SendPunchPacketCone { + listener_mapped_addr: remote_mapped_addr, + dest_addr: local_mapped_addr, + transaction_id: tid, + packet_count_per_batch: 2, + packet_batch_count: 5, + packet_interval_ms: 400, + }, + ) + .await + { + tracing::error!(?e, "failed to call remote send punch packet"); + } + })); + + let mut finish_time: Option = None; + while finish_time.is_none() || finish_time.as_ref().unwrap().elapsed().as_millis() < 1000 { + crate::foundation::time::sleep(std::time::Duration::from_millis(200)).await; + + if finish_time.is_none() && punch_task.is_finished() { + finish_time = Some(Instant::now()); + } + + let Some(socket) = udp_array.try_fetch_punched_socket(tid) else { + tracing::debug!("no punched socket found, send some more hole punch packets"); + send_from_local(&udp_array, &punch_packet, remote_mapped_addr).await?; + continue; + }; + + tracing::debug!(?socket, ?tid, "punched socket found, try connect with it"); + + for _ in 0..2 { + match runtime + .connect_with_socket(socket.socket.clone(), remote_mapped_addr) + .await + { + Ok(socket) => { + tracing::info!(?socket, "hole punched"); + return Ok(Some(socket)); + } + Err(e) => { + tracing::error!(?e, "failed to connect with socket"); + } + } + } + } + + Ok(None) +} + +async fn send_from_local( + udp_array: &UdpSocketArray, + punch_packet: &[u8], + remote_mapped_addr: SocketAddr, +) -> UdpHolePunchClientResult<()> +where + R: VirtualUdpSocketFactory, +{ + udp_array + .send_with_all(punch_packet, remote_mapped_addr) + .await?; + Ok(()) +} + +pub struct UdpSymToConePunchClient +where + R: UdpHolePunchRuntime, + S: UdpHolePunchSignaling + 'static, +{ + runtime: Arc, + signaling: Arc, + stun: Arc, + udp_array: RwLock>>>, + try_direct_connect: AtomicBool, + punch_predictably: AtomicBool, +} + +impl UdpSymToConePunchClient +where + R: UdpHolePunchRuntime, + S: UdpHolePunchSignaling + 'static, +{ + pub fn new(runtime: Arc, signaling: Arc, stun: Arc) -> Self { + Self { + runtime, + signaling, + stun, + udp_array: RwLock::new(None), + try_direct_connect: AtomicBool::new(true), + punch_predictably: AtomicBool::new(true), + } + } + + pub async fn clear_udp_array(&self) { + let mut wlocked = self.udp_array.write().await; + wlocked.take(); + } + + async fn prepare_udp_array(&self) -> anyhow::Result>> { + let rlocked = self.udp_array.read().await; + if let Some(udp_array) = rlocked.clone() { + return Ok(udp_array); + } + + drop(rlocked); + let mut wlocked = self.udp_array.write().await; + if let Some(udp_array) = wlocked.clone() { + return Ok(udp_array); + } + + let udp_array = Arc::new(UdpSocketArray::new_with_context( + UDP_ARRAY_SIZE_FOR_HARD_SYM, + self.runtime.clone(), + self.runtime.socket_context(), + )); + udp_array.start().await?; + wlocked.replace(udp_array.clone()); + Ok(udp_array) + } + + async fn get_base_port_for_easy_sym(&self, my_nat_info: UdpNatType) -> Option { + if my_nat_info.is_easy_sym() { + match self.stun.get_udp_port_mapping(0).await { + Ok(addr) => Some(addr.port()), + ret => { + tracing::warn!(?ret, "failed to get udp port mapping for easy sym"); + None + } + } + } else { + None + } + } + + async fn remote_send_hole_punch_packet_predictable( + signaling: Arc, + dst_peer_id: PeerId, + base_port_for_easy_sym: Option, + my_nat_info: UdpNatType, + remote_mapped_addr: SocketAddr, + public_ips: Vec, + tid: u32, + ) { + let Some(inc) = my_nat_info.get_inc_of_easy_sym() else { + return; + }; + let req = SendPunchPacketEasySym { + listener_mapped_addr: remote_mapped_addr, + public_ips, + transaction_id: tid, + base_port_num: base_port_for_easy_sym.unwrap() as u32, + max_port_num: 50, + is_incremental: inc, + }; + tracing::debug!(?req, "send punch packet for easy sym start"); + let ret = signaling.send_punch_packet_easy_sym(dst_peer_id, req).await; + tracing::debug!(?ret, "send punch packet for easy sym return"); + } + + async fn remote_send_hole_punch_packet_random( + signaling: Arc, + dst_peer_id: PeerId, + remote_mapped_addr: SocketAddr, + public_ips: Vec, + tid: u32, + round: u32, + port_index: u32, + ) -> Option { + let req = SendPunchPacketHardSym { + listener_mapped_addr: remote_mapped_addr, + public_ips, + transaction_id: tid, + round, + port_index, + }; + tracing::debug!(?req, "send punch packet for hard sym start"); + match signaling.send_punch_packet_hard_sym(dst_peer_id, req).await { + Err(e) => { + tracing::error!(?e, "failed to send punch packet for hard sym"); + None + } + Ok(resp) => Some(resp.next_port_index), + } + } + + async fn check_hole_punch_result( + &self, + udp_array: &Arc>, + packet: &[u8], + tid: u32, + remote_mapped_addr: SocketAddr, + punch_task: &AbortOnDropHandle, + ) -> anyhow::Result> { + let mut ret_socket = None; + let mut finish_time: Option = None; + while finish_time.is_none() || finish_time.as_ref().unwrap().elapsed().as_millis() < 1000 { + udp_array.send_with_all(packet, remote_mapped_addr).await?; + + crate::foundation::time::sleep(std::time::Duration::from_millis(200)).await; + + if finish_time.is_none() && punch_task.is_finished() { + finish_time = Some(Instant::now()); + } + + let Some(socket) = udp_array.try_fetch_punched_socket(tid) else { + tracing::debug!("no punched socket found, wait for more time"); + continue; + }; + + match self + .runtime + .connect_with_socket(socket.socket.clone(), remote_mapped_addr) + .await + { + Ok(socket) => { + ret_socket.replace(socket); + break; + } + Err(e) => { + tracing::error!(?e, "failed to connect with socket"); + udp_array.add_new_socket(socket.socket).await?; + continue; + } + } + } + + Ok(ret_socket) + } + + #[tracing::instrument(err(level = tracing::Level::ERROR), skip(self))] + pub async fn do_hole_punching( + &self, + dst_peer_id: PeerId, + round: u32, + last_port_idx: &mut usize, + my_nat_info: UdpNatType, + ) -> UdpHolePunchClientResult> { + let udp_array = self.prepare_udp_array().await?; + + let resp = self + .signaling + .select_punch_listener( + dst_peer_id, + SelectPunchListener { + force_new: false, + prefer_port_mapping: true, + }, + ) + .await?; + + let remote_mapped_addr = resp.listener_mapped_addr; + + if self.try_direct_connect.load(Ordering::Relaxed) { + let socket = self.runtime.bind_direct_connect_udp().await?; + if let Ok(socket) = self + .runtime + .connect_with_socket(socket, remote_mapped_addr) + .await + { + return Ok(Some(socket)); + } + } + + let stun_info = self.stun.get_stun_info(); + let public_ips: Vec = stun_info + .public_ip + .iter() + .filter_map(|x| x.parse().ok()) + .collect(); + if public_ips.is_empty() { + return Err(anyhow::anyhow!("failed to get public ips").into()); + } + + let tid = rand::thread_rng().r#gen(); + let packet = new_hole_punch_packet(tid, HOLE_PUNCH_PACKET_BODY_LEN).into_bytes(); + udp_array.add_intreast_tid(tid); + defer! { udp_array.remove_intreast_tid(tid); } + + let port_index = *last_port_idx as u32; + let base_port_for_easy_sym = self.get_base_port_for_easy_sym(my_nat_info).await; + udp_array.send_with_all(&packet, remote_mapped_addr).await?; + + if self.punch_predictably.load(Ordering::Relaxed) && base_port_for_easy_sym.is_some() { + let signaling = self.signaling.clone(); + let punch_task = AbortOnDropHandle::new(tokio::spawn( + Self::remote_send_hole_punch_packet_predictable( + signaling, + dst_peer_id, + base_port_for_easy_sym, + my_nat_info, + remote_mapped_addr, + public_ips.clone(), + tid, + ), + )); + let ret_socket = self + .check_hole_punch_result(&udp_array, &packet, tid, remote_mapped_addr, &punch_task) + .await?; + + let task_ret = punch_task.await; + tracing::debug!(?ret_socket, ?task_ret, "predictable punch task got result"); + if let Some(socket) = ret_socket { + return Ok(Some(socket)); + } + } + + let signaling = self.signaling.clone(); + let punch_task = + AbortOnDropHandle::new(tokio::spawn(Self::remote_send_hole_punch_packet_random( + signaling, + dst_peer_id, + remote_mapped_addr, + public_ips.clone(), + tid, + round, + port_index, + ))); + let ret_socket = self + .check_hole_punch_result(&udp_array, &packet, tid, remote_mapped_addr, &punch_task) + .await?; + + let punch_task_result = punch_task.await; + tracing::debug!(?punch_task_result, ?ret_socket, "punch task got result"); + + if let Ok(Some(next_port_idx)) = punch_task_result { + *last_port_idx = next_port_idx as usize; + } else { + *last_port_idx = rand::random(); + } + + Ok(ret_socket) + } +} + +pub struct UdpBothEasySymPunchClient +where + R: UdpHolePunchRuntime, + S: UdpHolePunchSignaling + 'static, +{ + runtime: Arc, + signaling: Arc, + stun: Arc, +} + +impl UdpBothEasySymPunchClient +where + R: UdpHolePunchRuntime, + S: UdpHolePunchSignaling + 'static, +{ + pub fn new(runtime: Arc, signaling: Arc, stun: Arc) -> Self { + Self { + runtime, + signaling, + stun, + } + } + + #[tracing::instrument(ret, skip(self))] + pub async fn do_hole_punching( + &self, + dst_peer_id: PeerId, + my_nat_info: UdpNatType, + peer_nat_info: UdpNatType, + is_busy: &mut bool, + ) -> UdpHolePunchClientResult> { + *is_busy = false; + + let udp_array = UdpSocketArray::new_with_context( + UDP_ARRAY_SIZE_FOR_BOTH_EASY_SYM, + self.runtime.clone(), + self.runtime.socket_context(), + ); + udp_array.start().await?; + + let cur_mapped_addr = self.stun.get_udp_port_mapping(0).await?; + let my_public_ip = match cur_mapped_addr.ip() { + IpAddr::V4(v4) => v4, + _ => { + return Err(anyhow::anyhow!("ipv6 is not supported").into()); + } + }; + let me_is_incremental = my_nat_info + .get_inc_of_easy_sym() + .ok_or(anyhow::anyhow!("me_is_incremental is required"))?; + let peer_is_incremental = peer_nat_info + .get_inc_of_easy_sym() + .ok_or(anyhow::anyhow!("peer_is_incremental is required"))?; + + let tid = rand::random(); + udp_array.add_intreast_tid(tid); + + let remote_ret = self + .signaling + .send_punch_packet_both_easy_sym( + dst_peer_id, + SendPunchPacketBothEasySym { + transaction_id: tid, + public_ip: my_public_ip, + dst_port_num: if me_is_incremental { + cur_mapped_addr.port().saturating_add(DST_PORT_OFFSET) + } else { + cur_mapped_addr.port().saturating_sub(DST_PORT_OFFSET) + } as u32, + udp_socket_count: UDP_ARRAY_SIZE_FOR_BOTH_EASY_SYM as u32, + wait_time_ms: REMOTE_WAIT_TIME_MS as u32, + }, + ) + .await?; + + if remote_ret.is_busy { + *is_busy = true; + return Err(anyhow::anyhow!("remote is busy").into()); + } + + let mut remote_mapped_addr = remote_ret + .base_mapped_addr + .ok_or(anyhow::anyhow!("remote_mapped_addr is required"))?; + + let now = Instant::now(); + remote_mapped_addr.set_port(apply_peer_easy_sym_port_offset( + remote_mapped_addr.port(), + peer_is_incremental, + )); + tracing::debug!( + ?remote_mapped_addr, + ?remote_ret, + "start send hole punch packet for both easy sym" + ); + + while now.elapsed().as_millis() < (REMOTE_WAIT_TIME_MS + 1000).into() { + udp_array + .send_with_all( + &new_hole_punch_packet(tid, HOLE_PUNCH_PACKET_BODY_LEN).into_bytes(), + remote_mapped_addr, + ) + .await?; + + crate::foundation::time::sleep(std::time::Duration::from_millis(100)).await; + + let Some(socket) = udp_array.try_fetch_punched_socket(tid) else { + tracing::trace!( + ?remote_mapped_addr, + ?tid, + "no punched socket found, send some more hole punch packets" + ); + continue; + }; + + tracing::info!( + ?socket, + ?remote_mapped_addr, + ?tid, + "got punched socket in both easy sym" + ); + + for _ in 0..2 { + match self + .runtime + .connect_with_socket(socket.socket.clone(), remote_mapped_addr) + .await + { + Ok(socket) => { + return Ok(Some(socket)); + } + Err(e) => { + tracing::error!(?e, "failed to connect with socket"); + continue; + } + } + } + udp_array.add_new_socket(socket.socket).await?; + } + + Ok(None) + } +} + +#[cfg(test)] +mod tests { + use std::{ + io, + sync::{ + Arc, + atomic::{AtomicUsize, Ordering}, + }, + }; + + use async_trait::async_trait; + + use super::*; + use crate::{ + proto::common::{NatType, StunInfo}, + socket::udp::VirtualUdpSocket, + }; + + impl UdpSymToConePunchClient + where + R: UdpHolePunchRuntime, + S: UdpHolePunchSignaling + 'static, + { + fn set_try_direct_connect(&self, enabled: bool) { + self.try_direct_connect.store(enabled, Ordering::Relaxed); + } + } + + struct MockSocket { + local_addr: SocketAddr, + } + + #[async_trait] + impl VirtualUdpSocket for MockSocket { + fn local_addr(&self) -> io::Result { + Ok(self.local_addr) + } + + async fn send_to(&self, _data: &[u8], _addr: SocketAddr) -> io::Result { + Ok(0) + } + + async fn recv_from(&self, _buf: &mut [u8]) -> io::Result<(usize, SocketAddr)> { + std::future::pending().await + } + } + + struct MockRuntime { + bind_count: AtomicUsize, + resolve_count: AtomicUsize, + bind_options: tokio::sync::Mutex>, + } + + impl MockRuntime { + fn new() -> Self { + Self { + bind_count: AtomicUsize::new(0), + resolve_count: AtomicUsize::new(0), + bind_options: tokio::sync::Mutex::new(Vec::new()), + } + } + } + + #[async_trait] + impl UdpHolePunchRuntime for MockRuntime { + type Socket = MockSocket; + + async fn bind_udp(&self, options: UdpBindOptions) -> anyhow::Result> { + self.bind_options.lock().await.push(options); + let bind_idx = self.bind_count.fetch_add(1, Ordering::Relaxed); + Ok(Arc::new(MockSocket { + local_addr: SocketAddr::from(([127, 0, 0, 1], 10000 + bind_idx as u16)), + })) + } + + async fn resolve_udp_public_addr( + &self, + _socket: Arc, + ) -> anyhow::Result { + self.resolve_count.fetch_add(1, Ordering::Relaxed); + Ok(super::super::UdpResolvedPublicAddr { + mapped_addr: SocketAddr::from(([203, 0, 113, 1], 10000)), + port_mapping_lease: None, + }) + } + + async fn create_listener( + &self, + _prefer_port_mapping: bool, + ) -> anyhow::Result> { + unimplemented!("not used by cone client tests") + } + + async fn create_port_bound_listener( + &self, + _port: u16, + ) -> anyhow::Result> { + unimplemented!("not used by cone client tests") + } + + async fn connect_with_socket( + &self, + _socket: Arc, + _remote: SocketAddr, + ) -> anyhow::Result { + unimplemented!("not used by cone client tests") + } + } + + #[derive(Default)] + struct MockStunInfoProvider { + port_mapping_count: AtomicUsize, + } + + #[async_trait] + impl StunInfoProvider for MockStunInfoProvider { + fn get_stun_info(&self) -> StunInfo { + StunInfo { + public_ip: vec!["127.0.0.1".to_string()], + ..Default::default() + } + } + + async fn get_udp_port_mapping(&self, _port: u16) -> anyhow::Result { + self.port_mapping_count.fetch_add(1, Ordering::Relaxed); + Ok(SocketAddr::from(([203, 0, 113, 1], 10000))) + } + + async fn get_tcp_port_mapping(&self, _port: u16) -> anyhow::Result { + unreachable!("TCP mapping is not used by UDP hole-punch tests") + } + + fn update_stun_info(&self) {} + } + + struct RejectingSignaling; + + #[async_trait] + impl UdpHolePunchSignaling for RejectingSignaling { + async fn select_punch_listener( + &self, + _dst_peer_id: PeerId, + _request: SelectPunchListener, + ) -> Result { + Err(UdpHolePunchSignalError::InvalidServiceKey) + } + + async fn send_punch_packet_cone( + &self, + _dst_peer_id: PeerId, + _request: SendPunchPacketCone, + ) -> Result<(), UdpHolePunchSignalError> { + Ok(()) + } + + async fn send_punch_packet_hard_sym( + &self, + _dst_peer_id: PeerId, + _request: super::super::SendPunchPacketHardSym, + ) -> Result { + unimplemented!("not used by cone client tests") + } + + async fn send_punch_packet_easy_sym( + &self, + _dst_peer_id: PeerId, + _request: super::super::SendPunchPacketEasySym, + ) -> Result<(), UdpHolePunchSignalError> { + unimplemented!("not used by cone client tests") + } + + async fn send_punch_packet_both_easy_sym( + &self, + _dst_peer_id: PeerId, + _request: super::super::SendPunchPacketBothEasySym, + ) -> Result + { + unimplemented!("not used by cone client tests") + } + } + + #[tokio::test] + async fn cone_punch_does_not_bind_or_resolve_before_listener_rpc_succeeds() { + let runtime = Arc::new(MockRuntime::new()); + let signaling = Arc::new(RejectingSignaling); + + let err = punch_cone_to_cone(runtime.clone(), signaling, 2) + .await + .unwrap_err(); + + assert!(matches!( + err, + UdpHolePunchClientError::Signaling(UdpHolePunchSignalError::InvalidServiceKey) + )); + assert_eq!(runtime.bind_count.load(Ordering::Relaxed), 0); + assert_eq!(runtime.resolve_count.load(Ordering::Relaxed), 0); + } + + #[tokio::test] + async fn default_direct_connect_bind_uses_direct_connect_purpose() { + let runtime = MockRuntime::new(); + + let socket = runtime.bind_direct_connect_udp().await.unwrap(); + + assert_eq!(socket.local_addr().unwrap().port(), 10000); + let bind_options = runtime.bind_options.lock().await; + assert_eq!( + bind_options.as_slice(), + &[UdpBindOptions::direct_connect().with_ip_version(crate::socket::IpVersion::V4)] + ); + } + + struct RecordingSignaling { + easy_requests: tokio::sync::Mutex>, + hard_requests: tokio::sync::Mutex>, + both_requests: tokio::sync::Mutex>, + both_response: super::super::SendPunchPacketBothEasySymResponse, + next_port_index: u32, + } + + impl RecordingSignaling { + fn new(next_port_index: u32) -> Self { + Self { + easy_requests: tokio::sync::Mutex::new(Vec::new()), + hard_requests: tokio::sync::Mutex::new(Vec::new()), + both_requests: tokio::sync::Mutex::new(Vec::new()), + both_response: super::super::SendPunchPacketBothEasySymResponse { + is_busy: false, + base_mapped_addr: Some(SocketAddr::from(([127, 0, 0, 1], 40144))), + }, + next_port_index, + } + } + + fn with_both_response( + mut self, + both_response: super::super::SendPunchPacketBothEasySymResponse, + ) -> Self { + self.both_response = both_response; + self + } + } + + #[async_trait] + impl UdpHolePunchSignaling for RecordingSignaling { + async fn select_punch_listener( + &self, + _dst_peer_id: PeerId, + _request: SelectPunchListener, + ) -> Result { + Ok(super::super::SelectPunchListenerResponse { + listener_mapped_addr: SocketAddr::from(([127, 0, 0, 1], 30000)), + }) + } + + async fn send_punch_packet_cone( + &self, + _dst_peer_id: PeerId, + _request: SendPunchPacketCone, + ) -> Result<(), UdpHolePunchSignalError> { + Ok(()) + } + + async fn send_punch_packet_hard_sym( + &self, + _dst_peer_id: PeerId, + request: SendPunchPacketHardSym, + ) -> Result { + self.hard_requests.lock().await.push(request); + Ok(super::super::SendPunchPacketHardSymResponse { + next_port_index: self.next_port_index, + }) + } + + async fn send_punch_packet_easy_sym( + &self, + _dst_peer_id: PeerId, + request: SendPunchPacketEasySym, + ) -> Result<(), UdpHolePunchSignalError> { + self.easy_requests.lock().await.push(request); + Ok(()) + } + + async fn send_punch_packet_both_easy_sym( + &self, + _dst_peer_id: PeerId, + request: SendPunchPacketBothEasySym, + ) -> Result + { + self.both_requests.lock().await.push(request); + Ok(self.both_response.clone()) + } + } + + #[tokio::test] + async fn sym_to_cone_easy_sym_uses_port_mapping_in_predictable_request() { + let runtime = Arc::new(MockRuntime::new()); + let signaling = Arc::new(RecordingSignaling::new(42)); + let stun = Arc::new(MockStunInfoProvider::default()); + let client = UdpSymToConePunchClient::new(runtime.clone(), signaling.clone(), stun.clone()); + client.set_try_direct_connect(false); + + let mut last_port_idx = 7; + let ret = client + .do_hole_punching(2, 3, &mut last_port_idx, NatType::SymmetricEasyInc.into()) + .await + .unwrap(); + + assert!(ret.is_none()); + assert_eq!(stun.port_mapping_count.load(Ordering::Relaxed), 1); + + let easy_requests = signaling.easy_requests.lock().await; + assert_eq!(easy_requests.len(), 1); + let req = &easy_requests[0]; + assert_eq!( + req.listener_mapped_addr, + SocketAddr::from(([127, 0, 0, 1], 30000)) + ); + assert_eq!(req.public_ips, vec![Ipv4Addr::new(127, 0, 0, 1)]); + assert_eq!(req.base_port_num, 10000); + assert_eq!(req.max_port_num, 50); + assert!(req.is_incremental); + + let hard_requests = signaling.hard_requests.lock().await; + assert_eq!(hard_requests.len(), 1); + assert_eq!(last_port_idx, 42); + } + + #[tokio::test] + async fn sym_to_cone_hard_sym_sends_random_request_and_updates_port_index() { + let runtime = Arc::new(MockRuntime::new()); + let signaling = Arc::new(RecordingSignaling::new(321)); + let stun = Arc::new(MockStunInfoProvider::default()); + let client = UdpSymToConePunchClient::new(runtime.clone(), signaling.clone(), stun.clone()); + client.set_try_direct_connect(false); + + let mut last_port_idx = 123; + let ret = client + .do_hole_punching(2, 4, &mut last_port_idx, NatType::Symmetric.into()) + .await + .unwrap(); + + assert!(ret.is_none()); + assert_eq!(stun.port_mapping_count.load(Ordering::Relaxed), 0); + assert!(signaling.easy_requests.lock().await.is_empty()); + + let hard_requests = signaling.hard_requests.lock().await; + assert_eq!(hard_requests.len(), 1); + let req = &hard_requests[0]; + assert_eq!( + req.listener_mapped_addr, + SocketAddr::from(([127, 0, 0, 1], 30000)) + ); + assert_eq!(req.public_ips, vec![Ipv4Addr::new(127, 0, 0, 1)]); + assert_eq!(req.round, 4); + assert_eq!(req.port_index, 123); + assert_eq!(last_port_idx, 321); + } + + #[test] + fn both_easy_sym_port_offset_preserves_old_proto_cast_semantics() { + assert_eq!(apply_peer_easy_sym_port_offset(65530, true), 14); + assert_eq!(apply_peer_easy_sym_port_offset(10, false), 0); + } + + #[tokio::test] + async fn both_easy_sym_sends_remote_request_and_reports_busy() { + let runtime = Arc::new(MockRuntime::new()); + let stun = Arc::new(MockStunInfoProvider::default()); + let signaling = Arc::new(RecordingSignaling::new(0).with_both_response( + super::super::SendPunchPacketBothEasySymResponse { + is_busy: true, + base_mapped_addr: None, + }, + )); + let client = + UdpBothEasySymPunchClient::new(runtime.clone(), signaling.clone(), stun.clone()); + + let mut is_busy = false; + let err = client + .do_hole_punching( + 2, + NatType::SymmetricEasyInc.into(), + NatType::SymmetricEasyDec.into(), + &mut is_busy, + ) + .await + .unwrap_err(); + + assert!(is_busy); + assert!(err.to_string().contains("remote is busy")); + assert_eq!(stun.port_mapping_count.load(Ordering::Relaxed), 1); + assert_eq!( + runtime.bind_count.load(Ordering::Relaxed), + UDP_ARRAY_SIZE_FOR_BOTH_EASY_SYM + ); + + let both_requests = signaling.both_requests.lock().await; + assert_eq!(both_requests.len(), 1); + let req = &both_requests[0]; + assert_eq!(req.public_ip, Ipv4Addr::new(203, 0, 113, 1)); + assert_eq!(req.dst_port_num, 10020); + assert_eq!( + req.udp_socket_count, + UDP_ARRAY_SIZE_FOR_BOTH_EASY_SYM as u32 + ); + assert_eq!(req.wait_time_ms, REMOTE_WAIT_TIME_MS as u32); + } +} diff --git a/easytier-core/src/connectivity/hole_punch/udp/common.rs b/easytier-core/src/connectivity/hole_punch/udp/common.rs new file mode 100644 index 00000000..aada7425 --- /dev/null +++ b/easytier-core/src/connectivity/hole_punch/udp/common.rs @@ -0,0 +1,249 @@ +use crate::{config::PeerId, proto::common::NatType}; + +pub const BLACKLIST_TIMEOUT_SEC: u64 = 3600; + +#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] +pub enum UdpPunchClientMethod { + None, + ConeToCone, + SymToCone, + EasySymToEasySym, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] +pub enum UdpNatType { + Unknown, + Open(NatType), + Cone(NatType), + EasySymmetric(NatType, bool), + HardSymmetric(NatType), +} + +impl From for UdpNatType { + fn from(nat_type: NatType) -> Self { + match nat_type { + NatType::Unknown => UdpNatType::Unknown, + NatType::OpenInternet => UdpNatType::Open(nat_type), + NatType::NoPat | NatType::FullCone | NatType::Restricted | NatType::PortRestricted => { + UdpNatType::Cone(nat_type) + } + NatType::Symmetric | NatType::SymUdpFirewall => UdpNatType::HardSymmetric(nat_type), + NatType::SymmetricEasyInc => UdpNatType::EasySymmetric(nat_type, true), + NatType::SymmetricEasyDec => UdpNatType::EasySymmetric(nat_type, false), + } + } +} + +impl From for NatType { + fn from(val: UdpNatType) -> Self { + match val { + UdpNatType::Unknown => NatType::Unknown, + UdpNatType::Open(nat_type) => nat_type, + UdpNatType::Cone(nat_type) => nat_type, + UdpNatType::EasySymmetric(nat_type, _) => nat_type, + UdpNatType::HardSymmetric(nat_type) => nat_type, + } + } +} + +impl UdpNatType { + pub fn is_open(&self) -> bool { + matches!(self, UdpNatType::Open(_)) + } + + pub fn is_unknown(&self) -> bool { + matches!(self, UdpNatType::Unknown) + } + + pub fn is_sym(&self) -> bool { + self.is_hard_sym() || self.is_easy_sym() + } + + pub fn is_hard_sym(&self) -> bool { + matches!(self, UdpNatType::HardSymmetric(_)) + } + + pub fn is_easy_sym(&self) -> bool { + matches!(self, UdpNatType::EasySymmetric(_, _)) + } + + pub fn is_cone(&self) -> bool { + matches!(self, UdpNatType::Cone(_)) + } + + pub fn get_inc_of_easy_sym(&self) -> Option { + match self { + UdpNatType::EasySymmetric(_, inc) => Some(*inc), + _ => None, + } + } + + pub fn get_punch_hole_method( + &self, + other: Self, + disable_sym_hole_punching: bool, + ) -> UdpPunchClientMethod { + if disable_sym_hole_punching && self.is_sym() { + if other.is_sym() { + return UdpPunchClientMethod::None; + } else { + return UdpPunchClientMethod::ConeToCone; + } + } + + if other.is_unknown() { + if self.is_sym() { + return UdpPunchClientMethod::SymToCone; + } else { + return UdpPunchClientMethod::ConeToCone; + } + } + + if self.is_unknown() { + if other.is_sym() { + return UdpPunchClientMethod::None; + } else { + return UdpPunchClientMethod::ConeToCone; + } + } + + if self.is_open() || other.is_open() { + return UdpPunchClientMethod::None; + } + + if self.is_cone() { + if other.is_sym() { + UdpPunchClientMethod::None + } else { + UdpPunchClientMethod::ConeToCone + } + } else if self.is_easy_sym() { + if other.is_hard_sym() { + UdpPunchClientMethod::None + } else if other.is_easy_sym() { + UdpPunchClientMethod::EasySymToEasySym + } else { + UdpPunchClientMethod::SymToCone + } + } else if self.is_hard_sym() { + if other.is_sym() { + UdpPunchClientMethod::None + } else { + UdpPunchClientMethod::SymToCone + } + } else { + unreachable!("invalid nat type"); + } + } + + pub fn can_punch_hole_as_client( + &self, + other: Self, + my_peer_id: PeerId, + dst_peer_id: PeerId, + disable_sym_hole_punching: bool, + ) -> bool { + match self.get_punch_hole_method(other, disable_sym_hole_punching) { + UdpPunchClientMethod::None => false, + UdpPunchClientMethod::ConeToCone | UdpPunchClientMethod::SymToCone => true, + UdpPunchClientMethod::EasySymToEasySym => my_peer_id < dst_peer_id, + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + + fn nat(nat_type: NatType) -> UdpNatType { + nat_type.into() + } + + #[test] + fn nat_type_classification_matches_proto_values() { + assert_eq!(nat(NatType::Unknown), UdpNatType::Unknown); + assert_eq!( + nat(NatType::OpenInternet), + UdpNatType::Open(NatType::OpenInternet) + ); + assert_eq!(nat(NatType::FullCone), UdpNatType::Cone(NatType::FullCone)); + assert_eq!( + nat(NatType::Symmetric), + UdpNatType::HardSymmetric(NatType::Symmetric) + ); + assert_eq!( + nat(NatType::SymmetricEasyInc), + UdpNatType::EasySymmetric(NatType::SymmetricEasyInc, true) + ); + assert_eq!( + nat(NatType::SymmetricEasyDec), + UdpNatType::EasySymmetric(NatType::SymmetricEasyDec, false) + ); + } + + #[test] + fn punch_method_preserves_current_nat_matrix() { + let cone = nat(NatType::FullCone); + let hard_sym = nat(NatType::Symmetric); + let easy_sym = nat(NatType::SymmetricEasyInc); + let open = nat(NatType::OpenInternet); + let unknown = nat(NatType::Unknown); + + assert_eq!( + cone.get_punch_hole_method(cone, false), + UdpPunchClientMethod::ConeToCone + ); + assert_eq!( + hard_sym.get_punch_hole_method(cone, false), + UdpPunchClientMethod::SymToCone + ); + assert_eq!( + cone.get_punch_hole_method(hard_sym, false), + UdpPunchClientMethod::None + ); + assert_eq!( + easy_sym.get_punch_hole_method(easy_sym, false), + UdpPunchClientMethod::EasySymToEasySym + ); + assert_eq!( + easy_sym.get_punch_hole_method(hard_sym, false), + UdpPunchClientMethod::None + ); + assert_eq!( + hard_sym.get_punch_hole_method(unknown, false), + UdpPunchClientMethod::SymToCone + ); + assert_eq!( + unknown.get_punch_hole_method(hard_sym, false), + UdpPunchClientMethod::None + ); + assert_eq!( + open.get_punch_hole_method(cone, false), + UdpPunchClientMethod::None + ); + } + + #[test] + fn disabled_symmetric_punching_keeps_existing_fallback() { + let cone = nat(NatType::FullCone); + let hard_sym = nat(NatType::Symmetric); + let easy_sym = nat(NatType::SymmetricEasyInc); + + assert_eq!( + hard_sym.get_punch_hole_method(cone, true), + UdpPunchClientMethod::ConeToCone + ); + assert_eq!( + hard_sym.get_punch_hole_method(easy_sym, true), + UdpPunchClientMethod::None + ); + } + + #[test] + fn easy_sym_to_easy_sym_uses_lower_peer_id_as_client() { + let easy_sym = nat(NatType::SymmetricEasyInc); + + assert!(easy_sym.can_punch_hole_as_client(easy_sym, 1, 2, false)); + assert!(!easy_sym.can_punch_hole_as_client(easy_sym, 2, 1, false)); + } +} diff --git a/easytier-core/src/connectivity/hole_punch/udp/connector.rs b/easytier-core/src/connectivity/hole_punch/udp/connector.rs new file mode 100644 index 00000000..e79259a9 --- /dev/null +++ b/easytier-core/src/connectivity/hole_punch/udp/connector.rs @@ -0,0 +1,527 @@ +use std::sync::{ + Arc, + atomic::{AtomicBool, Ordering}, +}; + +use anyhow::Error; +use dashmap::DashMap; +use quanta::Instant; +use tokio::{ + sync::{Mutex, OwnedMutexGuard, TryLockError}, + task::JoinHandle, +}; + +use crate::{ + config::PeerId, + connectivity::stun::StunInfoProvider, + foundation::task::{ExternalTaskSignal, PeerTaskLauncher, PeerTaskManager}, + proto::common::NatType, +}; + +use crate::connectivity::hole_punch::policy::BackOff; + +use super::{ + BLACKLIST_TIMEOUT_SEC, UdpBothEasySymPunchClient, UdpHolePunchClientError, + UdpHolePunchPeerSource, UdpHolePunchRuntime, UdpHolePunchSignaling, UdpHolePunchTransportSink, + UdpNatType, UdpPunchClientMethod, UdpPunchSocket, UdpPunchTaskInfo, UdpSymToConePunchClient, + collect_udp_punch_tasks, punch_cone_to_cone, should_blacklist_signal_error, +}; + +#[derive(Clone, Default)] +pub struct UdpSymPunchLock { + inner: Arc>, +} + +impl UdpSymPunchLock { + pub(crate) async fn lock(&self) -> OwnedMutexGuard<()> { + self.inner.clone().lock_owned().await + } + + pub(crate) fn try_lock(&self) -> Result, TryLockError> { + self.inner.clone().try_lock_owned() + } +} + +struct UdpHolePunchBlacklist { + items: DashMap, +} + +impl UdpHolePunchBlacklist { + fn new() -> Self { + Self { + items: DashMap::new(), + } + } + + fn contains(&self, peer_id: PeerId) -> bool { + let Some(insert_time) = self.items.get(&peer_id) else { + return false; + }; + let expired = insert_time.elapsed().as_secs() >= BLACKLIST_TIMEOUT_SEC; + drop(insert_time); + + if expired { + self.items.remove(&peer_id); + false + } else { + true + } + } + + fn insert(&self, peer_id: PeerId) { + self.items.insert(peer_id, Instant::now()); + } + + fn cleanup(&self) { + self.items + .retain(|_, insert_time| insert_time.elapsed().as_secs() < BLACKLIST_TIMEOUT_SEC); + } +} + +struct UdpHolePunchConnectorParts +where + P: UdpHolePunchPeerSource + 'static, + S: UdpHolePunchSignaling + 'static, + T: UdpHolePunchTransportSink + 'static, + R: UdpHolePunchRuntime, +{ + peer_source: Arc

, + signaling: Arc, + transport_sink: Arc, + runtime: Arc, + stun: Arc, + sym_punch_lock: UdpSymPunchLock, + try_cone_before_sym: AtomicBool, +} + +pub struct UdpHolePunchConnectorData +where + P: UdpHolePunchPeerSource + 'static, + S: UdpHolePunchSignaling + 'static, + T: UdpHolePunchTransportSink + 'static, + R: UdpHolePunchRuntime, +{ + peer_source: Arc

, + signaling: Arc, + transport_sink: Arc, + runtime: Arc, + stun: Arc, + sym_punch_lock: UdpSymPunchLock, + blacklist: UdpHolePunchBlacklist, + try_cone_before_sym: Arc, + pub sym_to_cone_client: UdpSymToConePunchClient, + pub both_easy_sym_client: UdpBothEasySymPunchClient, +} + +impl UdpHolePunchConnectorData +where + P: UdpHolePunchPeerSource + 'static, + S: UdpHolePunchSignaling + 'static, + T: UdpHolePunchTransportSink + 'static, + R: UdpHolePunchRuntime, +{ + fn new(parts: Arc>) -> Arc { + Arc::new(Self { + peer_source: parts.peer_source.clone(), + signaling: parts.signaling.clone(), + transport_sink: parts.transport_sink.clone(), + runtime: parts.runtime.clone(), + stun: parts.stun.clone(), + sym_punch_lock: parts.sym_punch_lock.clone(), + blacklist: UdpHolePunchBlacklist::new(), + try_cone_before_sym: Arc::new(AtomicBool::new( + parts.try_cone_before_sym.load(Ordering::Relaxed), + )), + sym_to_cone_client: UdpSymToConePunchClient::new( + parts.runtime.clone(), + parts.signaling.clone(), + parts.stun.clone(), + ), + both_easy_sym_client: UdpBothEasySymPunchClient::new( + parts.runtime.clone(), + parts.signaling.clone(), + parts.stun.clone(), + ), + }) + } + + fn should_skip_blacklisted(&self, peer_id: PeerId) -> bool { + if self.blacklist.contains(peer_id) { + tracing::debug!( + dst_peer_id = peer_id, + "peer is blacklisted, skipping hole punching" + ); + true + } else { + false + } + } + + fn map_client_result( + &self, + dst_peer_id: PeerId, + ret: Result, UdpHolePunchClientError>, + ) -> Result, Error> { + match ret { + Ok(ret) => Ok(ret), + Err(UdpHolePunchClientError::Signaling(err)) => { + if should_blacklist_signal_error(&err) { + self.blacklist.insert(dst_peer_id); + } + Err(err.into()) + } + Err(err) => Err(err.into()), + } + } + + #[tracing::instrument(skip(self))] + async fn handle_punch_result( + &self, + ret: Result, Error>, + backoff: Option<&mut BackOff>, + round: Option<&mut u32>, + ) -> bool { + let op = |rollback: bool| { + if rollback { + if let Some(backoff) = backoff { + backoff.rollback(); + } + if let Some(round) = round { + *round = round.saturating_sub(1); + } + } else if let Some(round) = round { + *round += 1; + } + }; + + match ret { + Ok(Some(socket)) => { + let (connected, requested_url) = socket.into_connected(); + if let Err(err) = self + .transport_sink + .add_client_transport(connected, requested_url) + .await + { + tracing::warn!(?err, "upgrade or add UDP hole-punch transport failed"); + op(true); + false + } else { + tracing::info!("hole punching transport admitted successfully"); + true + } + } + Ok(None) => { + tracing::info!("hole punching failed, no punched socket"); + op(false); + false + } + Err(err) => { + tracing::info!(?err, "hole punching failed"); + op(true); + false + } + } + } + + #[tracing::instrument(skip(self))] + async fn cone_to_cone(self: Arc, task_info: UdpPunchTaskInfo) -> Result<(), Error> { + let mut backoff = BackOff::new(vec![1000, 1000, 2000, 4000, 4000, 8000, 8000, 16000]); + + loop { + backoff.sleep_for_next_backoff().await; + + if self.should_skip_blacklisted(task_info.dst_peer_id) { + break; + } + + let ret = punch_cone_to_cone( + self.runtime.clone(), + self.signaling.clone(), + task_info.dst_peer_id, + ) + .await; + let ret = self.map_client_result(task_info.dst_peer_id, ret); + + if self + .handle_punch_result(ret, Some(&mut backoff), None) + .await + { + break; + } + } + + Ok(()) + } + + #[tracing::instrument(skip(self))] + async fn sym_to_cone(self: Arc, task_info: UdpPunchTaskInfo) -> Result<(), Error> { + let mut backoff = + BackOff::new(vec![1000, 1000, 2000, 4000, 4000, 8000, 8000, 16000, 64000]); + let mut round = 0; + let mut port_idx = rand::random(); + + loop { + backoff.sleep_for_next_backoff().await; + + if self.should_skip_blacklisted(task_info.dst_peer_id) { + break; + } + + if self.try_cone_before_sym.load(Ordering::Relaxed) { + let ret = punch_cone_to_cone( + self.runtime.clone(), + self.signaling.clone(), + task_info.dst_peer_id, + ) + .await; + let ret = self.map_client_result(task_info.dst_peer_id, ret); + if self.handle_punch_result(ret, None, None).await { + break; + } + if self.should_skip_blacklisted(task_info.dst_peer_id) { + break; + } + } + + let ret = { + let _lock = self.sym_punch_lock.lock().await; + self.sym_to_cone_client + .do_hole_punching( + task_info.dst_peer_id, + round, + &mut port_idx, + task_info.my_nat_type, + ) + .await + }; + let ret = self.map_client_result(task_info.dst_peer_id, ret); + + if self + .handle_punch_result(ret, Some(&mut backoff), Some(&mut round)) + .await + { + break; + } + } + + Ok(()) + } + + #[tracing::instrument(skip(self))] + async fn both_easy_sym(self: Arc, task_info: UdpPunchTaskInfo) -> Result<(), Error> { + let mut backoff = + BackOff::new(vec![1000, 1000, 2000, 4000, 4000, 8000, 8000, 16000, 64000]); + + loop { + backoff.sleep_for_next_backoff().await; + + if self.should_skip_blacklisted(task_info.dst_peer_id) { + break; + } + + if self.try_cone_before_sym.load(Ordering::Relaxed) { + let ret = punch_cone_to_cone( + self.runtime.clone(), + self.signaling.clone(), + task_info.dst_peer_id, + ) + .await; + let ret = self.map_client_result(task_info.dst_peer_id, ret); + if self.handle_punch_result(ret, None, None).await { + break; + } + if self.should_skip_blacklisted(task_info.dst_peer_id) { + break; + } + } + + let mut is_busy = false; + let ret = { + let _lock = self.sym_punch_lock.lock().await; + self.both_easy_sym_client + .do_hole_punching( + task_info.dst_peer_id, + task_info.my_nat_type, + task_info.dst_nat_type, + &mut is_busy, + ) + .await + }; + let ret = self.map_client_result(task_info.dst_peer_id, ret); + + if is_busy { + backoff.rollback(); + } else if self + .handle_punch_result(ret, Some(&mut backoff), None) + .await + { + break; + } + } + + Ok(()) + } +} + +struct UdpHolePunchPeerTaskLauncher(Arc>) +where + P: UdpHolePunchPeerSource + 'static, + S: UdpHolePunchSignaling + 'static, + T: UdpHolePunchTransportSink + 'static, + R: UdpHolePunchRuntime; + +impl Clone for UdpHolePunchPeerTaskLauncher +where + P: UdpHolePunchPeerSource + 'static, + S: UdpHolePunchSignaling + 'static, + T: UdpHolePunchTransportSink + 'static, + R: UdpHolePunchRuntime, +{ + fn clone(&self) -> Self { + Self(self.0.clone()) + } +} + +#[async_trait::async_trait] +impl PeerTaskLauncher for UdpHolePunchPeerTaskLauncher +where + P: UdpHolePunchPeerSource + 'static, + S: UdpHolePunchSignaling + 'static, + T: UdpHolePunchTransportSink + 'static, + R: UdpHolePunchRuntime, +{ + type CollectPeerItem = UdpPunchTaskInfo; + type TaskRet = (); + + async fn collect_peers_need_task(&self) -> Vec { + let data = &self.0; + let my_nat_type = data.stun.get_stun_info().udp_nat_type; + let my_nat_type: UdpNatType = NatType::try_from(my_nat_type) + .unwrap_or(NatType::Unknown) + .into(); + if !my_nat_type.is_sym() { + data.sym_to_cone_client.clear_udp_array().await; + } + + if my_nat_type.is_open() { + return Vec::new(); + } + + data.blacklist.cleanup(); + + let my_peer_id = data.peer_source.local_peer_id(); + let policy = data.peer_source.p2p_policy_flags(); + let candidates = data.peer_source.candidates().await; + let peers_to_connect = + collect_udp_punch_tasks(my_peer_id, my_nat_type, policy, candidates, |peer_id| { + data.blacklist.contains(peer_id) + }); + for task in &peers_to_connect { + tracing::info!( + peer_id = task.dst_peer_id, + peer_nat_type = ?task.dst_nat_type, + ?my_nat_type, + "found peer to do hole punching" + ); + } + + peers_to_connect + } + + async fn launch_task( + &self, + item: Self::CollectPeerItem, + ) -> JoinHandle> { + let data = self.0.clone(); + let disable_sym_hole_punching = data + .peer_source + .p2p_policy_flags() + .disable_sym_hole_punching; + let punch_method = item + .my_nat_type + .get_punch_hole_method(item.dst_nat_type, disable_sym_hole_punching); + match punch_method { + UdpPunchClientMethod::ConeToCone => tokio::spawn(data.cone_to_cone(item)), + UdpPunchClientMethod::SymToCone => tokio::spawn(data.sym_to_cone(item)), + UdpPunchClientMethod::EasySymToEasySym => tokio::spawn(data.both_easy_sym(item)), + _ => unreachable!(), + } + } + + async fn all_task_done(&self) { + self.0.sym_to_cone_client.clear_udp_array().await; + } + + fn loop_interval_ms(&self) -> u64 { + 5000 + } +} + +pub struct UdpHolePunchConnector +where + P: UdpHolePunchPeerSource + 'static, + S: UdpHolePunchSignaling + 'static, + T: UdpHolePunchTransportSink + 'static, + R: UdpHolePunchRuntime, +{ + client: PeerTaskManager>, +} + +impl UdpHolePunchConnector +where + P: UdpHolePunchPeerSource + 'static, + S: UdpHolePunchSignaling + 'static, + T: UdpHolePunchTransportSink + 'static, + R: UdpHolePunchRuntime, +{ + pub fn new( + peer_source: Arc

, + signaling: Arc, + transport_sink: Arc, + runtime: Arc, + stun: Arc, + sym_punch_lock: UdpSymPunchLock, + external_signal: Option>, + ) -> Self { + let parts = Arc::new(UdpHolePunchConnectorParts { + peer_source, + signaling, + transport_sink, + runtime, + stun, + sym_punch_lock, + try_cone_before_sym: AtomicBool::new(true), + }); + let data = UdpHolePunchConnectorData::new(parts); + Self { + client: PeerTaskManager::new_with_external_signal( + UdpHolePunchPeerTaskLauncher(data), + external_signal, + ), + } + } + + pub fn run_as_client(&self) { + self.client.start(); + } + + pub async fn stop(&self) { + self.client.stop().await; + } +} + +#[cfg(test)] +mod tests { + use super::UdpSymPunchLock; + + #[test] + fn symmetric_punch_locks_are_scoped_per_instance() { + let first = UdpSymPunchLock::default(); + let first_clone = first.clone(); + let second = UdpSymPunchLock::default(); + + let _first_guard = first.try_lock().unwrap(); + assert!(first_clone.try_lock().is_err()); + assert!(second.try_lock().is_ok()); + } +} diff --git a/easytier-core/src/connectivity/hole_punch/udp/mod.rs b/easytier-core/src/connectivity/hole_punch/udp/mod.rs new file mode 100644 index 00000000..22c09fe7 --- /dev/null +++ b/easytier-core/src/connectivity/hole_punch/udp/mod.rs @@ -0,0 +1,39 @@ +mod binding; +mod client; +mod common; +mod connector; +mod punch_listener; +mod rpc; +mod runtime; +mod server; +mod socket_array; +mod task; + +pub(crate) use binding::CoreUdpHolePunchService; +pub(crate) use client::{ + UdpBothEasySymPunchClient, UdpHolePunchClientError, UdpSymToConePunchClient, punch_cone_to_cone, +}; +pub(crate) use common::{BLACKLIST_TIMEOUT_SEC, UdpNatType, UdpPunchClientMethod}; +pub(crate) use connector::{UdpHolePunchConnector, UdpSymPunchLock}; +pub(crate) use punch_listener::{ + MAX_PUBLIC_UDP_HOLE_PUNCH_LISTENERS, ReusableUdpPunchListener, can_reuse_port_mapping_listener, + can_reuse_public_listener, select_reusable_port_mapping_listener_idx, + select_reusable_public_listener_idx, should_create_public_listener, + should_retry_public_listener_selection, +}; +pub(crate) use rpc::UdpHolePunchRpcSource; +pub(crate) use runtime::{ + ProtocolUdpHolePunchTransportSink, SelectPunchListener, SelectPunchListenerResponse, + SendPunchPacketBothEasySym, SendPunchPacketBothEasySymResponse, SendPunchPacketCone, + SendPunchPacketEasySym, SendPunchPacketHardSym, SendPunchPacketHardSymResponse, + UdpHolePunchInbound, UdpHolePunchPeerSource, UdpHolePunchRuntime, UdpHolePunchSignalError, + UdpHolePunchSignaling, UdpHolePunchTransportSink, UdpPunchAcceptor, UdpPunchListener, + UdpPunchSocket, UdpResolvedPublicAddr, should_blacklist_signal_error, +}; +pub(crate) use server::UdpHolePunchServer; +pub(crate) use socket_array::UdpSocketArray; +pub(crate) use task::{UdpPunchCandidate, UdpPunchTaskInfo, collect_udp_punch_tasks}; + +const fn udp_packet_len(body_len: u16) -> usize { + crate::packet::UDP_TUNNEL_HEADER_SIZE + body_len as usize +} diff --git a/easytier-core/src/connectivity/hole_punch/udp/punch_listener.rs b/easytier-core/src/connectivity/hole_punch/udp/punch_listener.rs new file mode 100644 index 00000000..fd0b776f --- /dev/null +++ b/easytier-core/src/connectivity/hole_punch/udp/punch_listener.rs @@ -0,0 +1,195 @@ +use std::net::SocketAddr; + +use quanta::Instant; + +pub const MAX_PUBLIC_UDP_HOLE_PUNCH_LISTENERS: usize = 4; + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub struct ReusableUdpPunchListener { + pub running: bool, + pub mapped_addr: SocketAddr, + pub has_port_mapping_lease: bool, + pub last_active_time: Instant, +} + +pub fn can_reuse_public_listener(listener: &ReusableUdpPunchListener) -> bool { + listener.running && !listener.mapped_addr.ip().is_unspecified() +} + +pub fn can_reuse_port_mapping_listener(listener: &ReusableUdpPunchListener) -> bool { + can_reuse_public_listener(listener) && listener.has_port_mapping_lease +} + +pub fn select_reusable_public_listener_idx( + listeners: &[ReusableUdpPunchListener], +) -> Option { + listeners + .iter() + .enumerate() + .filter(|(_, listener)| can_reuse_public_listener(listener)) + .max_by_key(|(_, listener)| listener.last_active_time) + .map(|(idx, _)| idx) +} + +pub fn select_reusable_port_mapping_listener_idx( + listeners: &[ReusableUdpPunchListener], +) -> Option { + listeners + .iter() + .enumerate() + .filter(|(_, listener)| can_reuse_port_mapping_listener(listener)) + .max_by_key(|(_, listener)| listener.last_active_time) + .map(|(idx, _)| idx) +} + +pub fn should_create_public_listener( + current_listener_count: usize, + has_reusable_listener: bool, + has_port_mapping_listener: bool, + force_new_listener: bool, + prefer_port_mapping: bool, +) -> bool { + if current_listener_count >= MAX_PUBLIC_UDP_HOLE_PUNCH_LISTENERS { + return false; + } + + if current_listener_count == 0 { + return true; + } + + if force_new_listener { + return true; + } + + if prefer_port_mapping && !has_port_mapping_listener { + return true; + } + + !has_reusable_listener +} + +pub fn should_retry_public_listener_selection( + force_new_listener: bool, + current_listener_count: usize, + prefer_port_mapping: bool, + has_port_mapping_listener: bool, +) -> bool { + if prefer_port_mapping && has_port_mapping_listener { + return false; + } + + !force_new_listener && current_listener_count < MAX_PUBLIC_UDP_HOLE_PUNCH_LISTENERS +} + +#[cfg(test)] +mod tests { + use std::{ + net::{Ipv4Addr, SocketAddr, SocketAddrV4}, + time::Duration, + }; + + use super::*; + + fn listener( + port: u16, + running: bool, + has_port_mapping_lease: bool, + active_age: Duration, + ) -> ReusableUdpPunchListener { + ReusableUdpPunchListener { + running, + mapped_addr: SocketAddr::V4(SocketAddrV4::new(Ipv4Addr::LOCALHOST, port)), + has_port_mapping_lease, + last_active_time: Instant::now() - active_age, + } + } + + #[test] + fn listener_selection_prefers_reuse_before_cap() { + assert!(!should_create_public_listener(1, true, true, false, false)); + assert!(!should_create_public_listener( + MAX_PUBLIC_UDP_HOLE_PUNCH_LISTENERS, + true, + true, + false, + false + )); + } + + #[test] + fn listener_selection_creates_when_empty_or_no_reusable_listener() { + assert!(should_create_public_listener(0, false, false, false, false)); + assert!(should_create_public_listener(1, false, false, false, false)); + } + + #[test] + fn listener_selection_force_new_respects_cap() { + assert!(should_create_public_listener(1, true, true, true, false)); + assert!(!should_create_public_listener( + MAX_PUBLIC_UDP_HOLE_PUNCH_LISTENERS, + true, + true, + true, + false + )); + } + + #[test] + fn listener_selection_prefers_port_mapping_until_available() { + assert!(should_create_public_listener(1, true, false, false, true)); + assert!(!should_create_public_listener(1, true, true, false, true)); + } + + #[test] + fn listener_selection_retry_respects_cap() { + assert!(should_retry_public_listener_selection( + false, 1, false, false + )); + assert!(!should_retry_public_listener_selection( + false, + MAX_PUBLIC_UDP_HOLE_PUNCH_LISTENERS, + false, + false + )); + assert!(!should_retry_public_listener_selection( + true, 1, false, false + )); + assert!(!should_retry_public_listener_selection( + false, 1, true, true + )); + } + + #[test] + fn selects_most_recent_reusable_public_listener() { + let listeners = vec![ + listener(1000, true, false, Duration::from_secs(10)), + listener(1001, false, false, Duration::from_secs(1)), + listener(1002, true, false, Duration::from_secs(2)), + ]; + + assert_eq!(select_reusable_public_listener_idx(&listeners), Some(2)); + } + + #[test] + fn selects_most_recent_reusable_port_mapping_listener() { + let listeners = vec![ + listener(1000, true, false, Duration::from_secs(1)), + listener(1001, true, true, Duration::from_secs(10)), + listener(1002, true, true, Duration::from_secs(2)), + ]; + + assert_eq!( + select_reusable_port_mapping_listener_idx(&listeners), + Some(2) + ); + } + + #[test] + fn unspecified_addr_is_not_reusable() { + let mut listener = listener(1000, true, true, Duration::ZERO); + listener.mapped_addr = SocketAddr::V4(SocketAddrV4::new(Ipv4Addr::UNSPECIFIED, 1000)); + + assert!(!can_reuse_public_listener(&listener)); + assert!(!can_reuse_port_mapping_listener(&listener)); + } +} diff --git a/easytier-core/src/connectivity/hole_punch/udp/rpc.rs b/easytier-core/src/connectivity/hole_punch/udp/rpc.rs new file mode 100644 index 00000000..e5ea1d4a --- /dev/null +++ b/easytier-core/src/connectivity/hole_punch/udp/rpc.rs @@ -0,0 +1,828 @@ +use std::{fmt, net::SocketAddr, sync::Arc}; + +use async_trait::async_trait; + +use crate::{ + config::PeerId, + connectivity::hole_punch::udp::{ + SelectPunchListener, SelectPunchListenerResponse as CoreSelectPunchListenerResponse, + SendPunchPacketBothEasySym, + SendPunchPacketBothEasySymResponse as CoreSendPunchPacketBothEasySymResponse, + SendPunchPacketCone, SendPunchPacketEasySym, SendPunchPacketHardSym, + SendPunchPacketHardSymResponse as CoreSendPunchPacketHardSymResponse, UdpHolePunchInbound, + UdpHolePunchRuntime, UdpHolePunchServer as CoreUdpHolePunchServer, UdpHolePunchSignalError, + UdpHolePunchSignaling, UdpHolePunchTransportSink, UdpSymPunchLock, + }, + connectivity::stun::StunInfoProvider, + proto::{ + common::Void, + peer_rpc::{ + SelectPunchListenerRequest, SelectPunchListenerResponse, + SendPunchPacketBothEasySymRequest, SendPunchPacketBothEasySymResponse, + SendPunchPacketConeRequest, SendPunchPacketEasySymRequest, + SendPunchPacketHardSymRequest, SendPunchPacketHardSymResponse, UdpHolePunchRpc, + }, + rpc_types::{self, controller::BaseController}, + }, +}; + +const CONE_RPC_TIMEOUT_MS: i32 = 4000; +const SYMMETRIC_RPC_TIMEOUT_MS: i32 = 4000; +const BOTH_EASY_SYMMETRIC_RPC_TIMEOUT_MS: i32 = 2000; + +fn cone_controller() -> BaseController { + BaseController { + timeout_ms: CONE_RPC_TIMEOUT_MS, + ..Default::default() + } +} + +fn symmetric_controller() -> BaseController { + BaseController { + timeout_ms: SYMMETRIC_RPC_TIMEOUT_MS, + trace_id: 0, + ..Default::default() + } +} + +fn both_easy_symmetric_controller() -> BaseController { + BaseController { + timeout_ms: BOTH_EASY_SYMMETRIC_RPC_TIMEOUT_MS, + ..Default::default() + } +} + +fn select_listener_request_to_rpc(request: SelectPunchListener) -> SelectPunchListenerRequest { + SelectPunchListenerRequest { + force_new: request.force_new, + prefer_port_mapping: request.prefer_port_mapping, + } +} + +fn select_listener_request_from_rpc(input: SelectPunchListenerRequest) -> SelectPunchListener { + SelectPunchListener { + force_new: input.force_new, + prefer_port_mapping: input.prefer_port_mapping, + } +} + +fn select_listener_response_from_rpc( + response: SelectPunchListenerResponse, +) -> Result { + Ok(CoreSelectPunchListenerResponse { + listener_mapped_addr: SocketAddr::from( + response + .listener_mapped_addr + .ok_or_else(|| missing_field("listener_mapped_addr"))?, + ), + }) +} + +fn select_listener_response_to_rpc( + response: CoreSelectPunchListenerResponse, +) -> SelectPunchListenerResponse { + SelectPunchListenerResponse { + listener_mapped_addr: Some(response.listener_mapped_addr.into()), + } +} + +fn cone_request_to_rpc(request: SendPunchPacketCone) -> SendPunchPacketConeRequest { + SendPunchPacketConeRequest { + listener_mapped_addr: Some(request.listener_mapped_addr.into()), + dest_addr: Some(request.dest_addr.into()), + transaction_id: request.transaction_id, + packet_count_per_batch: request.packet_count_per_batch, + packet_batch_count: request.packet_batch_count, + packet_interval_ms: request.packet_interval_ms, + } +} + +fn hard_symmetric_request_to_rpc(request: SendPunchPacketHardSym) -> SendPunchPacketHardSymRequest { + SendPunchPacketHardSymRequest { + listener_mapped_addr: Some(request.listener_mapped_addr.into()), + public_ips: request.public_ips.into_iter().map(Into::into).collect(), + transaction_id: request.transaction_id, + port_index: request.port_index, + round: request.round, + } +} + +fn easy_symmetric_request_to_rpc(request: SendPunchPacketEasySym) -> SendPunchPacketEasySymRequest { + SendPunchPacketEasySymRequest { + listener_mapped_addr: Some(request.listener_mapped_addr.into()), + public_ips: request.public_ips.into_iter().map(Into::into).collect(), + transaction_id: request.transaction_id, + base_port_num: request.base_port_num, + max_port_num: request.max_port_num, + is_incremental: request.is_incremental, + } +} + +fn both_easy_symmetric_request_to_rpc( + request: SendPunchPacketBothEasySym, +) -> SendPunchPacketBothEasySymRequest { + SendPunchPacketBothEasySymRequest { + transaction_id: request.transaction_id, + public_ip: Some(request.public_ip.into()), + dst_port_num: request.dst_port_num, + udp_socket_count: request.udp_socket_count, + wait_time_ms: request.wait_time_ms, + } +} + +fn hard_symmetric_response_from_rpc( + response: SendPunchPacketHardSymResponse, +) -> CoreSendPunchPacketHardSymResponse { + CoreSendPunchPacketHardSymResponse { + next_port_index: response.next_port_index, + } +} + +fn hard_symmetric_response_to_rpc( + response: CoreSendPunchPacketHardSymResponse, +) -> SendPunchPacketHardSymResponse { + SendPunchPacketHardSymResponse { + next_port_index: response.next_port_index, + } +} + +fn both_easy_symmetric_response_from_rpc( + response: SendPunchPacketBothEasySymResponse, +) -> CoreSendPunchPacketBothEasySymResponse { + CoreSendPunchPacketBothEasySymResponse { + is_busy: response.is_busy, + base_mapped_addr: response.base_mapped_addr.map(SocketAddr::from), + } +} + +fn both_easy_symmetric_response_to_rpc( + response: CoreSendPunchPacketBothEasySymResponse, +) -> SendPunchPacketBothEasySymResponse { + SendPunchPacketBothEasySymResponse { + is_busy: response.is_busy, + base_mapped_addr: response.base_mapped_addr.map(Into::into), + } +} + +/// Narrow source of peer-scoped UDP hole-punch RPC stubs. +/// +/// Implemented only by the sealed peer adapter in `super::peer_adapters`. +pub trait UdpHolePunchRpcSource: Send + Sync + 'static { + fn local_peer_id(&self) -> PeerId; + + fn rpc_stub( + &self, + dst_peer_id: PeerId, + ) -> Box + Send + Sync + 'static>; +} + +#[derive(Clone)] +pub(super) struct PeerRpcUdpHolePunchSignaling

+where + P: UdpHolePunchRpcSource, +{ + rpc_source: Arc

, +} + +impl

fmt::Debug for PeerRpcUdpHolePunchSignaling

+where + P: UdpHolePunchRpcSource, +{ + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter + .debug_struct("PeerRpcUdpHolePunchSignaling") + .field("my_peer_id", &self.rpc_source.local_peer_id()) + .finish_non_exhaustive() + } +} + +impl

PeerRpcUdpHolePunchSignaling

+where + P: UdpHolePunchRpcSource, +{ + pub(super) fn new(rpc_source: Arc

) -> Self { + Self { rpc_source } + } + + fn rpc_stub( + &self, + dst_peer_id: PeerId, + ) -> Box + Send + Sync + 'static> { + self.rpc_source.rpc_stub(dst_peer_id) + } +} + +fn map_rpc_error(error: rpc_types::error::Error) -> UdpHolePunchSignalError { + match error { + rpc_types::error::Error::InvalidServiceKey(_, _) => { + UdpHolePunchSignalError::InvalidServiceKey + } + rpc_types::error::Error::Timeout(_) => UdpHolePunchSignalError::Timeout, + rpc_types::error::Error::ExecutionError(error) => { + UdpHolePunchSignalError::RemoteRejected(error.to_string()) + } + other => UdpHolePunchSignalError::Transport(other.to_string()), + } +} + +fn missing_field(field: &str) -> UdpHolePunchSignalError { + UdpHolePunchSignalError::RemoteRejected(format!("missing {field}")) +} + +#[async_trait] +impl

UdpHolePunchSignaling for PeerRpcUdpHolePunchSignaling

+where + P: UdpHolePunchRpcSource, +{ + async fn select_punch_listener( + &self, + dst_peer_id: PeerId, + request: SelectPunchListener, + ) -> Result { + let response = self + .rpc_stub(dst_peer_id) + .select_punch_listener( + BaseController::default(), + select_listener_request_to_rpc(request), + ) + .await + .map_err(map_rpc_error)?; + + select_listener_response_from_rpc(response) + } + + async fn send_punch_packet_cone( + &self, + dst_peer_id: PeerId, + request: SendPunchPacketCone, + ) -> Result<(), UdpHolePunchSignalError> { + self.rpc_stub(dst_peer_id) + .send_punch_packet_cone(cone_controller(), cone_request_to_rpc(request)) + .await + .map(|_| ()) + .map_err(map_rpc_error) + } + + async fn send_punch_packet_hard_sym( + &self, + dst_peer_id: PeerId, + request: SendPunchPacketHardSym, + ) -> Result { + let response = self + .rpc_stub(dst_peer_id) + .send_punch_packet_hard_sym( + symmetric_controller(), + hard_symmetric_request_to_rpc(request), + ) + .await + .map_err(map_rpc_error)?; + + Ok(hard_symmetric_response_from_rpc(response)) + } + + async fn send_punch_packet_easy_sym( + &self, + dst_peer_id: PeerId, + request: SendPunchPacketEasySym, + ) -> Result<(), UdpHolePunchSignalError> { + self.rpc_stub(dst_peer_id) + .send_punch_packet_easy_sym( + symmetric_controller(), + easy_symmetric_request_to_rpc(request), + ) + .await + .map(|_| ()) + .map_err(map_rpc_error) + } + + async fn send_punch_packet_both_easy_sym( + &self, + dst_peer_id: PeerId, + request: SendPunchPacketBothEasySym, + ) -> Result { + let response = self + .rpc_stub(dst_peer_id) + .send_punch_packet_both_easy_sym( + both_easy_symmetric_controller(), + both_easy_symmetric_request_to_rpc(request), + ) + .await + .map_err(map_rpc_error)?; + + Ok(both_easy_symmetric_response_from_rpc(response)) + } +} + +pub(super) struct UdpHolePunchRpcEndpoint +where + R: UdpHolePunchRuntime, + T: UdpHolePunchTransportSink + 'static, +{ + inner: CoreUdpHolePunchServer, +} + +impl UdpHolePunchRpcEndpoint +where + R: UdpHolePunchRuntime, + T: UdpHolePunchTransportSink + 'static, +{ + pub(super) fn new( + stun: Arc, + transport_sink: Arc, + sym_punch_lock: UdpSymPunchLock, + runtime: Arc, + ) -> Arc { + let inner = CoreUdpHolePunchServer::new(runtime, stun, transport_sink, sym_punch_lock); + Arc::new(Self { inner }) + } + + pub(super) async fn start(&self) { + self.inner.start().await; + } + + pub(super) fn begin_stop(&self) { + self.inner.begin_stop(); + } + + pub(super) async fn stop(&self) { + self.inner.stop().await; + } +} + +fn signal_error_to_rpc_error(error: UdpHolePunchSignalError) -> rpc_types::error::Error { + match error { + UdpHolePunchSignalError::InvalidServiceKey => rpc_types::error::Error::InvalidServiceKey( + "UdpHolePunchRpc".to_owned(), + "UdpHolePunchRpc".to_owned(), + ), + UdpHolePunchSignalError::Timeout => anyhow::anyhow!("timeout").into(), + UdpHolePunchSignalError::RemoteRejected(message) + | UdpHolePunchSignalError::Transport(message) => anyhow::anyhow!(message).into(), + } +} + +fn cone_request_from_rpc( + input: SendPunchPacketConeRequest, +) -> rpc_types::error::Result { + let listener_addr = input.listener_mapped_addr.ok_or(anyhow::anyhow!( + "send_punch_packet_for_cone request missing listener_mapped_addr" + ))?; + let dest_addr = input.dest_addr.ok_or(anyhow::anyhow!( + "send_punch_packet_for_cone request missing dest_addr" + ))?; + Ok(SendPunchPacketCone { + listener_mapped_addr: listener_addr.into(), + dest_addr: dest_addr.into(), + transaction_id: input.transaction_id, + packet_count_per_batch: input.packet_count_per_batch, + packet_batch_count: input.packet_batch_count, + packet_interval_ms: input.packet_interval_ms, + }) +} + +fn hard_symmetric_request_from_rpc( + input: SendPunchPacketHardSymRequest, +) -> rpc_types::error::Result { + let listener_addr = input.listener_mapped_addr.ok_or(anyhow::anyhow!( + "try_punch_symmetric request missing listener_addr" + ))?; + Ok(SendPunchPacketHardSym { + listener_mapped_addr: listener_addr.into(), + public_ips: input.public_ips.into_iter().map(Into::into).collect(), + transaction_id: input.transaction_id, + port_index: input.port_index, + round: input.round, + }) +} + +fn easy_symmetric_request_from_rpc( + input: SendPunchPacketEasySymRequest, +) -> rpc_types::error::Result { + let listener_addr = input.listener_mapped_addr.ok_or(anyhow::anyhow!( + "send_punch_packet_easy_sym request missing listener_addr" + ))?; + Ok(SendPunchPacketEasySym { + listener_mapped_addr: listener_addr.into(), + public_ips: input.public_ips.into_iter().map(Into::into).collect(), + transaction_id: input.transaction_id, + base_port_num: input.base_port_num, + max_port_num: input.max_port_num, + is_incremental: input.is_incremental, + }) +} + +fn both_easy_symmetric_request_from_rpc( + input: SendPunchPacketBothEasySymRequest, +) -> rpc_types::error::Result { + let public_ip = input + .public_ip + .ok_or(anyhow::anyhow!("public_ip is required"))?; + Ok(SendPunchPacketBothEasySym { + transaction_id: input.transaction_id, + public_ip: public_ip.into(), + dst_port_num: input.dst_port_num, + udp_socket_count: input.udp_socket_count, + wait_time_ms: input.wait_time_ms, + }) +} + +#[async_trait] +impl UdpHolePunchInbound for UdpHolePunchRpcEndpoint +where + R: UdpHolePunchRuntime, + T: UdpHolePunchTransportSink + 'static, +{ + async fn select_punch_listener( + &self, + request: SelectPunchListener, + ) -> Result { + self.inner.select_punch_listener(request).await + } + + async fn send_punch_packet_cone( + &self, + request: SendPunchPacketCone, + ) -> Result<(), UdpHolePunchSignalError> { + self.inner.send_punch_packet_cone(request).await + } + + async fn send_punch_packet_hard_sym( + &self, + request: SendPunchPacketHardSym, + ) -> Result { + self.inner.send_punch_packet_hard_sym(request).await + } + + async fn send_punch_packet_easy_sym( + &self, + request: SendPunchPacketEasySym, + ) -> Result<(), UdpHolePunchSignalError> { + self.inner.send_punch_packet_easy_sym(request).await + } + + async fn send_punch_packet_both_easy_sym( + &self, + request: SendPunchPacketBothEasySym, + ) -> Result { + self.inner.send_punch_packet_both_easy_sym(request).await + } +} + +#[async_trait] +impl UdpHolePunchRpc for UdpHolePunchRpcEndpoint +where + R: UdpHolePunchRuntime, + T: UdpHolePunchTransportSink + 'static, +{ + type Controller = BaseController; + + async fn select_punch_listener( + &self, + _controller: Self::Controller, + input: SelectPunchListenerRequest, + ) -> rpc_types::error::Result { + let response = UdpHolePunchInbound::select_punch_listener( + self, + select_listener_request_from_rpc(input), + ) + .await + .map_err(signal_error_to_rpc_error)?; + + Ok(select_listener_response_to_rpc(response)) + } + + async fn send_punch_packet_cone( + &self, + _controller: Self::Controller, + input: SendPunchPacketConeRequest, + ) -> rpc_types::error::Result { + UdpHolePunchInbound::send_punch_packet_cone(self, cone_request_from_rpc(input)?) + .await + .map_err(signal_error_to_rpc_error)?; + + Ok(Void::default()) + } + + async fn send_punch_packet_hard_sym( + &self, + _controller: Self::Controller, + input: SendPunchPacketHardSymRequest, + ) -> rpc_types::error::Result { + let response = UdpHolePunchInbound::send_punch_packet_hard_sym( + self, + hard_symmetric_request_from_rpc(input)?, + ) + .await + .map_err(signal_error_to_rpc_error)?; + + Ok(hard_symmetric_response_to_rpc(response)) + } + + async fn send_punch_packet_easy_sym( + &self, + _controller: Self::Controller, + input: SendPunchPacketEasySymRequest, + ) -> rpc_types::error::Result { + UdpHolePunchInbound::send_punch_packet_easy_sym( + self, + easy_symmetric_request_from_rpc(input)?, + ) + .await + .map_err(signal_error_to_rpc_error)?; + + Ok(Void::default()) + } + + async fn send_punch_packet_both_easy_sym( + &self, + _controller: Self::Controller, + input: SendPunchPacketBothEasySymRequest, + ) -> rpc_types::error::Result { + let response = UdpHolePunchInbound::send_punch_packet_both_easy_sym( + self, + both_easy_symmetric_request_from_rpc(input)?, + ) + .await + .map_err(signal_error_to_rpc_error)?; + + Ok(both_easy_symmetric_response_to_rpc(response)) + } +} + +#[cfg(test)] +mod tests { + use std::{future, net::Ipv4Addr, time::Duration}; + + use super::*; + + #[test] + fn outbound_rpc_controllers_preserve_timeouts() { + assert_eq!(cone_controller().timeout_ms, 4000); + assert_eq!(symmetric_controller().timeout_ms, 4000); + assert_eq!(symmetric_controller().trace_id, 0); + assert_eq!(both_easy_symmetric_controller().timeout_ms, 2000); + } + + #[test] + fn select_listener_dto_preserves_fields_and_requires_response_addr() { + let domain_request = SelectPunchListener { + force_new: true, + prefer_port_mapping: false, + }; + let rpc_request = select_listener_request_to_rpc(domain_request.clone()); + assert!(rpc_request.force_new); + assert!(!rpc_request.prefer_port_mapping); + assert_eq!( + select_listener_request_from_rpc(SelectPunchListenerRequest { + force_new: true, + prefer_port_mapping: false, + }), + domain_request + ); + + let mapped_addr: SocketAddr = "198.51.100.1:31001".parse().unwrap(); + let core_response = CoreSelectPunchListenerResponse { + listener_mapped_addr: mapped_addr, + }; + let rpc_response = select_listener_response_to_rpc(core_response.clone()); + assert_eq!( + SocketAddr::from(rpc_response.listener_mapped_addr.unwrap()), + mapped_addr + ); + assert_eq!( + select_listener_response_from_rpc(rpc_response).unwrap(), + core_response + ); + + let error = + select_listener_response_from_rpc(SelectPunchListenerResponse::default()).unwrap_err(); + assert_eq!( + error, + UdpHolePunchSignalError::RemoteRejected("missing listener_mapped_addr".to_owned()) + ); + } + + #[test] + fn cone_dto_round_trip_preserves_all_fields() { + let request = SendPunchPacketCone { + listener_mapped_addr: "198.51.100.2:31002".parse().unwrap(), + dest_addr: "203.0.113.2:32002".parse().unwrap(), + transaction_id: 12, + packet_count_per_batch: 3, + packet_batch_count: 4, + packet_interval_ms: 500, + }; + + let rpc = cone_request_to_rpc(request.clone()); + assert_eq!( + SocketAddr::from(rpc.listener_mapped_addr.unwrap()), + request.listener_mapped_addr + ); + assert_eq!(SocketAddr::from(rpc.dest_addr.unwrap()), request.dest_addr); + assert_eq!(rpc.transaction_id, 12); + assert_eq!(rpc.packet_count_per_batch, 3); + assert_eq!(rpc.packet_batch_count, 4); + assert_eq!(rpc.packet_interval_ms, 500); + assert_eq!(cone_request_from_rpc(rpc).unwrap(), request); + } + + #[test] + fn hard_symmetric_dto_round_trip_preserves_all_fields() { + let request = SendPunchPacketHardSym { + listener_mapped_addr: "198.51.100.3:31003".parse().unwrap(), + public_ips: vec![Ipv4Addr::new(203, 0, 113, 3), Ipv4Addr::new(203, 0, 113, 4)], + transaction_id: 13, + port_index: 17, + round: 19, + }; + + let rpc = hard_symmetric_request_to_rpc(request.clone()); + assert_eq!( + SocketAddr::from(rpc.listener_mapped_addr.unwrap()), + request.listener_mapped_addr + ); + assert_eq!( + rpc.public_ips + .iter() + .cloned() + .map(Ipv4Addr::from) + .collect::>(), + request.public_ips + ); + assert_eq!(rpc.transaction_id, 13); + assert_eq!(rpc.port_index, 17); + assert_eq!(rpc.round, 19); + assert_eq!(hard_symmetric_request_from_rpc(rpc).unwrap(), request); + } + + #[test] + fn easy_symmetric_dto_round_trip_preserves_all_fields() { + let request = SendPunchPacketEasySym { + listener_mapped_addr: "198.51.100.5:31005".parse().unwrap(), + public_ips: vec![Ipv4Addr::new(203, 0, 113, 5)], + transaction_id: 15, + base_port_num: 33000, + max_port_num: 51, + is_incremental: true, + }; + + let rpc = easy_symmetric_request_to_rpc(request.clone()); + assert_eq!( + SocketAddr::from(rpc.listener_mapped_addr.unwrap()), + request.listener_mapped_addr + ); + assert_eq!( + rpc.public_ips + .iter() + .cloned() + .map(Ipv4Addr::from) + .collect::>(), + request.public_ips + ); + assert_eq!(rpc.transaction_id, 15); + assert_eq!(rpc.base_port_num, 33000); + assert_eq!(rpc.max_port_num, 51); + assert!(rpc.is_incremental); + assert_eq!(easy_symmetric_request_from_rpc(rpc).unwrap(), request); + } + + #[test] + fn both_easy_symmetric_dto_round_trip_preserves_all_fields() { + let request = SendPunchPacketBothEasySym { + udp_socket_count: 25, + public_ip: Ipv4Addr::new(203, 0, 113, 6), + transaction_id: 16, + dst_port_num: 34000, + wait_time_ms: 2500, + }; + + let rpc = both_easy_symmetric_request_to_rpc(request.clone()); + assert_eq!(rpc.udp_socket_count, 25); + assert_eq!(Ipv4Addr::from(rpc.public_ip.unwrap()), request.public_ip); + assert_eq!(rpc.transaction_id, 16); + assert_eq!(rpc.dst_port_num, 34000); + assert_eq!(rpc.wait_time_ms, 2500); + assert_eq!(both_easy_symmetric_request_from_rpc(rpc).unwrap(), request); + } + + #[test] + fn rpc_response_dtos_preserve_all_fields() { + let hard_response = CoreSendPunchPacketHardSymResponse { + next_port_index: 41, + }; + let hard_rpc = hard_symmetric_response_to_rpc(hard_response.clone()); + assert_eq!(hard_rpc.next_port_index, 41); + assert_eq!(hard_symmetric_response_from_rpc(hard_rpc), hard_response); + + let both_response = CoreSendPunchPacketBothEasySymResponse { + is_busy: true, + base_mapped_addr: Some("198.51.100.8:31008".parse().unwrap()), + }; + let both_rpc = both_easy_symmetric_response_to_rpc(both_response.clone()); + assert!(both_rpc.is_busy); + assert_eq!( + SocketAddr::from(both_rpc.base_mapped_addr.unwrap()), + both_response.base_mapped_addr.unwrap() + ); + assert_eq!( + both_easy_symmetric_response_from_rpc(both_rpc), + both_response + ); + } + + #[test] + fn inbound_rpc_dtos_reject_missing_required_fields() { + let cone_listener_error = cone_request_from_rpc(SendPunchPacketConeRequest::default()) + .unwrap_err() + .to_string(); + assert_eq!( + cone_listener_error, + "Rust error: send_punch_packet_for_cone request missing listener_mapped_addr" + ); + + let cone_dest_error = cone_request_from_rpc(SendPunchPacketConeRequest { + listener_mapped_addr: Some("198.51.100.7:31007".parse::().unwrap().into()), + ..Default::default() + }) + .unwrap_err() + .to_string(); + assert_eq!( + cone_dest_error, + "Rust error: send_punch_packet_for_cone request missing dest_addr" + ); + + assert_eq!( + hard_symmetric_request_from_rpc(SendPunchPacketHardSymRequest::default()) + .unwrap_err() + .to_string(), + "Rust error: try_punch_symmetric request missing listener_addr" + ); + assert_eq!( + easy_symmetric_request_from_rpc(SendPunchPacketEasySymRequest::default()) + .unwrap_err() + .to_string(), + "Rust error: send_punch_packet_easy_sym request missing listener_addr" + ); + assert_eq!( + both_easy_symmetric_request_from_rpc(SendPunchPacketBothEasySymRequest::default()) + .unwrap_err() + .to_string(), + "Rust error: public_ip is required" + ); + } + + #[tokio::test] + async fn rpc_errors_keep_domain_classification() { + assert_eq!( + map_rpc_error(rpc_types::error::Error::InvalidServiceKey( + "service".to_owned(), + "proto".to_owned() + )), + UdpHolePunchSignalError::InvalidServiceKey + ); + assert_eq!( + map_rpc_error(rpc_types::error::Error::ExecutionError(anyhow::anyhow!( + "rejected" + ))), + UdpHolePunchSignalError::RemoteRejected("rejected".to_owned()) + ); + assert_eq!( + map_rpc_error(rpc_types::error::Error::TunnelError("closed".to_owned())), + UdpHolePunchSignalError::Transport("Tunnel error: closed".to_owned()) + ); + + let elapsed = tokio::time::timeout(Duration::ZERO, future::pending::<()>()) + .await + .unwrap_err(); + assert_eq!( + map_rpc_error(rpc_types::error::Error::Timeout(elapsed)), + UdpHolePunchSignalError::Timeout + ); + } + + #[test] + fn domain_errors_keep_rpc_classification() { + assert!(matches!( + signal_error_to_rpc_error(UdpHolePunchSignalError::InvalidServiceKey), + rpc_types::error::Error::InvalidServiceKey(_, _) + )); + for (domain_error, expected_message) in [ + (UdpHolePunchSignalError::Timeout, "timeout"), + ( + UdpHolePunchSignalError::RemoteRejected("rejected".to_owned()), + "rejected", + ), + ( + UdpHolePunchSignalError::Transport("closed".to_owned()), + "closed", + ), + ] { + let rpc_types::error::Error::ExecutionError(error) = + signal_error_to_rpc_error(domain_error) + else { + panic!("domain error should map to execution error"); + }; + assert_eq!(error.to_string(), expected_message); + } + } +} diff --git a/easytier-core/src/connectivity/hole_punch/udp/runtime.rs b/easytier-core/src/connectivity/hole_punch/udp/runtime.rs new file mode 100644 index 00000000..b23f16c1 --- /dev/null +++ b/easytier-core/src/connectivity/hole_punch/udp/runtime.rs @@ -0,0 +1,509 @@ +use std::{ + net::{IpAddr, Ipv4Addr, Ipv6Addr, SocketAddr}, + sync::Arc, +}; + +use async_trait::async_trait; + +use super::super::{HolePunchTunnelSink, port_mapping::UdpPortMappingLease}; + +use crate::{ + config::P2pPolicyFlags, + connectivity::{ + protocol::ClientProtocolUpgrader, + transport::{ConnectedTransport, ConnectedUdpSession}, + }, + foundation::task::ExternalTaskSignal, + socket::{ + ListenerConnectionCounter, SocketContext, + udp::{UdpBindOptions, UdpSession, VirtualUdpSocket, VirtualUdpSocketFactory}, + }, + tunnel::Tunnel, +}; + +#[async_trait] +pub trait UdpPunchAcceptor: Send { + async fn accept(&mut self) -> anyhow::Result; +} + +pub struct UdpPunchSocket { + session: UdpSession, + requested_remote_addr: SocketAddr, + lifetime_guard: Box, +} + +impl UdpPunchSocket { + pub fn new(session: UdpSession, requested_remote_addr: SocketAddr, lifetime_guard: G) -> Self + where + G: Send + Sync + 'static, + { + Self { + session, + requested_remote_addr, + lifetime_guard: Box::new(lifetime_guard), + } + } + + pub(crate) fn into_connected(self) -> (ConnectedUdpSession, url::Url) { + ( + ConnectedUdpSession::new(self.session, self.lifetime_guard), + udp_url(self.requested_remote_addr), + ) + } +} + +impl std::fmt::Debug for UdpPunchSocket { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("UdpPunchSocket") + .field("session", &self.session) + .field("requested_remote_addr", &self.requested_remote_addr) + .finish_non_exhaustive() + } +} + +fn udp_url(addr: SocketAddr) -> url::Url { + let mut url = url::Url::parse("udp://0.0.0.0").expect("static UDP URL should be valid"); + url.set_ip_host(addr.ip()) + .expect("socket IP should be a valid URL host"); + url.set_port(Some(addr.port())) + .expect("UDP URL should accept a port"); + url +} + +pub struct UdpPunchListener { + pub socket: Arc, + pub mapped_addr: SocketAddr, + pub conn_counter: Arc, + pub acceptor: Box, + pub(crate) port_mapping_lease: Option>, +} + +pub struct UdpResolvedPublicAddr { + pub mapped_addr: SocketAddr, + pub(crate) port_mapping_lease: Option>, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct SelectPunchListener { + pub force_new: bool, + pub prefer_port_mapping: bool, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct SelectPunchListenerResponse { + pub listener_mapped_addr: SocketAddr, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct SendPunchPacketCone { + pub listener_mapped_addr: SocketAddr, + pub dest_addr: SocketAddr, + pub transaction_id: u32, + pub packet_count_per_batch: u32, + pub packet_batch_count: u32, + pub packet_interval_ms: u32, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct SendPunchPacketHardSym { + pub listener_mapped_addr: SocketAddr, + pub public_ips: Vec, + pub transaction_id: u32, + pub port_index: u32, + pub round: u32, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct SendPunchPacketHardSymResponse { + pub next_port_index: u32, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct SendPunchPacketEasySym { + pub listener_mapped_addr: SocketAddr, + pub public_ips: Vec, + pub transaction_id: u32, + pub base_port_num: u32, + pub max_port_num: u32, + pub is_incremental: bool, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct SendPunchPacketBothEasySym { + pub udp_socket_count: u32, + pub public_ip: Ipv4Addr, + pub transaction_id: u32, + pub dst_port_num: u32, + pub wait_time_ms: u32, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct SendPunchPacketBothEasySymResponse { + pub is_busy: bool, + pub base_mapped_addr: Option, +} + +#[derive(Debug, thiserror::Error, PartialEq, Eq)] +pub enum UdpHolePunchSignalError { + #[error("invalid service key")] + InvalidServiceKey, + #[error("timeout")] + Timeout, + #[error("remote rejected: {0}")] + RemoteRejected(String), + #[error("transport: {0}")] + Transport(String), +} + +pub fn should_blacklist_signal_error(error: &UdpHolePunchSignalError) -> bool { + matches!(error, UdpHolePunchSignalError::InvalidServiceKey) +} + +#[async_trait] +pub trait UdpHolePunchSignaling: Send + Sync { + async fn select_punch_listener( + &self, + dst_peer_id: crate::config::PeerId, + request: SelectPunchListener, + ) -> Result; + + async fn send_punch_packet_cone( + &self, + dst_peer_id: crate::config::PeerId, + request: SendPunchPacketCone, + ) -> Result<(), UdpHolePunchSignalError>; + + async fn send_punch_packet_hard_sym( + &self, + dst_peer_id: crate::config::PeerId, + request: SendPunchPacketHardSym, + ) -> Result; + + async fn send_punch_packet_easy_sym( + &self, + dst_peer_id: crate::config::PeerId, + request: SendPunchPacketEasySym, + ) -> Result<(), UdpHolePunchSignalError>; + + async fn send_punch_packet_both_easy_sym( + &self, + dst_peer_id: crate::config::PeerId, + request: SendPunchPacketBothEasySym, + ) -> Result; +} + +#[async_trait] +pub trait UdpHolePunchInbound: Send + Sync { + async fn select_punch_listener( + &self, + request: SelectPunchListener, + ) -> Result; + + async fn send_punch_packet_cone( + &self, + request: SendPunchPacketCone, + ) -> Result<(), UdpHolePunchSignalError>; + + async fn send_punch_packet_hard_sym( + &self, + request: SendPunchPacketHardSym, + ) -> Result; + + async fn send_punch_packet_easy_sym( + &self, + request: SendPunchPacketEasySym, + ) -> Result<(), UdpHolePunchSignalError>; + + async fn send_punch_packet_both_easy_sym( + &self, + request: SendPunchPacketBothEasySym, + ) -> Result; +} + +#[async_trait] +pub trait UdpHolePunchTransportSink: Send + Sync { + async fn add_client_transport( + &self, + connected: ConnectedUdpSession, + requested_url: url::Url, + ) -> anyhow::Result<()>; + + async fn add_server_transport( + &self, + connected: ConnectedUdpSession, + requested_url: url::Url, + ) -> anyhow::Result<()>; +} + +pub struct ProtocolUdpHolePunchTransportSink { + protocol: Arc>, + tunnel_sink: Arc, +} + +impl ProtocolUdpHolePunchTransportSink { + pub fn new(protocol: Arc>, tunnel_sink: Arc) -> Self { + Self { + protocol, + tunnel_sink, + } + } + + async fn upgrade( + &self, + connected: ConnectedUdpSession, + requested_url: url::Url, + ) -> anyhow::Result> { + self.protocol + .upgrade_client(ConnectedTransport::Udp(connected), requested_url) + .await + } +} + +#[async_trait] +impl UdpHolePunchTransportSink for ProtocolUdpHolePunchTransportSink +where + TcpSocket: 'static, + T: HolePunchTunnelSink, +{ + async fn add_client_transport( + &self, + connected: ConnectedUdpSession, + requested_url: url::Url, + ) -> anyhow::Result<()> { + let tunnel = self.upgrade(connected, requested_url).await?; + self.tunnel_sink.add_client_tunnel(tunnel).await + } + + async fn add_server_transport( + &self, + connected: ConnectedUdpSession, + requested_url: url::Url, + ) -> anyhow::Result<()> { + let tunnel = self.upgrade(connected, requested_url).await?; + self.tunnel_sink.add_server_tunnel(tunnel).await + } +} + +#[async_trait] +pub trait UdpHolePunchPeerSource: Send + Sync { + fn local_peer_id(&self) -> crate::config::PeerId; + fn p2p_policy_flags(&self) -> P2pPolicyFlags; + + async fn candidates(&self) -> Vec; + + fn p2p_demand_notify(&self) -> Arc; + + fn is_local_virtual_ip(&self, ip: &IpAddr) -> bool; + + async fn is_easytier_managed_ipv6(&self, ip: &Ipv6Addr) -> bool; +} + +#[async_trait] +pub trait UdpHolePunchRuntime: Send + Sync + 'static { + type Socket: VirtualUdpSocket + 'static; + + fn socket_context(&self) -> SocketContext { + SocketContext::default() + } + + async fn bind_udp(&self, options: UdpBindOptions) -> anyhow::Result>; + + async fn bind_direct_connect_udp(&self) -> anyhow::Result> { + UdpHolePunchRuntime::bind_udp( + self, + UdpBindOptions::direct_connect().with_context( + self.socket_context() + .with_ip_version(crate::socket::IpVersion::V4), + ), + ) + .await + } + + async fn resolve_udp_public_addr( + &self, + socket: Arc, + ) -> anyhow::Result; + + async fn create_listener( + &self, + prefer_port_mapping: bool, + ) -> anyhow::Result>; + + async fn create_port_bound_listener( + &self, + port: u16, + ) -> anyhow::Result>; + + async fn connect_with_socket( + &self, + socket: Arc, + remote: SocketAddr, + ) -> anyhow::Result; +} + +#[async_trait] +impl VirtualUdpSocketFactory for T +where + T: UdpHolePunchRuntime + Send + Sync + 'static, +{ + type Socket = T::Socket; + + async fn bind_udp(&self, options: UdpBindOptions) -> anyhow::Result> { + UdpHolePunchRuntime::bind_udp(self, options).await + } +} + +#[cfg(test)] +mod tests { + use std::{ + io, + sync::{ + Arc, + atomic::{AtomicBool, AtomicUsize, Ordering}, + }, + }; + + use super::*; + use crate::socket::udp::UdpSessionKind; + + struct MockSocket { + local_addr: SocketAddr, + } + + #[async_trait] + impl VirtualUdpSocket for MockSocket { + fn local_addr(&self) -> io::Result { + Ok(self.local_addr) + } + + async fn send_to(&self, data: &[u8], _addr: SocketAddr) -> io::Result { + Ok(data.len()) + } + + async fn recv_from(&self, _buf: &mut [u8]) -> io::Result<(usize, SocketAddr)> { + std::future::pending().await + } + } + + struct DropSignal(Arc); + + impl Drop for DropSignal { + fn drop(&mut self) { + self.0.store(true, Ordering::Relaxed); + } + } + + #[derive(Default)] + struct MockProtocol { + upgrades: AtomicUsize, + } + + #[async_trait] + impl ClientProtocolUpgrader<()> for MockProtocol { + fn supports_scheme(&self, scheme: &str) -> bool { + scheme == "udp" + } + + async fn upgrade_client( + &self, + connected: ConnectedTransport<()>, + requested_url: url::Url, + ) -> anyhow::Result> { + self.upgrades.fetch_add(1, Ordering::Relaxed); + let ConnectedTransport::Udp(connected) = connected else { + anyhow::bail!("expected UDP transport"); + }; + Ok(crate::connectivity::protocol::raw::upgrade_connected_udp( + connected, + requested_url, + )?) + } + } + + #[derive(Default)] + struct MockTunnelSink { + clients: AtomicUsize, + servers: AtomicUsize, + } + + #[async_trait] + impl HolePunchTunnelSink for MockTunnelSink { + async fn add_client_tunnel(&self, _tunnel: Box) -> anyhow::Result<()> { + self.clients.fetch_add(1, Ordering::Relaxed); + Ok(()) + } + + async fn add_server_tunnel(&self, _tunnel: Box) -> anyhow::Result<()> { + self.servers.fetch_add(1, Ordering::Relaxed); + Ok(()) + } + } + + fn punched_socket(local_port: u16, remote_port: u16) -> UdpPunchSocket { + let remote_addr = SocketAddr::from(([203, 0, 113, 1], remote_port)); + let session = UdpSession::identity_standalone( + Arc::new(MockSocket { + local_addr: SocketAddr::from(([127, 0, 0, 1], local_port)), + }), + remote_addr, + UdpSessionKind::EasyTierMux, + ) + .unwrap(); + UdpPunchSocket::new(session, remote_addr, ()) + } + + #[tokio::test] + async fn protocol_sink_upgrades_before_role_specific_admission() { + let protocol = Arc::new(MockProtocol::default()); + let tunnel_sink = Arc::new(MockTunnelSink::default()); + let sink = + ProtocolUdpHolePunchTransportSink::<(), _>::new(protocol.clone(), tunnel_sink.clone()); + + let (client, client_url) = punched_socket(1000, 2000).into_connected(); + sink.add_client_transport(client, client_url).await.unwrap(); + let (server, server_url) = punched_socket(1001, 2001).into_connected(); + sink.add_server_transport(server, server_url).await.unwrap(); + + assert_eq!(protocol.upgrades.load(Ordering::Relaxed), 2); + assert_eq!(tunnel_sink.clients.load(Ordering::Relaxed), 1); + assert_eq!(tunnel_sink.servers.load(Ordering::Relaxed), 1); + } + + #[tokio::test] + async fn punched_socket_preserves_requested_and_resolved_addresses() { + let local_addr = SocketAddr::from(([127, 0, 0, 1], 1000)); + let requested_remote_addr = SocketAddr::from(([198, 51, 100, 1], 2000)); + let resolved_remote_addr = SocketAddr::from(([203, 0, 113, 1], 3000)); + let session = UdpSession::identity_standalone( + Arc::new(MockSocket { local_addr }), + resolved_remote_addr, + UdpSessionKind::EasyTierMux, + ) + .unwrap(); + let guard_dropped = Arc::new(AtomicBool::new(false)); + let socket = UdpPunchSocket::new( + session, + requested_remote_addr, + DropSignal(guard_dropped.clone()), + ); + + let (connected, requested_url) = socket.into_connected(); + let tunnel = + crate::connectivity::protocol::raw::upgrade_connected_udp(connected, requested_url) + .unwrap(); + let info = tunnel.info().unwrap(); + let local_url: url::Url = info.local_addr.unwrap().into(); + let remote_url: url::Url = info.remote_addr.unwrap().into(); + let resolved_url: url::Url = info.resolved_remote_addr.unwrap().into(); + + assert_eq!(local_url.host_str(), Some("127.0.0.1")); + assert_eq!(local_url.port(), Some(local_addr.port())); + assert_eq!(remote_url.host_str(), Some("198.51.100.1")); + assert_eq!(remote_url.port(), Some(requested_remote_addr.port())); + assert_eq!(resolved_url.host_str(), Some("203.0.113.1")); + assert_eq!(resolved_url.port(), Some(resolved_remote_addr.port())); + assert!(!guard_dropped.load(Ordering::Relaxed)); + drop(tunnel); + assert!(guard_dropped.load(Ordering::Relaxed)); + } +} diff --git a/easytier-core/src/connectivity/hole_punch/udp/server.rs b/easytier-core/src/connectivity/hole_punch/udp/server.rs new file mode 100644 index 00000000..5da74fb8 --- /dev/null +++ b/easytier-core/src/connectivity/hole_punch/udp/server.rs @@ -0,0 +1,1386 @@ +use std::{ + net::{Ipv4Addr, SocketAddr, SocketAddrV4}, + sync::{ + Arc, + atomic::{AtomicBool, Ordering}, + }, + time::Duration, +}; + +use anyhow::Context; +use crossbeam::atomic::AtomicCell; +use quanta::Instant; +use rand::{Rng, seq::SliceRandom as _}; +use tokio::{ + sync::{Mutex, RwLock, RwLockReadGuard}, + task::JoinSet, +}; +use tokio_util::task::AbortOnDropHandle; + +use crate::{ + connectivity::{hole_punch::port_mapping::UdpPortMappingLease, stun::StunInfoProvider}, + packet::{HOLE_PUNCH_PACKET_BODY_LEN, new_hole_punch_packet}, + socket::{ListenerConnectionCounter, udp::VirtualUdpSocket}, +}; + +use super::{ + MAX_PUBLIC_UDP_HOLE_PUNCH_LISTENERS, ReusableUdpPunchListener, SelectPunchListener, + SelectPunchListenerResponse, SendPunchPacketBothEasySym, SendPunchPacketBothEasySymResponse, + SendPunchPacketCone, SendPunchPacketEasySym, SendPunchPacketHardSym, + SendPunchPacketHardSymResponse, UdpHolePunchInbound, UdpHolePunchRuntime, + UdpHolePunchSignalError, UdpHolePunchTransportSink, UdpPunchListener, UdpSocketArray, + UdpSymPunchLock, can_reuse_port_mapping_listener, can_reuse_public_listener, + select_reusable_port_mapping_listener_idx, select_reusable_public_listener_idx, + should_create_public_listener, should_retry_public_listener_selection, +}; + +const MAX_K1_FOR_RANDOM_HARD_SYM: u32 = 180; + +pub struct SelectedUdpPunchListener { + /// Kept for listener-selection side effects and tests; production callers + /// only consume `mapped_addr`. + #[allow(dead_code)] + pub socket: Arc, + pub mapped_addr: SocketAddr, +} + +pub struct UdpHolePunchServer +where + R: UdpHolePunchRuntime, + T: UdpHolePunchTransportSink + 'static, +{ + sym_punch_lock: UdpSymPunchLock, + common: Arc>, + both_easy_sym_server: UdpBothEasySymPunchServer, + shuffled_port_vec: Arc>, + admission: RwLock<()>, + stopping: AtomicBool, +} + +impl UdpHolePunchServer +where + R: UdpHolePunchRuntime, + T: UdpHolePunchTransportSink + 'static, +{ + pub fn new( + runtime: Arc, + stun: Arc, + transport_sink: Arc, + sym_punch_lock: UdpSymPunchLock, + ) -> Self { + let common = Arc::new(UdpHolePunchServerCommon::new( + runtime.clone(), + stun.clone(), + transport_sink.clone(), + )); + let both_easy_sym_common = + Arc::new(UdpHolePunchServerCommon::new(runtime, stun, transport_sink)); + let both_easy_sym_server = UdpBothEasySymPunchServer::new(both_easy_sym_common); + let mut shuffled_port_vec: Vec = (1..=65535).collect(); + shuffled_port_vec.shuffle(&mut rand::thread_rng()); + + Self { + sym_punch_lock, + common, + both_easy_sym_server, + shuffled_port_vec: Arc::new(shuffled_port_vec), + admission: RwLock::new(()), + stopping: AtomicBool::new(true), + } + } + + pub async fn start(&self) { + let _admission = self.admission.write().await; + if !self.stopping.load(Ordering::Acquire) { + return; + } + self.stop_inner().await; + self.common.start().await; + self.both_easy_sym_server.common.start().await; + self.stopping.store(false, Ordering::Release); + } + + pub fn begin_stop(&self) { + self.stopping.store(true, Ordering::Release); + } + + pub async fn stop(&self) { + self.begin_stop(); + let _admission = self.admission.write().await; + self.stopping.store(true, Ordering::Release); + self.stop_inner().await; + } + + async fn stop_inner(&self) { + self.both_easy_sym_server.stop().await; + self.both_easy_sym_server.common.stop().await; + self.common.stop().await; + } + + async fn admit(&self) -> Result, UdpHolePunchSignalError> { + if self.stopping.load(Ordering::Acquire) { + return Err(UdpHolePunchSignalError::Transport( + "udp hole punch server is stopping".into(), + )); + } + let guard = self.admission.read().await; + if self.stopping.load(Ordering::Acquire) { + return Err(UdpHolePunchSignalError::Transport( + "udp hole punch server is stopping".into(), + )); + } + Ok(guard) + } + + fn busy_signal_error() -> UdpHolePunchSignalError { + UdpHolePunchSignalError::RemoteRejected("sym punch lock is busy".into()) + } + + fn anyhow_to_signal_error(error: anyhow::Error) -> UdpHolePunchSignalError { + UdpHolePunchSignalError::RemoteRejected(error.to_string()) + } + + async fn send_punch_packet_easy_sym_inner( + &self, + request: SendPunchPacketEasySym, + ) -> anyhow::Result<()> { + tracing::info!("send_punch_packet_easy_sym start"); + + let listener = self + .common + .find_listener(&request.listener_mapped_addr) + .await + .ok_or(anyhow::anyhow!( + "send_punch_packet_easy_sym failed to find listener" + ))?; + + if request.public_ips.is_empty() { + tracing::warn!("send_punch_packet_easy_sym got zero len public ip"); + anyhow::bail!("send_punch_packet_easy_sym got zero len public ip"); + } + + let base_port_num = request.base_port_num; + let max_port_num = request.max_port_num.max(1); + let port_start = if request.is_incremental { + base_port_num.saturating_add(1) + } else { + base_port_num.saturating_sub(max_port_num) + }; + let port_end = if request.is_incremental { + base_port_num.saturating_add(max_port_num) + } else { + base_port_num.saturating_sub(1) + }; + + if port_end <= port_start { + anyhow::bail!("send_punch_packet_easy_sym invalid port range"); + } + + let ports = (port_start..=port_end) + .map(|port| port as u16) + .collect::>(); + tracing::debug!( + ?ports, + public_ips = ?request.public_ips, + "send_punch_packet_easy_sym send to ports" + ); + + for _ in 0..2 { + send_symmetric_hole_punch_packet( + &ports, + listener.clone(), + request.transaction_id, + &request.public_ips, + 0, + ports.len(), + ) + .await + .with_context(|| "failed to send symmetric hole punch packet")?; + } + + Ok(()) + } + + async fn send_punch_packet_hard_sym_inner( + &self, + request: SendPunchPacketHardSym, + ) -> anyhow::Result { + tracing::info!("try_punch_symmetric start"); + + let listener = self + .common + .find_listener(&request.listener_mapped_addr) + .await + .ok_or(anyhow::anyhow!( + "send_punch_packet_for_cone failed to find listener" + ))?; + + if request.public_ips.is_empty() { + tracing::warn!("try_punch_symmetric got zero len public ip"); + anyhow::bail!("try_punch_symmetric got zero len public ip"); + } + + let last_port_index = request.port_index as usize; + let round = request.round.max(1); + let mut max_k2: u32 = rand::thread_rng().gen_range(600..800); + if round > 2 { + max_k2 = (max_k2 * 2 / round).max(MAX_K1_FOR_RANDOM_HARD_SYM); + } + + let mut next_port_index = 0; + for _ in 0..2 { + next_port_index = send_symmetric_hole_punch_packet( + &self.shuffled_port_vec, + listener.clone(), + request.transaction_id, + &request.public_ips, + last_port_index, + max_k2 as usize, + ) + .await + .with_context(|| "failed to send symmetric hole punch packet randomly")?; + } + + Ok(SendPunchPacketHardSymResponse { + next_port_index: next_port_index as u32, + }) + } +} + +#[async_trait::async_trait] +impl UdpHolePunchInbound for UdpHolePunchServer +where + R: UdpHolePunchRuntime, + T: UdpHolePunchTransportSink + 'static, +{ + async fn select_punch_listener( + &self, + request: SelectPunchListener, + ) -> Result { + let _admission = self.admit().await?; + let selected = self + .common + .select_listener(request.force_new, request.prefer_port_mapping) + .await + .ok_or_else(|| { + UdpHolePunchSignalError::RemoteRejected("no listener available".into()) + })?; + + Ok(SelectPunchListenerResponse { + listener_mapped_addr: selected.mapped_addr, + }) + } + + async fn send_punch_packet_cone( + &self, + request: SendPunchPacketCone, + ) -> Result<(), UdpHolePunchSignalError> { + let _admission = self.admit().await?; + let listener = self + .common + .find_listener(&request.listener_mapped_addr) + .await + .ok_or_else(|| { + UdpHolePunchSignalError::RemoteRejected( + "send_punch_packet_for_cone failed to find listener".into(), + ) + })?; + + send_cone_hole_punch_packets(listener, &request) + .await + .map_err(Self::anyhow_to_signal_error) + } + + async fn send_punch_packet_hard_sym( + &self, + request: SendPunchPacketHardSym, + ) -> Result { + let _admission = self.admit().await?; + let _locked = self + .sym_punch_lock + .try_lock() + .map_err(|_| Self::busy_signal_error())?; + self.send_punch_packet_hard_sym_inner(request) + .await + .map_err(Self::anyhow_to_signal_error) + } + + async fn send_punch_packet_easy_sym( + &self, + request: SendPunchPacketEasySym, + ) -> Result<(), UdpHolePunchSignalError> { + let _admission = self.admit().await?; + let _locked = self + .sym_punch_lock + .try_lock() + .map_err(|_| Self::busy_signal_error())?; + self.send_punch_packet_easy_sym_inner(request) + .await + .map_err(Self::anyhow_to_signal_error) + } + + async fn send_punch_packet_both_easy_sym( + &self, + request: SendPunchPacketBothEasySym, + ) -> Result { + let _admission = self.admit().await?; + let _locked = self + .sym_punch_lock + .try_lock() + .map_err(|_| Self::busy_signal_error())?; + self.both_easy_sym_server + .send_punch_packet_both_easy_sym(request) + .await + .map_err(Self::anyhow_to_signal_error) + } +} + +type UdpPunchListenerRecords = Arc>>>>; + +pub struct UdpHolePunchServerCommon +where + R: UdpHolePunchRuntime, + T: UdpHolePunchTransportSink + 'static, +{ + runtime: Arc, + stun: Arc, + transport_sink: Arc, + listeners: UdpPunchListenerRecords, + pending: UdpPunchListenerRecords, + retiring: UdpPunchListenerRecords, + cleanup_task: Mutex>>, +} + +impl UdpHolePunchServerCommon +where + R: UdpHolePunchRuntime, + T: UdpHolePunchTransportSink + 'static, +{ + pub fn new(runtime: Arc, stun: Arc, transport_sink: Arc) -> Self { + let listeners = Arc::new(Mutex::new(Vec::new())); + + Self { + runtime, + stun, + transport_sink, + listeners, + pending: Arc::new(Mutex::new(Vec::new())), + retiring: Arc::new(Mutex::new(Vec::new())), + cleanup_task: Mutex::new(None), + } + } + + pub async fn start(&self) { + let mut task_slot = self.cleanup_task.lock().await; + if task_slot.as_ref().is_some_and(|task| !task.is_finished()) { + return; + } + let listeners = self.listeners.clone(); + let retiring = self.retiring.clone(); + task_slot.replace(AbortOnDropHandle::new(tokio::spawn(async move { + loop { + crate::foundation::time::sleep(Duration::from_secs(5)).await; + { + let mut retiring = retiring.lock().await; + let mut listeners = listeners.lock().await; + let mut index = 0; + while index < listeners.len() { + let listener = &listeners[index]; + let active = listener.last_active_time.load().elapsed().as_secs() < 40 + || listener.last_select_time.load().elapsed().as_secs() < 30; + if active { + index += 1; + } else { + retiring.push(listeners.remove(index)); + } + } + } + drain_retiring_listeners(retiring.as_ref()).await; + } + }))); + } + + pub async fn stop(&self) { + let mut cleanup_task = self.cleanup_task.lock().await; + if let Some(cleanup_task) = cleanup_task.as_mut() { + cleanup_task.abort(); + let _ = cleanup_task.await; + } + cleanup_task.take(); + drop(cleanup_task); + + { + let mut pending = self.pending.lock().await; + let mut retiring = self.retiring.lock().await; + let mut listeners = self.listeners.lock().await; + retiring.extend(std::mem::take(&mut *pending)); + retiring.extend(std::mem::take(&mut *listeners)); + } + drain_retiring_listeners(self.retiring.as_ref()).await; + } + + pub async fn add_listener(&self, listener: UdpPunchListener) { + let mut listeners = self.listeners.lock().await; + listeners.push(Arc::new(UdpPunchListenerRecord::new( + listener, + self.transport_sink.clone(), + ))); + } + + async fn track_pending_listener( + &self, + listener: UdpPunchListener, + ) -> Arc> { + let mut pending = self.pending.lock().await; + let listener = Arc::new(UdpPunchListenerRecord::new( + listener, + self.transport_sink.clone(), + )); + pending.push(listener.clone()); + listener + } + + async fn promote_pending_listener(&self, listener: Arc>) { + let mut pending = self.pending.lock().await; + let mut listeners = self.listeners.lock().await; + if let Some(index) = pending + .iter() + .position(|candidate| Arc::ptr_eq(candidate, &listener)) + { + pending.remove(index); + listeners.push(listener); + } + } + + async fn retire_pending_listener(&self, listener: Arc>) { + let mut pending = self.pending.lock().await; + let mut retiring = self.retiring.lock().await; + if let Some(index) = pending + .iter() + .position(|candidate| Arc::ptr_eq(candidate, &listener)) + { + retiring.push(pending.remove(index)); + } + } + + pub async fn find_listener(&self, addr: &SocketAddr) -> Option> { + let listeners = self.listeners.lock().await; + + let listener = listeners + .iter() + .find(|listener| listener.mapped_addr == *addr && listener.running.load())?; + + Some(listener.get_socket()) + } + + pub async fn select_listener( + &self, + force_new_listener: bool, + prefer_port_mapping: bool, + ) -> Option> { + let mut force_new_listener = force_new_listener; + + loop { + let (listener_count, has_reusable_listener, has_port_mapping_listener) = { + let listeners = self.listeners.lock().await; + let states = listener_reuse_states(listeners.as_slice()); + ( + states.len(), + states.iter().any(can_reuse_public_listener), + states.iter().any(can_reuse_port_mapping_listener), + ) + }; + let should_create = should_create_public_listener( + listener_count, + has_reusable_listener, + has_port_mapping_listener, + force_new_listener, + prefer_port_mapping, + ); + + if should_create { + tracing::warn!( + max_listeners = MAX_PUBLIC_UDP_HOLE_PUNCH_LISTENERS, + "creating udp hole punching listener" + ); + match self.runtime.create_listener(prefer_port_mapping).await { + Ok(listener) => self.add_listener(listener).await, + Err(err) => { + tracing::warn!(?err, "failed to create udp hole punching listener"); + } + } + } + + let mut listeners = self.listeners.lock().await; + let listener_count = listeners.len(); + let states = listener_reuse_states(listeners.as_slice()); + let listener_idx = if prefer_port_mapping { + select_reusable_port_mapping_listener_idx(&states) + .or_else(|| { + if should_create && states.last().is_some_and(can_reuse_public_listener) { + Some(states.len() - 1) + } else { + None + } + }) + .or_else(|| select_reusable_public_listener_idx(&states)) + } else if should_create { + listeners.len().checked_sub(1) + } else { + select_reusable_public_listener_idx(&states) + }; + + let Some(listener_idx) = listener_idx else { + tracing::warn!( + ?force_new_listener, + ?prefer_port_mapping, + listener_count, + max_listeners = MAX_PUBLIC_UDP_HOLE_PUNCH_LISTENERS, + "no available udp hole punching listener with mapped address" + ); + if should_retry_public_listener_selection( + force_new_listener, + listener_count, + prefer_port_mapping, + has_port_mapping_listener, + ) { + force_new_listener = true; + continue; + } + return None; + }; + + let listener = &mut listeners[listener_idx]; + if !can_reuse_public_listener(&listener.reuse_state()) { + tracing::warn!( + ?force_new_listener, + ?prefer_port_mapping, + listener_count, + max_listeners = MAX_PUBLIC_UDP_HOLE_PUNCH_LISTENERS, + "selected udp hole punching listener is not reusable" + ); + return None; + } + + return Some(SelectedUdpPunchListener { + socket: listener.get_socket(), + mapped_addr: listener.mapped_addr, + }); + } + } +} + +struct UdpPunchListenerRecord { + socket: Arc, + tasks: Mutex>, + connection_tasks: Arc>>, + running: Arc>, + mapped_addr: SocketAddr, + has_port_mapping_lease: bool, + _port_mapping_lease: Option>, + conn_counter: Arc, + + _listen_time: Instant, + last_select_time: AtomicCell, + last_active_time: Arc>, +} + +impl UdpPunchListenerRecord +where + S: VirtualUdpSocket + 'static, +{ + fn new(listener: UdpPunchListener, transport_sink: Arc) -> Self + where + T: UdpHolePunchTransportSink + 'static, + { + let UdpPunchListener { + socket, + mapped_addr, + conn_counter, + mut acceptor, + port_mapping_lease, + } = listener; + + let running = Arc::new(AtomicCell::new(true)); + let running_clone = running.clone(); + let mut tasks = JoinSet::new(); + let connection_tasks = Arc::new(Mutex::new(JoinSet::new())); + let accept_connection_tasks = connection_tasks.clone(); + + tasks.spawn(async move { + while let Ok(socket) = acceptor.accept().await { + tracing::warn!(?socket, "udp hole punching listener got peer connection"); + let (connected, requested_url) = socket.into_connected(); + let transport_sink = transport_sink.clone(); + let mut connection_tasks = accept_connection_tasks.lock().await; + while connection_tasks.try_join_next().is_some() {} + connection_tasks.spawn(async move { + if let Err(err) = transport_sink + .add_server_transport(connected, requested_url) + .await + { + tracing::error!( + ?err, + "failed to upgrade or add server UDP hole-punch transport" + ); + } + }); + } + + running_clone.store(false); + }); + + let last_active_time = Arc::new(AtomicCell::new(Instant::now())); + let conn_counter_clone = conn_counter.clone(); + let last_active_time_clone = last_active_time.clone(); + tasks.spawn(async move { + loop { + crate::foundation::time::sleep(Duration::from_secs(5)).await; + if conn_counter_clone.get().unwrap_or(0) != 0 { + last_active_time_clone.store(Instant::now()); + } + } + }); + + tracing::warn!(?mapped_addr, "udp hole punching listener started"); + + Self { + socket, + tasks: Mutex::new(tasks), + connection_tasks, + running, + mapped_addr, + has_port_mapping_lease: port_mapping_lease.is_some(), + _port_mapping_lease: port_mapping_lease, + conn_counter, + + _listen_time: Instant::now(), + last_select_time: AtomicCell::new(Instant::now()), + last_active_time, + } + } + + fn get_socket(&self) -> Arc { + self.last_select_time.store(Instant::now()); + self.socket.clone() + } + + fn conn_count(&self) -> usize { + self.conn_counter.get().unwrap_or(0) as usize + } + + fn reuse_state(&self) -> ReusableUdpPunchListener { + ReusableUdpPunchListener { + running: self.running.load(), + mapped_addr: self.mapped_addr, + has_port_mapping_lease: self.has_port_mapping_lease, + last_active_time: self.last_active_time.load(), + } + } + + async fn stop(&self) { + let mut tasks = self.tasks.lock().await; + tasks.abort_all(); + while tasks.join_next().await.is_some() {} + drop(tasks); + + let mut connection_tasks = self.connection_tasks.lock().await; + connection_tasks.abort_all(); + while connection_tasks.join_next().await.is_some() {} + self.running.store(false); + } +} + +async fn drain_retiring_listeners(retiring: &Mutex>>>) +where + S: VirtualUdpSocket + 'static, +{ + loop { + let listener = retiring.lock().await.first().cloned(); + let Some(listener) = listener else { + return; + }; + listener.stop().await; + let mut retiring = retiring.lock().await; + if let Some(index) = retiring + .iter() + .position(|candidate| Arc::ptr_eq(candidate, &listener)) + { + retiring.remove(index); + } + } +} + +fn listener_reuse_states( + listeners: &[Arc>], +) -> Vec +where + S: VirtualUdpSocket + 'static, +{ + listeners + .iter() + .map(|listener| listener.reuse_state()) + .collect() +} + +pub async fn send_cone_hole_punch_packets( + udp: Arc, + request: &SendPunchPacketCone, +) -> anyhow::Result<()> +where + S: VirtualUdpSocket + 'static, +{ + let dest_ip = request.dest_addr.ip(); + if dest_ip.is_unspecified() || dest_ip.is_multicast() { + anyhow::bail!( + "send_punch_packet_for_cone dest_ip is malformed: {:?}", + request + ); + } + + for _ in 0..request.packet_batch_count { + tracing::info!(?request, "sending hole punching packet"); + + for _ in 0..request.packet_count_per_batch { + let udp_packet = + new_hole_punch_packet(request.transaction_id, HOLE_PUNCH_PACKET_BODY_LEN); + if let Err(err) = udp + .send_to(&udp_packet.into_bytes(), request.dest_addr) + .await + { + tracing::error!(?err, "failed to send hole punch packet to dest addr"); + } + } + crate::foundation::time::sleep(Duration::from_millis(request.packet_interval_ms as u64)) + .await; + } + + Ok(()) +} + +#[tracing::instrument(err, ret(level = tracing::Level::DEBUG), skip(ports, udp))] +pub async fn send_symmetric_hole_punch_packet( + ports: &[u16], + udp: Arc, + transaction_id: u32, + public_ips: &[Ipv4Addr], + port_start_idx: usize, + max_packets: usize, +) -> anyhow::Result +where + S: VirtualUdpSocket + 'static, +{ + tracing::debug!("sending hard symmetric hole punching packet"); + let mut sent_packets = 0; + let mut cur_port_idx = port_start_idx; + while sent_packets < max_packets { + let port = ports[cur_port_idx % ports.len()]; + for pub_ip in public_ips { + let addr = SocketAddr::V4(SocketAddrV4::new(*pub_ip, port)); + for _ in 0..3 { + let packet = new_hole_punch_packet(transaction_id, HOLE_PUNCH_PACKET_BODY_LEN); + udp.send_to(&packet.into_bytes(), addr).await?; + } + sent_packets += 1; + } + cur_port_idx = cur_port_idx.wrapping_add(1); + crate::foundation::time::sleep(Duration::from_millis(1)).await; + } + Ok(cur_port_idx % ports.len()) +} + +pub struct UdpBothEasySymPunchServer +where + R: UdpHolePunchRuntime, + T: UdpHolePunchTransportSink + 'static, +{ + common: Arc>, + task: Mutex>>, +} + +impl UdpBothEasySymPunchServer +where + R: UdpHolePunchRuntime, + T: UdpHolePunchTransportSink + 'static, +{ + pub fn new(common: Arc>) -> Self { + Self { + common, + task: Mutex::new(None), + } + } + + pub async fn stop(&self) { + let mut task = self.task.lock().await; + if let Some(task) = task.as_mut() { + task.abort(); + let _ = task.await; + } + task.take(); + } + + #[tracing::instrument(skip(self), ret, err)] + pub async fn send_punch_packet_both_easy_sym( + &self, + request: SendPunchPacketBothEasySym, + ) -> anyhow::Result { + tracing::info!("send_punch_packet_both_easy_sym start"); + let busy_resp = Ok(SendPunchPacketBothEasySymResponse { + is_busy: true, + base_mapped_addr: None, + }); + let Ok(mut locked_task) = self.task.try_lock() else { + return busy_resp; + }; + if locked_task.is_some() && !locked_task.as_ref().unwrap().is_finished() { + return busy_resp; + } + + let cur_mapped_addr = self + .common + .stun + .get_udp_port_mapping(0) + .await + .with_context(|| "failed to get udp port mapping")?; + + tracing::info!("send_punch_packet_hard_sym start"); + let socket_count = request.udp_socket_count as usize; + let transaction_id = request.transaction_id; + + let udp_array = UdpSocketArray::new_with_context( + socket_count, + self.common.runtime.clone(), + self.common.runtime.socket_context(), + ); + udp_array.start().await?; + udp_array.add_intreast_tid(transaction_id); + + let punch_packet = + new_hole_punch_packet(transaction_id, HOLE_PUNCH_PACKET_BODY_LEN).into_bytes(); + let common = self.common.clone(); + + let task = tokio::spawn(async move { + let mut listeners = Vec::new(); + let mut punched = Vec::new(); + let start_time = Instant::now(); + let wait_time_ms = request.wait_time_ms.min(8000); + while start_time.elapsed() < Duration::from_millis(wait_time_ms as u64) { + if let Err(e) = udp_array + .send_with_all( + &punch_packet, + SocketAddr::V4(SocketAddrV4::new( + request.public_ip, + request.dst_port_num as u16, + )), + ) + .await + { + tracing::error!(?e, "failed to send hole punch packet"); + break; + } + + crate::foundation::time::sleep(Duration::from_millis(100)).await; + + if let Some(s) = udp_array.try_fetch_punched_socket(transaction_id) { + tracing::info!(?s, ?transaction_id, "got punched socket in both easy sym"); + assert!(Arc::strong_count(&s.socket) == 1); + let Some(port) = s.socket.local_addr().ok().map(|addr| addr.port()) else { + tracing::warn!("failed to get local addr from punched socket"); + continue; + }; + let remote_addr = s.remote_addr; + drop(s); + + let listener = match common.runtime.create_port_bound_listener(port).await { + Ok(listener) => listener, + Err(e) => { + tracing::warn!(?e, "failed to create listener"); + continue; + } + }; + let socket = listener.socket.clone(); + let record = common.track_pending_listener(listener).await; + punched.push((socket, remote_addr)); + listeners.push(record); + } + + for listener in &listeners { + if listener.conn_count() > 0 { + tracing::info!(?listener.mapped_addr, "got punched listener"); + break; + } + } + + if !punched.is_empty() { + tracing::debug!( + punched_count = punched.len(), + "got punched socket and keep sending punch packet" + ); + } + + for p in &punched { + let (socket, remote_addr) = p; + let send_remote_ret = socket.send_to(&punch_packet, *remote_addr).await; + tracing::debug!( + ?send_remote_ret, + ?remote_addr, + "send hole punch packet to punched remote" + ); + } + } + + for listener in listeners { + if listener.conn_count() > 0 { + common.promote_pending_listener(listener).await; + } else { + common.retire_pending_listener(listener).await; + } + } + drain_retiring_listeners(common.retiring.as_ref()).await; + }); + + *locked_task = Some(AbortOnDropHandle::new(task)); + Ok(SendPunchPacketBothEasySymResponse { + is_busy: false, + base_mapped_addr: Some(cur_mapped_addr), + }) + } +} + +#[cfg(test)] +mod tests { + use std::{ + collections::VecDeque, + io, + sync::{ + Mutex as StdMutex, + atomic::{AtomicUsize, Ordering}, + }, + }; + + use async_trait::async_trait; + + use super::*; + use crate::{ + connectivity::hole_punch::udp::{UdpPunchAcceptor, UdpPunchSocket, UdpResolvedPublicAddr}, + proto::common::StunInfo, + socket::udp::{UdpBindOptions, UdpSession, UdpSessionKind}, + }; + + struct MockSocket { + local_addr: SocketAddr, + sent: tokio::sync::Mutex, SocketAddr)>>, + fail_next_send: AtomicUsize, + } + + #[async_trait] + impl VirtualUdpSocket for MockSocket { + fn local_addr(&self) -> io::Result { + Ok(self.local_addr) + } + + async fn send_to(&self, data: &[u8], _addr: SocketAddr) -> io::Result { + if self.fail_next_send.load(Ordering::Relaxed) != 0 { + self.fail_next_send.fetch_sub(1, Ordering::Relaxed); + return Err(io::Error::other("mock send failure")); + } + self.sent.lock().await.push((data.to_vec(), _addr)); + Ok(data.len()) + } + + async fn recv_from(&self, _buf: &mut [u8]) -> io::Result<(usize, SocketAddr)> { + std::future::pending().await + } + } + + #[derive(Debug, Default)] + struct MockCounter { + count: AtomicCell, + } + + impl ListenerConnectionCounter for MockCounter { + fn get(&self) -> Option { + Some(self.count.load()) + } + } + + struct MockAcceptor { + sockets: VecDeque, + } + + #[async_trait] + impl UdpPunchAcceptor for MockAcceptor { + async fn accept(&mut self) -> anyhow::Result { + let Some(socket) = self.sockets.pop_front() else { + return std::future::pending().await; + }; + Ok(socket) + } + } + + struct MockRuntime { + listeners: StdMutex>>, + } + + impl MockRuntime { + fn new(listeners: Vec>) -> Self { + Self { + listeners: StdMutex::new(listeners.into()), + } + } + } + + #[async_trait] + impl UdpHolePunchRuntime for MockRuntime { + type Socket = MockSocket; + + async fn bind_udp(&self, _options: UdpBindOptions) -> anyhow::Result> { + Ok(Arc::new(MockSocket { + local_addr: SocketAddr::from(([127, 0, 0, 1], 0)), + sent: tokio::sync::Mutex::new(Vec::new()), + fail_next_send: AtomicUsize::new(0), + })) + } + + async fn resolve_udp_public_addr( + &self, + socket: Arc, + ) -> anyhow::Result { + Ok(UdpResolvedPublicAddr { + mapped_addr: socket.local_addr()?, + port_mapping_lease: None, + }) + } + + async fn create_listener( + &self, + _prefer_port_mapping: bool, + ) -> anyhow::Result> { + self.listeners + .lock() + .unwrap() + .pop_front() + .ok_or_else(|| anyhow::anyhow!("no listener")) + } + + async fn create_port_bound_listener( + &self, + _port: u16, + ) -> anyhow::Result> { + self.create_listener(false).await + } + + async fn connect_with_socket( + &self, + socket: Arc, + remote: SocketAddr, + ) -> anyhow::Result { + let session = + UdpSession::identity_standalone(socket, remote, UdpSessionKind::EasyTierMux)?; + Ok(UdpPunchSocket::new(session, remote, ())) + } + } + + struct MockStunInfoProvider; + + #[async_trait] + impl StunInfoProvider for MockStunInfoProvider { + fn get_stun_info(&self) -> StunInfo { + StunInfo::default() + } + + async fn get_udp_port_mapping(&self, _port: u16) -> anyhow::Result { + Ok(SocketAddr::from(([203, 0, 113, 1], 10000))) + } + + async fn get_tcp_port_mapping(&self, _port: u16) -> anyhow::Result { + unreachable!("TCP mapping is not used by UDP hole-punch tests") + } + + fn update_stun_info(&self) {} + } + + fn mock_stun() -> Arc { + Arc::new(MockStunInfoProvider) + } + + #[derive(Default)] + struct MockSink { + server_tunnels: AtomicUsize, + } + + #[async_trait] + impl UdpHolePunchTransportSink for MockSink { + async fn add_client_transport( + &self, + _connected: crate::connectivity::transport::ConnectedUdpSession, + _requested_url: url::Url, + ) -> anyhow::Result<()> { + Ok(()) + } + + async fn add_server_transport( + &self, + _connected: crate::connectivity::transport::ConnectedUdpSession, + _requested_url: url::Url, + ) -> anyhow::Result<()> { + self.server_tunnels.fetch_add(1, Ordering::Relaxed); + Ok(()) + } + } + + fn listener(port: u16, sockets: Vec) -> UdpPunchListener { + UdpPunchListener { + socket: Arc::new(MockSocket { + local_addr: SocketAddr::from(([127, 0, 0, 1], port)), + sent: tokio::sync::Mutex::new(Vec::new()), + fail_next_send: AtomicUsize::new(0), + }), + mapped_addr: SocketAddr::from(([203, 0, 113, 1], port)), + conn_counter: Arc::new(MockCounter::default()), + acceptor: Box::new(MockAcceptor { + sockets: sockets.into(), + }), + port_mapping_lease: None, + } + } + + #[tokio::test] + async fn server_keeps_both_easy_sym_listener_pool_separate() { + let runtime = Arc::new(MockRuntime::new(Vec::new())); + let sink = Arc::new(MockSink::default()); + let server = + UdpHolePunchServer::new(runtime, mock_stun(), sink, UdpSymPunchLock::default()); + + assert!(!Arc::ptr_eq( + &server.common, + &server.both_easy_sym_server.common + )); + } + + #[test] + fn server_constructor_is_cold() { + let runtime = Arc::new(MockRuntime::new(Vec::new())); + let sink = Arc::new(MockSink::default()); + + let _server = + UdpHolePunchServer::new(runtime, mock_stun(), sink, UdpSymPunchLock::default()); + } + + #[tokio::test] + async fn common_lifecycle_joins_cleanup_and_listener_tasks() { + let runtime = Arc::new(MockRuntime::new(vec![listener(10003, Vec::new())])); + let sink = Arc::new(MockSink::default()); + let common = Arc::new(UdpHolePunchServerCommon::new(runtime, mock_stun(), sink)); + + common.start().await; + common.select_listener(false, false).await.unwrap(); + common.stop().await; + + assert!(common.cleanup_task.lock().await.is_none()); + assert!(common.listeners.lock().await.is_empty()); + } + + #[tokio::test] + async fn select_listener_creates_and_finds_listener() { + let runtime = Arc::new(MockRuntime::new(vec![listener(10000, Vec::new())])); + let sink = Arc::new(MockSink::default()); + let common = UdpHolePunchServerCommon::new(runtime, mock_stun(), sink); + + let selected = common.select_listener(false, true).await.unwrap(); + + assert_eq!( + selected.mapped_addr, + SocketAddr::from(([203, 0, 113, 1], 10000)) + ); + assert_eq!( + selected.socket.local_addr().unwrap(), + SocketAddr::from(([127, 0, 0, 1], 10000)) + ); + assert!(common.find_listener(&selected.mapped_addr).await.is_some()); + } + + #[tokio::test] + async fn accepted_tunnel_is_forwarded_to_sink() { + let socket = Arc::new(MockSocket { + local_addr: SocketAddr::from(([127, 0, 0, 1], 10001)), + sent: tokio::sync::Mutex::new(Vec::new()), + fail_next_send: AtomicUsize::new(0), + }); + let remote_addr = SocketAddr::from(([127, 0, 0, 1], 20001)); + let session = + UdpSession::identity_standalone(socket, remote_addr, UdpSessionKind::EasyTierMux) + .unwrap(); + let punched_socket = UdpPunchSocket::new(session, remote_addr, ()); + let runtime = Arc::new(MockRuntime::new(vec![listener( + 10001, + vec![punched_socket], + )])); + let sink = Arc::new(MockSink::default()); + let common = UdpHolePunchServerCommon::new(runtime, mock_stun(), sink.clone()); + + common.select_listener(false, false).await.unwrap(); + + for _ in 0..10 { + if sink.server_tunnels.load(Ordering::Relaxed) == 1 { + return; + } + tokio::task::yield_now().await; + } + + assert_eq!(sink.server_tunnels.load(Ordering::Relaxed), 1); + } + + #[tokio::test] + async fn cone_packet_sender_keeps_old_batch_shape() { + let socket = Arc::new(MockSocket { + local_addr: SocketAddr::from(([127, 0, 0, 1], 10002)), + sent: tokio::sync::Mutex::new(Vec::new()), + fail_next_send: AtomicUsize::new(0), + }); + let request = SendPunchPacketCone { + listener_mapped_addr: SocketAddr::from(([203, 0, 113, 1], 10002)), + dest_addr: SocketAddr::from(([198, 51, 100, 1], 20000)), + transaction_id: 9, + packet_count_per_batch: 2, + packet_batch_count: 3, + packet_interval_ms: 0, + }; + + send_cone_hole_punch_packets(socket.clone(), &request) + .await + .unwrap(); + + let sent = socket.sent.lock().await; + assert_eq!(sent.len(), 6); + assert!(sent.iter().all(|(_, addr)| *addr == request.dest_addr)); + } + + #[tokio::test] + async fn cone_packet_sender_rejects_malformed_dest_ip() { + let socket = Arc::new(MockSocket { + local_addr: SocketAddr::from(([127, 0, 0, 1], 10004)), + sent: tokio::sync::Mutex::new(Vec::new()), + fail_next_send: AtomicUsize::new(0), + }); + let request = SendPunchPacketCone { + listener_mapped_addr: SocketAddr::from(([203, 0, 113, 1], 10004)), + dest_addr: SocketAddr::from(([0, 0, 0, 0], 20000)), + transaction_id: 9, + packet_count_per_batch: 2, + packet_batch_count: 3, + packet_interval_ms: 0, + }; + + let err = send_cone_hole_punch_packets(socket.clone(), &request) + .await + .unwrap_err(); + + assert!(err.to_string().contains("dest_ip is malformed")); + assert!(socket.sent.lock().await.is_empty()); + } + + #[tokio::test] + async fn cone_packet_sender_continues_after_send_error() { + let socket = Arc::new(MockSocket { + local_addr: SocketAddr::from(([127, 0, 0, 1], 10005)), + sent: tokio::sync::Mutex::new(Vec::new()), + fail_next_send: AtomicUsize::new(1), + }); + let request = SendPunchPacketCone { + listener_mapped_addr: SocketAddr::from(([203, 0, 113, 1], 10005)), + dest_addr: SocketAddr::from(([198, 51, 100, 1], 20000)), + transaction_id: 9, + packet_count_per_batch: 2, + packet_batch_count: 1, + packet_interval_ms: 0, + }; + + send_cone_hole_punch_packets(socket.clone(), &request) + .await + .unwrap(); + + let sent = socket.sent.lock().await; + assert_eq!(sent.len(), 1); + } + + #[tokio::test] + async fn symmetric_packet_sender_returns_next_port_index() { + let socket = Arc::new(MockSocket { + local_addr: SocketAddr::from(([127, 0, 0, 1], 10003)), + sent: tokio::sync::Mutex::new(Vec::new()), + fail_next_send: AtomicUsize::new(0), + }); + let next_idx = send_symmetric_hole_punch_packet( + &[10, 11, 12], + socket.clone(), + 9, + &[Ipv4Addr::new(198, 51, 100, 1)], + 1, + 2, + ) + .await + .unwrap(); + + assert_eq!(next_idx, 0); + let sent = socket.sent.lock().await; + assert_eq!(sent.len(), 6); + assert_eq!(sent[0].1, SocketAddr::from(([198, 51, 100, 1], 11))); + assert_eq!(sent[3].1, SocketAddr::from(([198, 51, 100, 1], 12))); + } + + #[tokio::test] + async fn symmetric_packet_sender_returns_send_error() { + let socket = Arc::new(MockSocket { + local_addr: SocketAddr::from(([127, 0, 0, 1], 10006)), + sent: tokio::sync::Mutex::new(Vec::new()), + fail_next_send: AtomicUsize::new(1), + }); + + let err = send_symmetric_hole_punch_packet( + &[10], + socket.clone(), + 9, + &[Ipv4Addr::new(198, 51, 100, 1)], + 0, + 1, + ) + .await + .unwrap_err(); + + assert!(err.to_string().contains("mock send failure")); + assert!(socket.sent.lock().await.is_empty()); + } + + #[tokio::test] + async fn both_easy_sym_server_reports_busy_while_task_running() { + let runtime = Arc::new(MockRuntime::new(Vec::new())); + let sink = Arc::new(MockSink::default()); + let common = Arc::new(UdpHolePunchServerCommon::new(runtime, mock_stun(), sink)); + let server = UdpBothEasySymPunchServer::new(common); + let request = SendPunchPacketBothEasySym { + udp_socket_count: 1, + public_ip: Ipv4Addr::new(198, 51, 100, 1), + transaction_id: 9, + dst_port_num: 20000, + wait_time_ms: 500, + }; + + let first_response = server + .send_punch_packet_both_easy_sym(request.clone()) + .await + .unwrap(); + assert!(!first_response.is_busy); + assert_eq!( + first_response.base_mapped_addr, + Some(SocketAddr::from(([203, 0, 113, 1], 10000))) + ); + + let busy_response = server + .send_punch_packet_both_easy_sym(request) + .await + .unwrap(); + assert!(busy_response.is_busy); + assert!(busy_response.base_mapped_addr.is_none()); + } +} diff --git a/easytier-core/src/connectivity/hole_punch/udp/socket_array.rs b/easytier-core/src/connectivity/hole_punch/udp/socket_array.rs new file mode 100644 index 00000000..00a5ef5f --- /dev/null +++ b/easytier-core/src/connectivity/hole_punch/udp/socket_array.rs @@ -0,0 +1,376 @@ +use std::{ + fmt::Debug, + net::SocketAddr, + sync::{Arc, Mutex}, +}; + +use dashmap::{DashMap, DashSet}; +use tokio::task::JoinSet; +use tracing::{Instrument, Level, instrument}; + +use crate::{ + foundation::task::reap_joinset_background, + packet::{HOLE_PUNCH_PACKET_BODY_LEN, hole_punch_packet_tid}, + socket::{ + IpVersion, SocketContext, + udp::{UdpBindOptions, VirtualUdpSocket, VirtualUdpSocketFactory}, + }, +}; + +pub struct PunchedUdpSocket { + pub socket: Arc, + pub tid: u32, + pub remote_addr: SocketAddr, +} + +impl Debug for PunchedUdpSocket { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("PunchedUdpSocket") + .field("tid", &self.tid) + .field("remote_addr", &self.remote_addr) + .finish_non_exhaustive() + } +} + +pub struct UdpSocketArray +where + R: VirtualUdpSocketFactory, +{ + sockets: Arc>>, + max_socket_count: usize, + socket_factory: Arc, + socket_context: SocketContext, + tasks: Arc>>, + + interest_tids: Arc>, + tid_to_socket: Arc>>>, +} + +impl UdpSocketArray +where + R: VirtualUdpSocketFactory, +{ + pub fn new_with_context( + max_socket_count: usize, + socket_factory: Arc, + socket_context: SocketContext, + ) -> Self { + let tasks = Arc::new(Mutex::new(JoinSet::new())); + tokio::spawn(reap_joinset_background(tasks.clone(), "UdpSocketArray")); + + Self { + sockets: Arc::new(DashMap::new()), + max_socket_count, + socket_factory, + socket_context, + tasks, + + interest_tids: Arc::new(DashSet::new()), + tid_to_socket: Arc::new(DashMap::new()), + } + } + + pub fn started(&self) -> bool { + !self.sockets.is_empty() + } + + pub async fn add_new_socket(&self, socket: Arc) -> anyhow::Result<()> { + let socket_map = self.sockets.clone(); + let local_addr = socket.local_addr()?; + let interest_tids = self.interest_tids.clone(); + let tid_to_socket = self.tid_to_socket.clone(); + socket_map.insert(local_addr, socket.clone()); + self.tasks.lock().unwrap().spawn( + async move { + let _socket_map_guard = RemoveSocketOnDrop { + sockets: socket_map, + local_addr, + }; + let mut buf = [0u8; super::udp_packet_len(HOLE_PUNCH_PACKET_BODY_LEN)]; + tracing::trace!(?local_addr, "udp socket added"); + loop { + let Ok((len, addr)) = socket.recv_from(&mut buf).await else { + break; + }; + + tracing::debug!(?len, ?addr, "got raw packet"); + + let packet = &buf[..len]; + let Some(tid) = hole_punch_packet_tid(packet, HOLE_PUNCH_PACKET_BODY_LEN) + else { + continue; + }; + + tracing::debug!(?addr, ?tid, "got udp hole punch packet"); + + if interest_tids.contains(&tid) { + tracing::info!(?addr, ?tid, "got hole punching packet with interest tid"); + tid_to_socket + .entry(tid) + .or_default() + .push(PunchedUdpSocket { + socket: socket.clone(), + tid, + remote_addr: addr, + }); + break; + } + } + tracing::debug!(?local_addr, "udp socket recv loop end"); + } + .instrument(tracing::info_span!("udp array socket recv loop")), + ); + Ok(()) + } + + #[instrument(err)] + pub async fn start(&self) -> anyhow::Result<()> { + tracing::info!("starting udp socket array"); + + while self.sockets.len() < self.max_socket_count { + let socket = self + .socket_factory + .bind_udp( + UdpBindOptions::hole_punch_candidate() + .with_context(self.socket_context.clone().with_ip_version(IpVersion::V4)), + ) + .await?; + self.add_new_socket(socket).await?; + } + + Ok(()) + } + + #[instrument(err)] + pub async fn send_with_all(&self, data: &[u8], addr: SocketAddr) -> anyhow::Result<()> { + tracing::info!(?addr, "sending hole punching packet"); + + let sockets = self + .sockets + .iter() + .map(|s| s.value().clone()) + .collect::>(); + + for socket in sockets.iter() { + for _ in 0..3 { + socket.send_to(data, addr).await?; + } + } + + Ok(()) + } + + #[instrument(ret(level = Level::DEBUG))] + pub fn try_fetch_punched_socket(&self, tid: u32) -> Option> { + tracing::debug!(?tid, "try fetch punched socket"); + self.tid_to_socket.get_mut(&tid)?.value_mut().pop() + } + + pub fn add_interest_tid(&self, tid: u32) { + self.interest_tids.insert(tid); + } + + pub fn add_intreast_tid(&self, tid: u32) { + self.add_interest_tid(tid); + } + + pub fn remove_interest_tid(&self, tid: u32) { + self.interest_tids.remove(&tid); + self.tid_to_socket.remove(&tid); + } + + pub fn remove_intreast_tid(&self, tid: u32) { + self.remove_interest_tid(tid); + } +} + +impl Debug for UdpSocketArray +where + R: VirtualUdpSocketFactory, +{ + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("UdpSocketArray") + .field("sockets", &self.sockets.len()) + .field("max_socket_count", &self.max_socket_count) + .field("started", &self.started()) + .field("interest_tids", &self.interest_tids.len()) + .field("tid_to_socket", &self.tid_to_socket.len()) + .finish() + } +} + +struct RemoveSocketOnDrop { + sockets: Arc>>, + local_addr: SocketAddr, +} + +impl Drop for RemoveSocketOnDrop { + fn drop(&mut self) { + self.sockets.remove(&self.local_addr); + } +} + +#[cfg(test)] +mod tests { + use std::{ + collections::VecDeque, + io, + sync::atomic::{AtomicU16, Ordering}, + }; + + use async_trait::async_trait; + use tokio::sync::Mutex as TokioMutex; + + use super::*; + use crate::{packet::new_hole_punch_packet, socket::NetNamespace}; + + impl UdpSocketArray + where + R: VirtualUdpSocketFactory, + { + fn new(max_socket_count: usize, socket_factory: Arc) -> Self { + Self::new_with_context(max_socket_count, socket_factory, SocketContext::default()) + } + } + + struct MockSocket { + local_addr: SocketAddr, + incoming: TokioMutex, SocketAddr)>>, + sent: TokioMutex, SocketAddr)>>, + } + + impl MockSocket { + fn new(local_addr: SocketAddr, incoming: Vec<(Vec, SocketAddr)>) -> Self { + Self { + local_addr, + incoming: TokioMutex::new(incoming.into()), + sent: TokioMutex::new(Vec::new()), + } + } + } + + #[async_trait] + impl VirtualUdpSocket for MockSocket { + fn local_addr(&self) -> io::Result { + Ok(self.local_addr) + } + + async fn send_to(&self, data: &[u8], addr: SocketAddr) -> io::Result { + self.sent.lock().await.push((data.to_vec(), addr)); + Ok(data.len()) + } + + async fn recv_from(&self, buf: &mut [u8]) -> io::Result<(usize, SocketAddr)> { + let packet = self.incoming.lock().await.pop_front(); + let Some((packet, addr)) = packet else { + return std::future::pending().await; + }; + buf[..packet.len()].copy_from_slice(&packet); + Ok((packet.len(), addr)) + } + } + + struct MockFactory { + next_port: AtomicU16, + bind_options: TokioMutex>, + } + + impl MockFactory { + fn new() -> Self { + Self { + next_port: AtomicU16::new(10000), + bind_options: TokioMutex::new(Vec::new()), + } + } + } + + #[async_trait] + impl VirtualUdpSocketFactory for MockFactory { + type Socket = MockSocket; + + async fn bind_udp(&self, options: UdpBindOptions) -> anyhow::Result> { + self.bind_options.lock().await.push(options); + let port = self.next_port.fetch_add(1, Ordering::Relaxed); + Ok(Arc::new(MockSocket::new( + SocketAddr::from(([127, 0, 0, 1], port)), + Vec::new(), + ))) + } + } + + #[tokio::test] + async fn fetches_socket_when_interested_tid_is_received() { + let runtime = Arc::new(MockFactory::new()); + let array = UdpSocketArray::new(0, runtime); + let tid = 7; + let remote_addr = SocketAddr::from(([10, 0, 0, 1], 1234)); + let packet = new_hole_punch_packet(tid, HOLE_PUNCH_PACKET_BODY_LEN) + .into_bytes() + .to_vec(); + let socket = Arc::new(MockSocket::new( + SocketAddr::from(([127, 0, 0, 1], 20000)), + vec![(packet, remote_addr)], + )); + + array.add_interest_tid(tid); + array.add_new_socket(socket.clone()).await.unwrap(); + + for _ in 0..10 { + if let Some(punched) = array.try_fetch_punched_socket(tid) { + assert_eq!(punched.tid, tid); + assert_eq!(punched.remote_addr, remote_addr); + assert!(Arc::ptr_eq(&punched.socket, &socket)); + return; + } + tokio::task::yield_now().await; + } + + panic!("punched socket was not recorded"); + } + + #[tokio::test] + async fn send_with_all_sends_three_packets_per_socket() { + let runtime = Arc::new(MockFactory::new()); + let array = UdpSocketArray::new(0, runtime); + let socket = Arc::new(MockSocket::new( + SocketAddr::from(([127, 0, 0, 1], 20001)), + Vec::new(), + )); + let remote_addr = SocketAddr::from(([10, 0, 0, 2], 1235)); + + array.add_new_socket(socket.clone()).await.unwrap(); + array.send_with_all(b"abc", remote_addr).await.unwrap(); + + let sent = socket.sent.lock().await; + assert_eq!(sent.len(), 3); + assert!( + sent.iter() + .all(|(data, addr)| data == b"abc" && *addr == remote_addr) + ); + } + + #[tokio::test] + async fn start_binds_up_to_max_socket_count() { + let runtime = Arc::new(MockFactory::new()); + let context = SocketContext::default() + .with_socket_mark(Some(0)) + .with_netns(Some(NetNamespace::new("instance-a"))); + let array = UdpSocketArray::new_with_context(2, runtime.clone(), context.clone()); + + array.start().await.unwrap(); + + assert!(array.started()); + assert_eq!(array.sockets.len(), 2); + + let bind_options = runtime.bind_options.lock().await; + assert_eq!( + bind_options.as_slice(), + &[ + UdpBindOptions::hole_punch_candidate() + .with_context(context.clone().with_ip_version(IpVersion::V4)), + UdpBindOptions::hole_punch_candidate() + .with_context(context.with_ip_version(IpVersion::V4)), + ] + ); + } +} diff --git a/easytier-core/src/connectivity/hole_punch/udp/task.rs b/easytier-core/src/connectivity/hole_punch/udp/task.rs new file mode 100644 index 00000000..4616dbf5 --- /dev/null +++ b/easytier-core/src/connectivity/hole_punch/udp/task.rs @@ -0,0 +1,217 @@ +use crate::{ + config::{P2pPolicyFlags, PeerId}, + proto::common::{NatType, PeerFeatureFlag}, +}; + +use super::{ + super::policy::{should_background_p2p_with_peer, should_try_p2p_with_peer}, + UdpNatType, +}; + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct UdpPunchCandidate { + pub peer_id: PeerId, + pub udp_nat_type: NatType, + pub feature_flag: Option, + pub has_direct_connection: bool, + pub has_recent_traffic: bool, +} + +#[derive(Clone, Copy, Debug, Hash, Eq, PartialEq)] +pub struct UdpPunchTaskInfo { + pub dst_peer_id: PeerId, + pub dst_nat_type: UdpNatType, + pub my_nat_type: UdpNatType, +} + +pub fn collect_udp_punch_tasks( + my_peer_id: PeerId, + my_nat_type: UdpNatType, + policy: P2pPolicyFlags, + candidates: I, + is_blacklisted: F, +) -> Vec +where + I: IntoIterator, + F: Fn(PeerId) -> bool, +{ + if my_nat_type.is_open() { + return Vec::new(); + } + + candidates + .into_iter() + .filter_map(|candidate| { + let static_allowed = should_background_p2p_with_peer( + candidate.feature_flag.as_ref(), + false, + policy.lazy_p2p, + policy.disable_p2p, + policy.need_p2p, + ); + let dynamic_allowed = should_try_p2p_with_peer( + candidate.feature_flag.as_ref(), + false, + policy.disable_p2p, + policy.need_p2p, + ) && candidate.has_recent_traffic; + if !static_allowed && !dynamic_allowed { + return None; + } + + let peer_id = candidate.peer_id; + if is_blacklisted(peer_id) || candidate.has_direct_connection { + return None; + } + + let peer_nat_type = candidate.udp_nat_type.into(); + if !my_nat_type.can_punch_hole_as_client( + peer_nat_type, + my_peer_id, + peer_id, + policy.disable_sym_hole_punching, + ) { + return None; + } + + Some(UdpPunchTaskInfo { + dst_peer_id: peer_id, + dst_nat_type: peer_nat_type, + my_nat_type, + }) + }) + .collect() +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::proto::common::PeerFeatureFlag; + + fn candidate(peer_id: PeerId, udp_nat_type: NatType) -> UdpPunchCandidate { + UdpPunchCandidate { + peer_id, + udp_nat_type, + feature_flag: Some(PeerFeatureFlag::default()), + has_direct_connection: false, + has_recent_traffic: false, + } + } + + fn collect( + my_peer_id: PeerId, + my_nat_type: NatType, + policy: P2pPolicyFlags, + candidates: Vec, + ) -> Vec { + collect_udp_punch_tasks(my_peer_id, my_nat_type.into(), policy, candidates, |_| { + false + }) + } + + #[test] + fn open_nat_does_not_start_udp_punch_tasks() { + let tasks = collect( + 1, + NatType::OpenInternet, + P2pPolicyFlags::default(), + vec![candidate(2, NatType::PortRestricted)], + ); + + assert!(tasks.is_empty()); + } + + #[test] + fn lazy_p2p_allows_recent_traffic_without_need_p2p_flag() { + let mut idle = candidate(2, NatType::PortRestricted); + idle.feature_flag = Some(PeerFeatureFlag { + need_p2p: false, + ..Default::default() + }); + + let mut active = idle.clone(); + active.peer_id = 3; + active.has_recent_traffic = true; + + let tasks = collect( + 1, + NatType::PortRestricted, + P2pPolicyFlags { + lazy_p2p: true, + ..Default::default() + }, + vec![idle, active], + ); + + assert_eq!(tasks.len(), 1); + assert_eq!(tasks[0].dst_peer_id, 3); + } + + #[test] + fn skips_blacklisted_and_directly_connected_candidates() { + let mut direct = candidate(2, NatType::PortRestricted); + direct.has_direct_connection = true; + + let tasks = collect_udp_punch_tasks( + 1, + NatType::PortRestricted.into(), + P2pPolicyFlags::default(), + vec![direct, candidate(3, NatType::PortRestricted)], + |peer_id| peer_id == 3, + ); + + assert!(tasks.is_empty()); + } + + #[test] + fn filters_candidates_by_udp_nat_method() { + let tasks = collect( + 1, + NatType::PortRestricted, + P2pPolicyFlags::default(), + vec![ + candidate(2, NatType::Symmetric), + candidate(3, NatType::PortRestricted), + ], + ); + + assert_eq!(tasks.len(), 1); + assert_eq!(tasks[0].dst_peer_id, 3); + assert_eq!(tasks[0].dst_nat_type, NatType::PortRestricted.into()); + } + + #[test] + fn easy_symmetric_pair_uses_lower_peer_id_as_initiator() { + let tasks = collect( + 1, + NatType::SymmetricEasyInc, + P2pPolicyFlags::default(), + vec![candidate(2, NatType::SymmetricEasyDec)], + ); + assert_eq!(tasks.len(), 1); + + let tasks = collect( + 2, + NatType::SymmetricEasyInc, + P2pPolicyFlags::default(), + vec![candidate(1, NatType::SymmetricEasyDec)], + ); + assert!(tasks.is_empty()); + } + + #[test] + fn disabling_symmetric_hole_punch_keeps_sym_to_cone_as_cone_method() { + let tasks = collect( + 1, + NatType::Symmetric, + P2pPolicyFlags { + disable_sym_hole_punching: true, + ..Default::default() + }, + vec![candidate(2, NatType::PortRestricted)], + ); + + assert_eq!(tasks.len(), 1); + assert_eq!(tasks[0].dst_peer_id, 2); + } +} diff --git a/easytier-core/src/connectivity/manual/discovery.rs b/easytier-core/src/connectivity/manual/discovery.rs new file mode 100644 index 00000000..ea52e7bd --- /dev/null +++ b/easytier-core/src/connectivity/manual/discovery.rs @@ -0,0 +1,73 @@ +use std::time::Duration; + +use serde::{Deserialize, Serialize}; + +use crate::socket::{IpVersion, SocketContext, tcp::TcpBindOptions}; + +#[cfg(feature = "endpoint-discovery")] +mod implementation; + +#[cfg(feature = "endpoint-discovery")] +pub(crate) use implementation::CoreManualEndpointResolver; + +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct ManualEndpointDiscoveryConfig { + pub user_agent: String, + pub network_name: String, + pub http_timeout: Duration, + pub http_ip_version: IpVersion, + pub http_tcp_bind: TcpBindOptions, + pub dns_record_context: SocketContext, + pub srv_protocols: Vec, +} + +impl Default for ManualEndpointDiscoveryConfig { + fn default() -> Self { + Self { + user_agent: "easytier-core".to_owned(), + network_name: String::new(), + http_timeout: Duration::from_secs(20), + http_ip_version: IpVersion::Both, + http_tcp_bind: TcpBindOptions::default(), + dns_record_context: SocketContext::default(), + srv_protocols: vec!["tcp".to_owned(), "udp".to_owned()], + } + } +} + +#[cfg(not(feature = "endpoint-discovery"))] +pub(crate) struct CoreManualEndpointResolver +where + H: crate::socket::tcp::VirtualTcpSocketFactory, +{ + _host: std::marker::PhantomData H>, +} + +#[cfg(not(feature = "endpoint-discovery"))] +impl CoreManualEndpointResolver +where + H: crate::socket::tcp::VirtualTcpSocketFactory, +{ + pub fn new( + host: std::sync::Arc, + dns: std::sync::Arc, + dns_records: std::sync::Arc, + config: ManualEndpointDiscoveryConfig, + ) -> Self { + let _ = (host, dns, dns_records, config); + Self { + _host: std::marker::PhantomData, + } + } +} + +#[cfg(not(feature = "endpoint-discovery"))] +#[async_trait::async_trait] +impl super::ManualEndpointResolver for CoreManualEndpointResolver +where + H: crate::socket::tcp::VirtualTcpSocketFactory, +{ + async fn resolve_endpoint(&self, url: &url::Url) -> anyhow::Result { + anyhow::bail!("endpoint discovery is disabled for {url}") + } +} diff --git a/easytier-core/src/connectivity/manual/discovery/implementation.rs b/easytier-core/src/connectivity/manual/discovery/implementation.rs new file mode 100644 index 00000000..17a2aafa --- /dev/null +++ b/easytier-core/src/connectivity/manual/discovery/implementation.rs @@ -0,0 +1,924 @@ +use std::{collections::HashSet, sync::Arc, time::Duration}; + +use anyhow::Context as _; +use bytes::Bytes; +use http_body_util::{BodyExt as _, Empty}; +use hyper::{Request, header}; +use hyper_util::rt::TokioIo; +use rand::{Rng as _, seq::SliceRandom}; +use rustls::pki_types::ServerName; +use tokio::io::{AsyncRead, AsyncWrite}; +use tokio_rustls::TlsConnector; +use tokio_util::task::AbortOnDropHandle; +use url::Url; + +use crate::{ + connectivity::transport, + host::dns::{DnsQuery, DnsRecordResolver, DnsResolver, DnsSrvRecord}, + socket::{ + IpVersion, SocketContext, + tcp::{TcpBindOptions, TcpSocketPurpose, VirtualTcpSocketFactory}, + }, +}; + +use super::super::{ManualEndpointResolver, resolve_url_addrs}; +use super::ManualEndpointDiscoveryConfig; + +const HTTP_DEFAULT_PORT: u16 = 80; +const HTTPS_DEFAULT_PORT: u16 = 443; + +#[derive(Debug, Clone)] +pub(crate) struct HttpDiscoveryRequest { + pub url: Url, + pub user_agent: String, + pub network_name: String, + pub timeout: Duration, + pub ip_version: IpVersion, + pub tcp_bind: TcpBindOptions, +} + +pub(crate) struct CoreManualEndpointResolver +where + H: VirtualTcpSocketFactory, +{ + host: Arc, + dns: Arc, + dns_records: Arc, + config: ManualEndpointDiscoveryConfig, +} + +impl CoreManualEndpointResolver +where + H: VirtualTcpSocketFactory, +{ + pub fn new( + host: Arc, + dns: Arc, + dns_records: Arc, + config: ManualEndpointDiscoveryConfig, + ) -> Self { + Self { + host, + dns, + dns_records, + config, + } + } +} + +#[async_trait::async_trait] +impl ManualEndpointResolver for CoreManualEndpointResolver +where + H: VirtualTcpSocketFactory, +{ + async fn resolve_endpoint(&self, url: &Url) -> anyhow::Result { + match url.scheme() { + "http" | "https" => { + let response = fetch_http_discovery( + self.host.clone(), + self.dns.as_ref(), + HttpDiscoveryRequest { + url: url.clone(), + user_agent: self.config.user_agent.clone(), + network_name: self.config.network_name.clone(), + timeout: self.config.http_timeout, + ip_version: self.config.http_ip_version, + tcp_bind: self.config.http_tcp_bind.clone(), + }, + ) + .await?; + resolve_http_endpoint(response) + .map(|endpoint| endpoint.url) + .map_err(|error| anyhow::anyhow!("Invalid Url: {error}")) + } + "txt" => { + let host = endpoint_host(url)?; + resolve_txt_endpoint( + self.dns_records.as_ref(), + host, + self.config.dns_record_context.clone(), + ) + .await + } + "srv" => { + let host = endpoint_host(url)?; + resolve_srv_endpoint( + self.dns_records.as_ref(), + host, + &self.config.srv_protocols, + self.config.dns_record_context.clone(), + ) + .await + } + scheme => anyhow::bail!("unsupported manual endpoint resolver scheme: {scheme}"), + } + } +} + +pub(super) fn endpoint_host(url: &Url) -> anyhow::Result<&str> { + url.host_str() + .ok_or_else(|| anyhow::anyhow!("host should not be empty in {url}")) +} + +pub(crate) async fn fetch_http_discovery( + host: Arc, + dns: &dyn DnsResolver, + request: HttpDiscoveryRequest, +) -> anyhow::Result +where + H: VirtualTcpSocketFactory, +{ + let timeout = request.timeout; + crate::foundation::time::timeout(timeout, fetch_http_discovery_inner(host, dns, request)) + .await + .map_err(|_| anyhow::anyhow!("HTTP discovery timed out after {timeout:?}"))? +} + +async fn fetch_http_discovery_inner( + host: Arc, + dns: &dyn DnsResolver, + request: HttpDiscoveryRequest, +) -> anyhow::Result +where + H: VirtualTcpSocketFactory, +{ + let default_port = match request.url.scheme() { + "http" => HTTP_DEFAULT_PORT, + "https" => HTTPS_DEFAULT_PORT, + scheme => anyhow::bail!("unsupported HTTP discovery scheme: {scheme}"), + }; + let addrs = resolve_url_addrs( + &request.url, + default_port, + request + .tcp_bind + .context + .clone() + .with_ip_version(request.ip_version), + dns, + ) + .await?; + + let mut last_error = None; + let mut socket = None; + for addr in addrs { + match transport::connect_tcp( + host.clone(), + addr, + Vec::new(), + request.tcp_bind.clone(), + TcpSocketPurpose::ManualConnect, + ) + .await + { + Ok(connected) => { + socket = Some(connected); + break; + } + Err(error) => last_error = Some(error), + } + } + let socket = socket.ok_or_else(|| { + last_error.unwrap_or_else(|| anyhow::anyhow!("no HTTP discovery address candidates")) + })?; + + if request.url.scheme() == "https" { + let server_name = tls_server_name(&request.url)?; + let root_store = rustls::RootCertStore { + roots: webpki_roots::TLS_SERVER_ROOTS.to_vec(), + }; + let tls_config = rustls::ClientConfig::builder() + .with_root_certificates(root_store) + .with_no_client_auth(); + let stream = TlsConnector::from(Arc::new(tls_config)) + .connect(server_name, socket) + .await + .with_context(|| format!("HTTPS handshake failed for {}", request.url))?; + send_http_discovery_request(stream, request).await + } else { + send_http_discovery_request(socket, request).await + } +} + +pub(super) fn tls_server_name(url: &Url) -> anyhow::Result> { + match url.host() { + Some(url::Host::Domain(host)) => ServerName::try_from(host.to_owned()) + .with_context(|| format!("invalid HTTPS server name in {url}")), + Some(url::Host::Ipv4(ip)) => Ok(ServerName::IpAddress(ip.into())), + Some(url::Host::Ipv6(ip)) => Ok(ServerName::IpAddress(ip.into())), + None => anyhow::bail!("HTTP discovery URL has no host: {url}"), + } +} + +async fn send_http_discovery_request( + stream: S, + request: HttpDiscoveryRequest, +) -> anyhow::Result +where + S: AsyncRead + AsyncWrite + Unpin + Send + 'static, +{ + let io = TokioIo::new(stream); + let (mut sender, connection) = hyper::client::conn::http1::handshake(io) + .await + .with_context(|| format!("starting HTTP connection failed for {}", request.url))?; + let connection_task = AbortOnDropHandle::new(tokio::spawn(connection)); + + let request_target = match request.url.query() { + Some(query) => format!("{}?{query}", request.url.path()), + None => request.url.path().to_owned(), + }; + let host_header = &request.url[url::Position::BeforeHost..url::Position::AfterPort]; + let outgoing = Request::builder() + .method("GET") + .uri(request_target) + .header(header::HOST, host_header) + .header(header::USER_AGENT, request.user_agent) + .header("X-Network-Name", request.network_name) + .header(header::CONNECTION, "close") + .body(Empty::::new())?; + let response = sender + .send_request(outgoing) + .await + .with_context(|| format!("sending HTTP request failed for {}", request.url))?; + let status_code = response.status().as_u16(); + let location = response + .headers() + .get(header::LOCATION) + .map(|value| String::from_utf8_lossy(value.as_bytes()).into_owned()); + let body = response + .into_body() + .collect() + .await + .with_context(|| format!("reading HTTP response failed for {}", request.url))? + .to_bytes(); + drop(sender); + connection_task + .await + .context("HTTP connection task failed")? + .with_context(|| format!("HTTP connection failed for {}", request.url))?; + + Ok(HttpDiscoveryResponse { + status_code, + location, + body: String::from_utf8_lossy(&body).into_owned(), + }) +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub(crate) enum HttpEndpointSource { + RedirectQuery, + RedirectUrl, + ResponseBody, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub(crate) struct HttpDiscoveryResponse { + pub status_code: u16, + pub location: Option, + pub body: String, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub(crate) struct ResolvedHttpEndpoint { + pub url: Url, + pub source: HttpEndpointSource, +} + +fn resolve_http_redirect(location: &str) -> anyhow::Result { + let url = Url::parse(location) + .with_context(|| format!("parsing redirect URL failed. url: {location}"))?; + if !matches!(url.scheme(), "http" | "https") { + return Ok(ResolvedHttpEndpoint { + url, + source: HttpEndpointSource::RedirectUrl, + }); + } + + let candidates = url + .query_pairs() + .filter_map(|(_, value)| Url::parse(&value).ok()) + .collect::>(); + if let Some(url) = candidates.choose(&mut rand::thread_rng()).cloned() { + return Ok(ResolvedHttpEndpoint { + url, + source: HttpEndpointSource::RedirectQuery, + }); + } + + if let Some(url) = location + .strip_prefix(&format!("{}://", url.scheme())) + .and_then(|value| Url::parse(value).ok()) + { + return Ok(ResolvedHttpEndpoint { + url, + source: HttpEndpointSource::RedirectUrl, + }); + } + + anyhow::bail!("no valid connector URL found in redirect location {location:?}") +} + +fn resolve_http_body(body: &str) -> anyhow::Result { + let mut candidates = body + .lines() + .map(str::trim) + .filter(|line| !line.is_empty()) + .collect::>(); + candidates.shuffle(&mut rand::thread_rng()); + for candidate in candidates { + if let Ok(url) = Url::parse(candidate) { + return Ok(ResolvedHttpEndpoint { + url, + source: HttpEndpointSource::ResponseBody, + }); + } + } + anyhow::bail!("no valid connector URL found in response body {body:?}") +} + +pub(crate) fn resolve_http_endpoint( + response: HttpDiscoveryResponse, +) -> anyhow::Result { + match response.status_code { + 300..=399 => resolve_http_redirect( + response + .location + .as_deref() + .ok_or_else(|| anyhow::anyhow!("HTTP redirect has no Location header"))?, + ), + 200..=299 => resolve_http_body(&response.body), + status_code => anyhow::bail!( + "unexpected HTTP discovery status {status_code}, body: {:?}", + response.body + ), + } +} + +fn choose_weighted(options: &[(T, u64)]) -> Option<&T> { + let total_weight = options.iter().map(|(_, weight)| *weight).sum(); + let mut rng = rand::thread_rng(); + let selected = rng.gen_range(0..total_weight); + let mut accumulated = 0; + + for (item, weight) in options { + accumulated += *weight; + if selected < accumulated { + return Some(item); + } + } + None +} + +pub(crate) async fn resolve_txt_endpoint( + resolver: &dyn DnsRecordResolver, + domain_name: &str, + context: SocketContext, +) -> anyhow::Result { + let txt_data = resolver + .resolve_txt(DnsQuery::new(domain_name, context)) + .await + .with_context(|| format!("resolve TXT record failed for {domain_name}"))?; + let candidates = txt_data + .split(' ') + .filter_map(|candidate| Url::parse(candidate).ok()) + .collect::>(); + candidates + .choose(&mut rand::thread_rng()) + .cloned() + .ok_or_else(|| { + anyhow::anyhow!( + "no valid URL found in TXT data {txt_data:?}; expected a space-separated URL list" + ) + }) +} + +fn srv_record_url(protocol: &str, record: DnsSrvRecord) -> anyhow::Result<(Url, u64)> { + if record.port == 0 { + anyhow::bail!("SRV port must be non-zero"); + } + let url = format!("{protocol}://{}:{}", record.target, record.port); + // Preserve the existing EasyTier selection rule, which treats SRV priority + // as the candidate weight. + Ok((Url::parse(&url)?, u64::from(record.priority))) +} + +pub(super) fn deduplicate_srv_candidates(candidates: Vec<(Url, u64)>) -> Vec<(Url, u64)> { + candidates + .into_iter() + .collect::>() + .into_iter() + .collect() +} + +pub(crate) async fn resolve_srv_endpoint( + resolver: &dyn DnsRecordResolver, + domain_name: &str, + supported_protocols: &[String], + context: SocketContext, +) -> anyhow::Result { + let lookups = supported_protocols.iter().map(|protocol| { + let protocol = protocol.clone(); + let query = DnsQuery::new( + format!("_easytier._{protocol}.{domain_name}"), + context.clone(), + ); + async move { (protocol, resolver.resolve_srv(query).await) } + }); + + let mut candidates = Vec::new(); + for (protocol, result) in futures::future::join_all(lookups).await { + let Ok(records) = result else { + continue; + }; + candidates.extend(records.into_iter().filter_map(|record| { + match srv_record_url(&protocol, record) { + Ok(candidate) => Some(candidate), + Err(error) => { + tracing::warn!( + ?error, + srv_domain = %format!("_easytier._{protocol}.{domain_name}"), + "ignore invalid SRV endpoint record" + ); + None + } + } + })); + } + if candidates.is_empty() { + anyhow::bail!("no SRV endpoint found for {domain_name}"); + } + let candidates = deduplicate_srv_candidates(candidates); + + choose_weighted(&candidates) + .cloned() + .ok_or_else(|| anyhow::anyhow!("failed to choose an SRV endpoint for {domain_name}")) +} + +#[cfg(test)] +mod tests { + use std::{ + io, + net::{IpAddr, Ipv4Addr, SocketAddr}, + pin::Pin, + sync::Mutex, + task::{Context, Poll}, + }; + + use async_trait::async_trait; + use tokio::io::{AsyncReadExt as _, AsyncWriteExt as _, DuplexStream, ReadBuf}; + + use crate::socket::tcp::{TcpConnectOptions, VirtualTcpSocket, VirtualTcpSocketFactory}; + + use super::*; + + struct HttpTestSocket { + stream: DuplexStream, + local_addr: SocketAddr, + peer_addr: SocketAddr, + } + + impl AsyncRead for HttpTestSocket { + fn poll_read( + mut self: Pin<&mut Self>, + cx: &mut Context<'_>, + buf: &mut ReadBuf<'_>, + ) -> Poll> { + Pin::new(&mut self.stream).poll_read(cx, buf) + } + } + + impl AsyncWrite for HttpTestSocket { + fn poll_write( + mut self: Pin<&mut Self>, + cx: &mut Context<'_>, + buf: &[u8], + ) -> Poll> { + Pin::new(&mut self.stream).poll_write(cx, buf) + } + + fn poll_flush(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + Pin::new(&mut self.stream).poll_flush(cx) + } + + fn poll_shutdown(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + Pin::new(&mut self.stream).poll_shutdown(cx) + } + } + + impl VirtualTcpSocket for HttpTestSocket { + fn local_addr(&self) -> io::Result { + Ok(self.local_addr) + } + + fn peer_addr(&self) -> io::Result { + Ok(self.peer_addr) + } + } + + struct HttpTestHost { + stream: Mutex>, + connects: Mutex>, + } + + #[async_trait] + impl VirtualTcpSocketFactory for HttpTestHost { + type Socket = HttpTestSocket; + + async fn connect_tcp(&self, options: TcpConnectOptions) -> anyhow::Result { + self.connects.lock().unwrap().push(options.clone()); + let stream = self + .stream + .lock() + .unwrap() + .take() + .ok_or_else(|| anyhow::anyhow!("test socket already connected"))?; + Ok(HttpTestSocket { + stream, + local_addr: "192.0.2.2:40000".parse().unwrap(), + peer_addr: options.remote_addr, + }) + } + } + + struct HttpTestDns { + queries: Mutex>, + } + + #[async_trait] + impl DnsResolver for HttpTestDns { + async fn resolve(&self, query: DnsQuery) -> anyhow::Result> { + self.queries.lock().unwrap().push(query); + Ok(vec![IpAddr::V4(Ipv4Addr::new(192, 0, 2, 1))]) + } + } + + #[tokio::test] + async fn http_fetch_uses_host_dns_and_socket_while_core_drives_io() { + let (client, mut server) = tokio::io::duplex(8192); + let host = Arc::new(HttpTestHost { + stream: Mutex::new(Some(client)), + connects: Mutex::new(Vec::new()), + }); + let dns = HttpTestDns { + queries: Mutex::new(Vec::new()), + }; + let server_task = tokio::spawn(async move { + let mut request = Vec::new(); + loop { + let mut chunk = [0; 1024]; + let len = server.read(&mut chunk).await.unwrap(); + assert_ne!(len, 0, "HTTP request ended before its headers"); + request.extend_from_slice(&chunk[..len]); + if request.windows(4).any(|window| window == b"\r\n\r\n") { + break; + } + } + let request = String::from_utf8(request).unwrap().to_ascii_lowercase(); + assert!(request.starts_with("get /lookup?kind=peer http/1.1\r\n")); + assert!(request.contains("host: discovery.example:18080\r\n")); + assert!(request.contains("user-agent: easytier/test\r\n")); + assert!(request.contains("x-network-name: test-network\r\n")); + assert!(request.contains("connection: close\r\n")); + + server + .write_all( + b"HTTP/1.1 302 Found\r\nLocation: tcp://192.0.2.10:11010\r\nTransfer-Encoding: chunked\r\nConnection: close\r\n\r\n5\r\nhello\r\n0\r\n\r\n", + ) + .await + .unwrap(); + server.shutdown().await.unwrap(); + }); + + let response = fetch_http_discovery( + host.clone(), + &dns, + HttpDiscoveryRequest { + url: "http://discovery.example:18080/lookup?kind=peer" + .parse() + .unwrap(), + user_agent: "easytier/test".to_owned(), + network_name: "test-network".to_owned(), + timeout: Duration::from_secs(1), + ip_version: IpVersion::V4, + tcp_bind: TcpBindOptions::default().with_socket_mark(Some(9)), + }, + ) + .await + .unwrap(); + crate::foundation::time::timeout(Duration::from_secs(1), server_task) + .await + .expect("HTTP fetch test server did not finish") + .unwrap(); + + assert_eq!(response.status_code, 302); + assert_eq!(response.location.as_deref(), Some("tcp://192.0.2.10:11010")); + assert_eq!(response.body, "hello"); + assert_eq!( + *dns.queries.lock().unwrap(), + [DnsQuery::new( + "discovery.example", + SocketContext { + ip_version: IpVersion::V4, + socket_mark: Some(9), + netns: None, + } + )] + ); + assert_eq!(host.connects.lock().unwrap().len(), 1); + let options = &host.connects.lock().unwrap()[0]; + assert_eq!(options.remote_addr, "192.0.2.1:18080".parse().unwrap()); + assert_eq!(options.purpose, TcpSocketPurpose::ManualConnect); + assert_eq!(options.bind.context.socket_mark, Some(9)); + } + + #[test] + fn http_discovery_interprets_redirect_and_body_forms() { + let query = resolve_http_endpoint(HttpDiscoveryResponse { + status_code: 302, + location: Some("https://discovery.example/?url=tcp://127.0.0.1:11010".to_owned()), + body: String::new(), + }) + .unwrap(); + assert_eq!(query.url.as_str(), "tcp://127.0.0.1:11010"); + assert_eq!(query.source, HttpEndpointSource::RedirectQuery); + + let nested = resolve_http_endpoint(HttpDiscoveryResponse { + status_code: 302, + location: Some("https://udp://127.0.0.1:11010".to_owned()), + body: String::new(), + }) + .unwrap(); + assert_eq!(nested.url.as_str(), "udp://127.0.0.1:11010"); + assert_eq!(nested.source, HttpEndpointSource::RedirectUrl); + + let direct = resolve_http_endpoint(HttpDiscoveryResponse { + status_code: 307, + location: Some("quic://127.0.0.1:11012".to_owned()), + body: String::new(), + }) + .unwrap(); + assert_eq!(direct.url.as_str(), "quic://127.0.0.1:11012"); + assert_eq!(direct.source, HttpEndpointSource::RedirectUrl); + + let body = resolve_http_endpoint(HttpDiscoveryResponse { + status_code: 200, + location: None, + body: "invalid\nwg://127.0.0.1:11011\n".to_owned(), + }) + .unwrap(); + assert_eq!(body.url.as_str(), "wg://127.0.0.1:11011"); + assert_eq!(body.source, HttpEndpointSource::ResponseBody); + } + + #[test] + fn http_discovery_reports_malformed_redirect_location() { + let error = resolve_http_endpoint(HttpDiscoveryResponse { + status_code: 302, + location: Some("not a URL".to_owned()), + body: String::new(), + }) + .unwrap_err(); + + let message = error.to_string(); + assert!(message.contains("parsing redirect URL failed")); + assert!(message.contains("not a URL")); + } + + #[test] + fn https_server_name_accepts_ip_literals_without_url_brackets() { + let ipv4: Url = "https://192.0.2.1/".parse().unwrap(); + let ipv6: Url = "https://[2001:db8::1]/".parse().unwrap(); + + assert_eq!( + tls_server_name(&ipv4).unwrap(), + ServerName::IpAddress(Ipv4Addr::new(192, 0, 2, 1).into()) + ); + assert_eq!( + tls_server_name(&ipv6).unwrap(), + ServerName::IpAddress("2001:db8::1".parse::().unwrap().into()) + ); + } + + struct TestResolver { + txt: String, + srv: Vec, + queries: Mutex>, + } + + #[async_trait] + impl DnsRecordResolver for TestResolver { + async fn resolve_txt(&self, query: DnsQuery) -> anyhow::Result { + self.queries.lock().unwrap().push(query); + Ok(self.txt.clone()) + } + + async fn resolve_srv(&self, query: DnsQuery) -> anyhow::Result> { + self.queries.lock().unwrap().push(query); + Ok(self.srv.clone()) + } + } + + #[tokio::test] + async fn core_endpoint_resolver_owns_record_scheme_dispatch() { + let host = Arc::new(HttpTestHost { + stream: Mutex::new(None), + connects: Mutex::new(Vec::new()), + }); + let dns: Arc = Arc::new(HttpTestDns { + queries: Mutex::new(Vec::new()), + }); + let records = Arc::new(TestResolver { + txt: "tcp://192.0.2.10:11010".to_owned(), + srv: vec![DnsSrvRecord { + priority: 1, + weight: 10, + port: 11012, + target: "peer.example.com.".to_owned(), + }], + queries: Mutex::new(Vec::new()), + }); + let record_context = SocketContext { + ip_version: IpVersion::V6, + socket_mark: Some(17), + netns: None, + }; + let resolver = CoreManualEndpointResolver::new( + host, + dns, + records.clone(), + ManualEndpointDiscoveryConfig { + user_agent: "easytier/test".to_owned(), + network_name: "test-network".to_owned(), + http_timeout: Duration::from_secs(1), + http_ip_version: IpVersion::Both, + http_tcp_bind: TcpBindOptions::default(), + dns_record_context: record_context.clone(), + srv_protocols: vec!["quic".to_owned()], + }, + ); + + let txt = resolver + .resolve_endpoint(&"txt://discovery.example".parse().unwrap()) + .await + .unwrap(); + let srv = resolver + .resolve_endpoint(&"srv://discovery.example".parse().unwrap()) + .await + .unwrap(); + + assert_eq!(txt.as_str(), "tcp://192.0.2.10:11010"); + assert_eq!(srv.as_str(), "quic://peer.example.com.:11012"); + assert_eq!( + *records.queries.lock().unwrap(), + [ + DnsQuery::new("discovery.example", record_context.clone()), + DnsQuery::new("_easytier._quic.discovery.example", record_context) + ] + ); + } + + #[tokio::test] + async fn core_endpoint_resolver_passes_http_config_and_error_classification() { + let (client, mut server) = tokio::io::duplex(8192); + let host = Arc::new(HttpTestHost { + stream: Mutex::new(Some(client)), + connects: Mutex::new(Vec::new()), + }); + let dns = Arc::new(HttpTestDns { + queries: Mutex::new(Vec::new()), + }); + let records = Arc::new(TestResolver { + txt: String::new(), + srv: Vec::new(), + queries: Mutex::new(Vec::new()), + }); + let server_task = tokio::spawn(async move { + let mut request = Vec::new(); + loop { + let mut chunk = [0; 1024]; + let len = server.read(&mut chunk).await.unwrap(); + assert_ne!(len, 0, "HTTP request ended before its headers"); + request.extend_from_slice(&chunk[..len]); + if request.windows(4).any(|window| window == b"\r\n\r\n") { + break; + } + } + let request = String::from_utf8(request).unwrap().to_ascii_lowercase(); + assert!(request.contains("user-agent: easytier/facade-test\r\n")); + assert!(request.contains("x-network-name: facade-network\r\n")); + server + .write_all( + b"HTTP/1.1 200 OK\r\nContent-Length: 9\r\nConnection: close\r\n\r\nnot a URL", + ) + .await + .unwrap(); + server.shutdown().await.unwrap(); + }); + let resolver = CoreManualEndpointResolver::new( + host.clone(), + dns.clone(), + records, + ManualEndpointDiscoveryConfig { + user_agent: "easytier/facade-test".to_owned(), + network_name: "facade-network".to_owned(), + http_timeout: Duration::from_secs(1), + http_ip_version: IpVersion::V4, + http_tcp_bind: TcpBindOptions::default().with_socket_mark(Some(23)), + dns_record_context: SocketContext::default(), + srv_protocols: Vec::new(), + }, + ); + + let error = resolver + .resolve_endpoint(&"http://discovery.example:18081/endpoint".parse().unwrap()) + .await + .unwrap_err(); + crate::foundation::time::timeout(Duration::from_secs(1), server_task) + .await + .expect("HTTP facade test server did not finish") + .unwrap(); + + assert!(error.to_string().starts_with("Invalid Url:")); + assert_eq!( + *dns.queries.lock().unwrap(), + [DnsQuery::new( + "discovery.example", + SocketContext { + ip_version: IpVersion::V4, + socket_mark: Some(23), + netns: None, + } + )] + ); + let connects = host.connects.lock().unwrap(); + assert_eq!(connects.len(), 1); + assert_eq!(connects[0].remote_addr, "192.0.2.1:18081".parse().unwrap()); + assert_eq!(connects[0].bind.context.socket_mark, Some(23)); + } + + #[tokio::test] + async fn txt_discovery_parses_easy_tier_url_candidates() { + let resolver = TestResolver { + txt: "invalid tcp://127.0.0.1:11010".to_owned(), + srv: Vec::new(), + queries: Mutex::new(Vec::new()), + }; + + let endpoint = + resolve_txt_endpoint(&resolver, "discovery.example", SocketContext::default()) + .await + .unwrap(); + + assert_eq!(endpoint.as_str(), "tcp://127.0.0.1:11010"); + assert_eq!( + *resolver.queries.lock().unwrap(), + [DnsQuery::new("discovery.example", SocketContext::default())] + ); + } + + #[tokio::test] + async fn srv_discovery_builds_protocol_specific_endpoint() { + let resolver = TestResolver { + txt: String::new(), + srv: vec![DnsSrvRecord { + priority: 1, + weight: 10, + port: 11012, + target: "peer.example.com.".to_owned(), + }], + queries: Mutex::new(Vec::new()), + }; + + let endpoint = resolve_srv_endpoint( + &resolver, + "discovery.example", + &["quic".to_owned()], + SocketContext::default(), + ) + .await + .unwrap(); + + assert_eq!(endpoint.as_str(), "quic://peer.example.com.:11012"); + assert_eq!( + *resolver.queries.lock().unwrap(), + [DnsQuery::new( + "_easytier._quic.discovery.example", + SocketContext::default() + )] + ); + } + + #[test] + fn srv_discovery_deduplicates_url_and_priority() { + let endpoint: Url = "tcp://peer.example.com:11010".parse().unwrap(); + let candidates = deduplicate_srv_candidates(vec![ + (endpoint.clone(), 10), + (endpoint.clone(), 10), + (endpoint.clone(), 20), + ]); + + assert_eq!(candidates.len(), 2); + assert!(candidates.contains(&(endpoint.clone(), 10))); + assert!(candidates.contains(&(endpoint, 20))); + } +} diff --git a/easytier-core/src/connectivity/manual/mod.rs b/easytier-core/src/connectivity/manual/mod.rs new file mode 100644 index 00000000..cf57feec --- /dev/null +++ b/easytier-core/src/connectivity/manual/mod.rs @@ -0,0 +1,1336 @@ +pub mod discovery; + +use std::{ + collections::BTreeSet, + future::Future, + net::{IpAddr, Ipv4Addr, Ipv6Addr, SocketAddr}, + sync::{Arc, Mutex, Weak}, + time::Duration, +}; + +use async_trait::async_trait; +use dashmap::DashSet; +use percent_encoding::percent_decode_str; +use quanta::Instant; +use rand::seq::SliceRandom; +use serde::{Deserialize, Serialize}; +use tokio::task::JoinSet; +use tokio_util::{sync::CancellationToken, task::AbortOnDropHandle}; +use url::Url; + +use crate::tunnel::ring::RingTunnelRegistry; +use crate::{ + connectivity::{ + protocol::{ClientProtocolUpgrader, ProtocolTransport, protocol_transport}, + transport::{self, ConnectedByteStream, ConnectedTransport, UdpSessionMode}, + }, + events::{CoreEvent, CoreEventSink}, + host::dns::{DnsQuery, DnsResolver}, + peers::peer_manager::PeerManagerCore, + proto::common::TunnelInfo, + socket::{ + IpVersion, SocketContext, + tcp::{TcpBindOptions, TcpSocketPurpose, VirtualTcpSocketFactory}, + udp::{UdpBindOptions, VirtualUdpSocketFactory}, + }, + tunnel::{SplitTunnel, Tunnel, TunnelError}, +}; + +const MANUAL_PREFLIGHT_DEFAULT_PORT: u16 = 1000; +const MAX_MANUAL_ENDPOINT_HOPS: usize = 16; + +fn manual_default_port(url: &Url) -> u16 { + crate::connectivity::protocol::protocol_default_port(url.scheme()).unwrap_or(11010) +} + +fn is_manual_endpoint_scheme(scheme: &str) -> bool { + matches!(scheme, "http" | "https" | "txt" | "srv") +} + +fn validate_manual_url(url: &Url) -> anyhow::Result<()> { + if ManualTransport::from_url(url).is_ok() || is_manual_endpoint_scheme(url.scheme()) { + Ok(()) + } else { + anyhow::bail!("unsupported core manual connector URL: {url}") + } +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub(crate) enum ManualTransport { + Tcp(TcpSocketPurpose), + Udp(UdpSessionMode), + ByteStream, +} + +impl ManualTransport { + pub(crate) fn from_url(url: &Url) -> anyhow::Result { + match protocol_transport(url.scheme()) { + Some(ProtocolTransport::Tcp) => Ok(Self::Tcp(TcpSocketPurpose::ManualConnect)), + Some(ProtocolTransport::FakeTcp) => Ok(Self::Tcp(TcpSocketPurpose::FakeTcp)), + Some(ProtocolTransport::Udp(mode)) => Ok(Self::Udp(mode)), + None if matches!(url.scheme(), "ring" | "unix") => Ok(Self::ByteStream), + None => anyhow::bail!("unsupported core manual connector URL: {url}"), + } + } + + fn is_udp(self) -> bool { + matches!(self, Self::Udp(_)) + } + + fn supports_interface_bind(self) -> bool { + !matches!( + self, + Self::Tcp(TcpSocketPurpose::FakeTcp) | Self::ByteStream + ) + } +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct ManualInterfaceAddrs { + pub interface_ipv4s: Vec, + pub interface_ipv6s: Vec, + pub public_ipv6: Option, +} + +#[async_trait] +pub trait ManualConnectorHost: VirtualTcpSocketFactory + VirtualUdpSocketFactory { + async fn local_addr_for_remote( + &self, + remote_addr: SocketAddr, + context: SocketContext, + ) -> anyhow::Result; + + async fn interface_addrs(&self) -> anyhow::Result; + + async fn connect_byte_stream( + &self, + url: &Url, + ) -> anyhow::Result::Socket>> { + anyhow::bail!("host does not support external byte stream: {url}") + } +} + +#[async_trait] +pub(crate) trait ManualEndpointResolver: Send + Sync + 'static { + async fn resolve_endpoint(&self, url: &Url) -> anyhow::Result; +} + +/// Connects one endpoint without owning peer-manager retry or admission policy. +/// +/// This is intended for application protocols such as standalone RPC and the +/// configuration client. The host creates transports, while core resolves the +/// endpoint chain and upgrades the resulting socket or session into a tunnel. +pub struct ManualTunnelConnector +where + H: ManualConnectorHost, +{ + host: Arc, + dns: Arc, + endpoint_resolver: Arc, + protocol: Arc::Socket>>, + options: ManualConnectorOptions, + ring_registry: Option>, +} + +impl ManualTunnelConnector +where + H: ManualConnectorHost, +{ + pub(crate) fn new( + host: Arc, + dns: Arc, + endpoint_resolver: Arc, + protocol: Arc::Socket>>, + options: ManualConnectorOptions, + ) -> Self { + Self { + host, + dns, + endpoint_resolver, + protocol, + options, + ring_registry: None, + } + } + + pub(crate) fn with_ring_registry(mut self, ring_registry: Arc) -> Self { + self.ring_registry = Some(ring_registry); + self + } + + pub async fn connect( + &self, + requested_url: Url, + ip_version: IpVersion, + ) -> anyhow::Result> { + validate_manual_url(&requested_url)?; + let endpoint = resolve_manual_endpoint( + self.endpoint_resolver.as_ref(), + convert_idn_to_ascii(requested_url.clone())?, + ) + .await?; + if endpoint.url.scheme() == "ring" { + let registry = self + .ring_registry + .as_ref() + .ok_or_else(|| anyhow::anyhow!("ring registry is not configured"))?; + let tunnel = connect_ring_tunnel(registry, &endpoint.url)?; + return Ok(apply_resolved_endpoint_info( + tunnel, + requested_url, + endpoint.tunnel_prefixes, + )); + } + + if !self.protocol.supports_scheme(endpoint.url.scheme()) { + anyhow::bail!( + "unsupported client protocol upgrader: {}", + endpoint.url.scheme() + ); + } + + let transport = ManualTransport::from_url(&endpoint.url)?; + let connected = if transport == ManualTransport::ByteStream { + ConnectedTransport::ByteStream(self.host.connect_byte_stream(&endpoint.url).await?) + } else { + let remote_addr = resolve_url_addrs( + &endpoint.url, + manual_default_port(&endpoint.url), + self.options.socket_context(transport, ip_version), + self.dns.as_ref(), + ) + .await? + .choose(&mut rand::thread_rng()) + .copied() + .ok_or(TunnelError::NoDnsRecordFound(ip_version))?; + connect_resolved( + self.host.clone(), + transport, + remote_addr, + Vec::new(), + self.options.tcp_bind.clone(), + self.options.udp_bind.clone(), + ) + .await? + }; + + let tunnel = self + .protocol + .upgrade_client(connected, endpoint.url) + .await?; + Ok(apply_resolved_endpoint_info( + tunnel, + requested_url, + endpoint.tunnel_prefixes, + )) + } +} + +struct ResolvedManualTunnel { + inner: Box, + info: TunnelInfo, +} + +impl Tunnel for ResolvedManualTunnel { + fn split(&self) -> SplitTunnel { + self.inner.split() + } + + fn info(&self) -> Option { + Some(self.info.clone()) + } +} + +#[derive(Debug)] +struct ResolvedManualEndpoint { + url: Url, + tunnel_prefixes: Vec, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum ManualConnectorStatus { + Connected, + Disconnected, + Connecting, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct ManualConnectorSnapshot { + pub url: Url, + pub status: ManualConnectorStatus, +} + +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct ManualConnectorOptions { + pub reconnect_interval: Duration, + pub connect_timeout: Duration, + pub endpoint_discovery_timeout: Duration, + pub bind_device: bool, + pub allow_interface_bind: bool, + pub tcp_bind: TcpBindOptions, + pub udp_bind: UdpBindOptions, +} + +impl Default for ManualConnectorOptions { + fn default() -> Self { + Self { + reconnect_interval: Duration::from_secs(1), + connect_timeout: Duration::from_secs(2), + endpoint_discovery_timeout: Duration::from_secs(20), + bind_device: false, + allow_interface_bind: true, + tcp_bind: TcpBindOptions::default(), + udp_bind: UdpBindOptions::direct_connect(), + } + } +} + +impl ManualConnectorOptions { + fn connect_timeout( + &self, + url: &Url, + protocol: &dyn ClientProtocolUpgrader, + ) -> Duration { + if is_manual_endpoint_scheme(url.scheme()) { + self.endpoint_discovery_timeout + } else { + protocol + .connect_timeout(url.scheme()) + .unwrap_or(self.connect_timeout) + } + } + + pub(crate) fn socket_context( + &self, + transport: ManualTransport, + ip_version: IpVersion, + ) -> SocketContext { + let context = match transport { + ManualTransport::Tcp(_) => self.tcp_bind.context.clone(), + ManualTransport::Udp(_) => self.udp_bind.context.clone(), + ManualTransport::ByteStream => SocketContext::default(), + }; + context.with_ip_version(ip_version) + } +} + +struct ManualConnectorData +where + H: ManualConnectorHost, +{ + connectors: DashSet, + reconnecting: DashSet, + removed: DashSet, + state_lock: Mutex<()>, + peer_manager: Weak, + host: Arc, + dns: Arc, + endpoint_resolver: Arc, + protocol: Arc::Socket>>, + ring_registry: Arc, + events: Arc, + options: ManualConnectorOptions, +} + +pub(crate) struct ManualConnectorManager +where + H: ManualConnectorHost, +{ + data: Arc>, + task: Mutex>, +} + +struct ManualConnectorTask { + cancel: CancellationToken, + handle: AbortOnDropHandle<()>, +} + +impl ManualConnectorManager +where + H: ManualConnectorHost, +{ + #[allow(clippy::too_many_arguments)] + pub(crate) fn new( + peer_manager: Arc, + host: Arc, + dns: Arc, + endpoint_resolver: Arc, + protocol: Arc::Socket>>, + ring_registry: Arc, + options: ManualConnectorOptions, + events: Arc, + ) -> Self { + let data = Arc::new(ManualConnectorData { + connectors: DashSet::new(), + reconnecting: DashSet::new(), + removed: DashSet::new(), + state_lock: Mutex::new(()), + peer_manager: Arc::downgrade(&peer_manager), + host, + dns, + endpoint_resolver, + protocol, + ring_registry, + events, + options, + }); + Self { + data, + task: Mutex::new(None), + } + } + + pub fn start(&self) { + let mut task_slot = self.task.lock().unwrap(); + if task_slot + .as_ref() + .is_some_and(|task| !task.handle.is_finished()) + { + return; + } + let cancel = CancellationToken::new(); + task_slot.replace(ManualConnectorTask { + handle: AbortOnDropHandle::new(tokio::spawn(Self::run( + self.data.clone(), + cancel.clone(), + ))), + cancel, + }); + } + + pub async fn stop(&self) { + let task = self.task.lock().unwrap().take(); + if let Some(task) = task { + task.cancel.cancel(); + let _ = task.handle.await; + } + let _state_guard = self.data.state_lock.lock().unwrap(); + restore_interrupted_connectors( + &self.data.connectors, + &self.data.reconnecting, + &self.data.removed, + ); + } + + pub fn add_connector(&self, url: Url) -> anyhow::Result<()> { + validate_manual_url(&url)?; + let _state_guard = self.data.state_lock.lock().unwrap(); + self.data.removed.remove(&url); + if !self.data.reconnecting.contains(&url) { + self.data.connectors.insert(url); + } + Ok(()) + } + + pub fn remove_connector(&self, url: &Url) -> bool { + let _state_guard = self.data.state_lock.lock().unwrap(); + if self.data.connectors.remove(url).is_some() { + tracing::warn!(%url, "manual connector removed"); + return true; + } + if self.data.reconnecting.contains(url) { + self.data.removed.insert(url.clone()); + return true; + } + false + } + + pub fn clear_connectors(&self) { + let _state_guard = self.data.state_lock.lock().unwrap(); + self.data.connectors.clear(); + for url in self.data.reconnecting.iter() { + self.data.removed.insert(url.key().clone()); + } + } + + pub fn list_connectors(&self) -> Vec { + let _state_guard = self.data.state_lock.lock().unwrap(); + let peer_manager = self.data.peer_manager.upgrade(); + let mut snapshots = self + .data + .connectors + .iter() + .map(|entry| { + let url = entry.key().clone(); + let connected = peer_manager + .as_ref() + .is_some_and(|peer_manager| client_url_is_alive(peer_manager, &url)); + ManualConnectorSnapshot { + url, + status: if connected { + ManualConnectorStatus::Connected + } else { + ManualConnectorStatus::Disconnected + }, + } + }) + .collect::>(); + snapshots.extend( + self.data + .reconnecting + .iter() + .map(|entry| ManualConnectorSnapshot { + url: entry.key().clone(), + status: ManualConnectorStatus::Connecting, + }), + ); + snapshots + } + + async fn run(data: Arc>, cancel: CancellationToken) { + let mut interval = crate::foundation::time::interval(data.options.reconnect_interval); + let mut reconnect_tasks = JoinSet::new(); + + loop { + tokio::select! { + _ = cancel.cancelled() => break, + _ = interval.tick() => { + for url in take_dead_connectors_for_reconnect(&data) { + let task_data = data.clone(); + reconnect_tasks.spawn(async move { + let result = reconnect(task_data, url.clone()).await; + (url, result) + }); + } + } + result = reconnect_tasks.join_next(), if !reconnect_tasks.is_empty() => { + let Some(result) = result else { + continue; + }; + match result { + Ok((url, reconnect_result)) => { + tracing::warn!(?url, ?reconnect_result, "manual reconnect task done"); + let _state_guard = data.state_lock.lock().unwrap(); + data.reconnecting.remove(&url); + if data.removed.remove(&url).is_some() { + tracing::warn!(%url, "manual connector removed after reconnect"); + } else { + data.connectors.insert(url); + } + } + Err(error) => { + tracing::error!(?error, "manual reconnect task failed"); + } + } + } + } + } + + reconnect_tasks.abort_all(); + while reconnect_tasks.join_next().await.is_some() {} + + let _state_guard = data.state_lock.lock().unwrap(); + restore_interrupted_connectors(&data.connectors, &data.reconnecting, &data.removed); + } +} + +fn restore_interrupted_connectors( + connectors: &DashSet, + reconnecting: &DashSet, + removed: &DashSet, +) { + let interrupted = reconnecting + .iter() + .map(|entry| entry.key().clone()) + .collect::>(); + for url in interrupted { + reconnecting.remove(&url); + if removed.remove(&url).is_none() { + connectors.insert(url); + } + } +} + +fn take_dead_connectors_for_reconnect(data: &ManualConnectorData) -> BTreeSet +where + H: ManualConnectorHost, +{ + let _state_guard = data.state_lock.lock().unwrap(); + let Some(peer_manager) = data.peer_manager.upgrade() else { + tracing::warn!("peer manager is gone, skip manual reconnect"); + return BTreeSet::new(); + }; + let dead_connectors = data + .connectors + .iter() + .filter_map(|entry| { + let url = entry.key(); + (!client_url_is_alive(&peer_manager, url)).then(|| url.clone()) + }) + .collect::>(); + for url in &dead_connectors { + let removed = data.connectors.remove(url); + debug_assert!(removed.is_some()); + let inserted = data.reconnecting.insert(url.clone()); + debug_assert!(inserted); + } + dead_connectors +} + +fn client_url_is_alive(peer_manager: &PeerManagerCore, url: &Url) -> bool { + peer_manager.get_peer_map().is_client_url_alive(url) + || peer_manager + .get_foreign_network_client() + .is_client_url_alive(url) +} + +async fn resolve_manual_endpoint( + resolver: &dyn ManualEndpointResolver, + requested_url: Url, +) -> anyhow::Result { + let mut url = requested_url; + let mut tunnel_prefixes = Vec::new(); + let mut visited = BTreeSet::new(); + loop { + if !visited.insert(url.clone()) { + anyhow::bail!("manual endpoint resolution cycle detected at {url}"); + } + if ManualTransport::from_url(&url).is_ok() { + return Ok(ResolvedManualEndpoint { + url, + tunnel_prefixes, + }); + } + if !is_manual_endpoint_scheme(url.scheme()) { + anyhow::bail!("unsupported resolved manual connector URL: {url}"); + } + if tunnel_prefixes.len() >= MAX_MANUAL_ENDPOINT_HOPS { + anyhow::bail!("manual endpoint resolution exceeded {MAX_MANUAL_ENDPOINT_HOPS} hops"); + } + tunnel_prefixes.push(url.scheme().to_owned()); + url = convert_idn_to_ascii(resolver.resolve_endpoint(&url).await?)?; + } +} + +fn apply_resolved_endpoint_info( + tunnel: Box, + requested_url: Url, + tunnel_prefixes: Vec, +) -> Box { + if tunnel_prefixes.is_empty() { + return tunnel; + } + let inner_info = tunnel.info().unwrap_or_default(); + let tunnel_type = format!("{}-{}", tunnel_prefixes.join("-"), inner_info.tunnel_type); + Box::new(ResolvedManualTunnel { + inner: tunnel, + info: TunnelInfo { + local_addr: inner_info.local_addr, + remote_addr: Some(requested_url.into()), + resolved_remote_addr: inner_info.resolved_remote_addr.or(inner_info.remote_addr), + tunnel_type, + }, + }) +} + +fn connect_ring_tunnel( + registry: &RingTunnelRegistry, + endpoint: &Url, +) -> anyhow::Result> { + let remote_id = endpoint + .host_str() + .ok_or_else(|| anyhow::anyhow!("ring URL has no peer id: {endpoint}"))? + .parse()?; + Ok(registry.connect(remote_id)?.into_tunnel()) +} + +async fn resolve_reconnect_ip_versions( + url: &Url, + connect_timeout: Duration, + context: SocketContext, + dns: &dyn DnsResolver, +) -> anyhow::Result> { + if matches!( + ManualTransport::from_url(url), + Ok(ManualTransport::ByteStream) + ) { + return Ok(vec![IpVersion::Both]); + } + if matches!(url.scheme(), "txt" | "srv") { + return Ok(vec![IpVersion::Both]); + } + + let addrs = with_timeout_budget( + "resolve", + Instant::now(), + connect_timeout, + resolve_url_addrs( + url, + MANUAL_PREFLIGHT_DEFAULT_PORT, + context.with_ip_version(IpVersion::Both), + dns, + ), + ) + .await?; + tracing::info!(?addrs, %url, "manual preflight resolve done"); + + let mut ip_versions = Vec::new(); + if addrs.iter().any(SocketAddr::is_ipv4) { + ip_versions.push(IpVersion::V4); + } + if addrs.iter().any(SocketAddr::is_ipv6) { + ip_versions.push(IpVersion::V6); + } + Ok(ip_versions) +} + +async fn reconnect(data: Arc>, url: Url) -> anyhow::Result<()> +where + H: ManualConnectorHost, +{ + validate_manual_url(&url)?; + let connect_timeout = data.options.connect_timeout(&url, data.protocol.as_ref()); + tracing::info!(%url, "manual reconnect start"); + let normalized_url = match convert_idn_to_ascii(url.clone()) { + Ok(url) => url, + Err(error) => { + emit_connect_error(&data, &url, IpVersion::Both, &error); + return Err(error); + } + }; + let ip_versions = match resolve_reconnect_ip_versions( + &normalized_url, + connect_timeout, + ManualTransport::from_url(&normalized_url) + .ok() + .map(|transport| data.options.socket_context(transport, IpVersion::Both)) + .unwrap_or_default(), + data.dns.as_ref(), + ) + .await + { + Ok(ip_versions) => ip_versions, + Err(error) => { + emit_connect_error(&data, &url, IpVersion::Both, &error); + return Err(error); + } + }; + + let mut last_error = anyhow::anyhow!("cannot get IP from URL"); + for ip_version in ip_versions { + let started_at = Instant::now(); + match reconnect_with_ip_version( + data.clone(), + url.clone(), + ip_version, + started_at, + connect_timeout, + ) + .await + { + Ok(()) => return Ok(()), + Err(error) => { + emit_connect_error(&data, &url, ip_version, &error); + last_error = error; + } + } + } + Err(last_error) +} + +async fn reconnect_with_ip_version( + data: Arc>, + requested_url: Url, + ip_version: IpVersion, + started_at: Instant, + connect_timeout: Duration, +) -> anyhow::Result<()> +where + H: ManualConnectorHost, +{ + let endpoint = with_timeout_budget( + "discover", + started_at, + connect_timeout, + resolve_manual_endpoint( + data.endpoint_resolver.as_ref(), + convert_idn_to_ascii(requested_url.clone())?, + ), + ) + .await?; + if endpoint.url.scheme() != "ring" && !data.protocol.supports_scheme(endpoint.url.scheme()) { + anyhow::bail!( + "unsupported client protocol upgrader: {}", + endpoint.url.scheme() + ); + } + let transport = (endpoint.url.scheme() != "ring") + .then(|| ManualTransport::from_url(&endpoint.url)) + .transpose()?; + let resolved = match transport { + None | Some(ManualTransport::ByteStream) => None, + Some(transport) => Some( + with_timeout_budget("resolve", started_at, connect_timeout, async { + let peer_manager = data.peer_manager.upgrade().ok_or_else(|| { + anyhow::anyhow!("peer manager is gone, cannot resolve connector") + })?; + let remote_addr = resolve_remote_addr( + peer_manager.as_ref(), + data.host.as_ref(), + data.dns.as_ref(), + &endpoint.url, + manual_default_port(&endpoint.url), + data.options.socket_context(transport, ip_version), + ) + .await?; + let bind_addrs = if data.options.bind_device + && data.options.allow_interface_bind + && transport.supports_interface_bind() + { + collect_bind_addrs( + peer_manager.as_ref(), + data.host.as_ref(), + transport.is_udp(), + remote_addr, + ) + .await? + } else { + Vec::new() + }; + Ok((remote_addr, bind_addrs)) + }) + .await?, + ), + }; + data.events.emit(CoreEvent::ManualConnecting { + url: requested_url.clone(), + }); + + let tunnel = with_timeout_budget("connect", started_at, connect_timeout, async { + if endpoint.url.scheme() == "ring" { + return connect_ring_tunnel(&data.ring_registry, &endpoint.url); + } + let transport = transport.expect("non-Ring endpoint should have a transport"); + let connected = match resolved { + Some((remote_addr, bind_addrs)) => { + connect_resolved( + data.host.clone(), + transport, + remote_addr, + bind_addrs, + data.options.tcp_bind.clone(), + data.options.udp_bind.clone(), + ) + .await? + } + None => { + ConnectedTransport::ByteStream(data.host.connect_byte_stream(&endpoint.url).await?) + } + }; + data.protocol.upgrade_client(connected, endpoint.url).await + }) + .await?; + let tunnel = + apply_resolved_endpoint_info(tunnel, requested_url.clone(), endpoint.tunnel_prefixes); + let peer_manager = data + .peer_manager + .upgrade() + .ok_or_else(|| anyhow::anyhow!("peer manager is gone, cannot reconnect"))?; + let (peer_id, conn_id) = + with_timeout_budget("handshake", started_at, connect_timeout, async move { + peer_manager + .add_client_tunnel_with_peer_id_hint(tunnel, true, None) + .await + .map_err(anyhow::Error::from) + }) + .await?; + tracing::info!(peer_id, %conn_id, %requested_url, "manual reconnect succeeded"); + Ok(()) +} + +pub(crate) async fn connect_resolved( + host: Arc, + transport: ManualTransport, + remote_addr: SocketAddr, + bind_addrs: Vec, + tcp_bind: TcpBindOptions, + udp_bind: UdpBindOptions, +) -> anyhow::Result::Socket>> +where + H: ManualConnectorHost, +{ + match transport { + ManualTransport::Tcp(purpose) => { + transport::connect_tcp(host, remote_addr, bind_addrs, tcp_bind, purpose) + .await + .map(ConnectedTransport::Tcp) + } + ManualTransport::Udp(mode) => { + transport::connect_udp(host, remote_addr, bind_addrs, udp_bind, mode) + .await + .map(ConnectedTransport::Udp) + } + ManualTransport::ByteStream => { + anyhow::bail!("external byte streams do not use an IP transport address") + } + } +} + +pub(crate) async fn resolve_remote_addr( + peer_manager: &PeerManagerCore, + host: &H, + dns: &dyn DnsResolver, + url: &Url, + default_port: u16, + context: SocketContext, +) -> anyhow::Result +where + H: ManualConnectorHost, +{ + let ip_version = context.ip_version; + let addrs = resolve_url_addrs(url, default_port, context.clone(), dns).await?; + let mut usable = Vec::new(); + let mut rejected_reason = None; + for addr in addrs { + let SocketAddr::V6(v6_addr) = addr else { + usable.push(addr); + continue; + }; + if peer_manager.is_easytier_managed_ipv6(v6_addr.ip()).await { + rejected_reason = Some(format!( + "{url} resolves to EasyTier-managed IPv6 {}", + v6_addr.ip() + )); + continue; + } + match host.local_addr_for_remote(addr, context.clone()).await { + Ok(SocketAddr::V6(local_addr)) + if peer_manager.is_easytier_managed_ipv6(local_addr.ip()).await => + { + rejected_reason = Some(format!( + "{url} would use EasyTier-managed IPv6 {} as local source for {v6_addr}", + local_addr.ip() + )); + } + Ok(_) => usable.push(addr), + Err(error) => return Err(error), + } + } + + if usable.is_empty() { + if let Some(reason) = rejected_reason { + anyhow::bail!("{reason}, refusing overlay-backed underlay connection"); + } + return Err(TunnelError::NoDnsRecordFound(ip_version).into()); + } + usable + .choose(&mut rand::thread_rng()) + .copied() + .ok_or_else(|| TunnelError::NoDnsRecordFound(ip_version).into()) +} + +pub(crate) async fn collect_bind_addrs( + peer_manager: &PeerManagerCore, + host: &H, + is_udp: bool, + remote_addr: SocketAddr, +) -> anyhow::Result> +where + H: ManualConnectorHost, +{ + if is_udp && remote_addr.is_ipv6() { + return Ok(Vec::new()); + } + + let addrs = host.interface_addrs().await?; + if remote_addr.is_ipv4() { + return Ok(addrs + .interface_ipv4s + .into_iter() + .map(|addr| SocketAddr::new(IpAddr::V4(addr), 0)) + .collect()); + } + + let mut ipv6s = addrs.interface_ipv6s; + ipv6s.extend(addrs.public_ipv6); + let mut ret = Vec::new(); + for addr in ipv6s { + if !peer_manager.is_easytier_managed_ipv6(&addr).await { + ret.push(SocketAddr::new(IpAddr::V6(addr), 0)); + } + } + Ok(ret) +} + +pub(crate) async fn resolve_url_addrs( + url: &Url, + default_port: u16, + context: SocketContext, + dns: &dyn DnsResolver, +) -> anyhow::Result> { + let ip_version = context.ip_version; + let host = url + .host() + .ok_or_else(|| anyhow::anyhow!("URL has no host: {url}"))?; + let port = url.port().unwrap_or(default_port); + let addrs = match host { + url::Host::Ipv4(addr) => vec![SocketAddr::new(IpAddr::V4(addr), port)], + url::Host::Ipv6(addr) => vec![SocketAddr::new(IpAddr::V6(addr), port)], + url::Host::Domain(host) => dns + .resolve(DnsQuery::new(host, context)) + .await? + .into_iter() + .map(|addr| SocketAddr::new(addr, port)) + .collect(), + }; + let addrs = addrs + .into_iter() + .filter(|addr| match ip_version { + IpVersion::V4 => addr.is_ipv4(), + IpVersion::V6 => addr.is_ipv6(), + IpVersion::Both => true, + }) + .collect::>(); + if addrs.is_empty() { + return Err(TunnelError::NoDnsRecordFound(ip_version).into()); + } + Ok(addrs) +} + +pub(crate) fn convert_idn_to_ascii(mut url: Url) -> anyhow::Result { + if url.is_special() { + return Ok(url); + } + if let Some(domain) = url.domain() { + let domain = percent_decode_str(domain).decode_utf8()?; + let domain = idna::domain_to_ascii(&domain)?; + url.set_host(Some(&domain))?; + } + Ok(url) +} + +fn emit_connect_error( + data: &ManualConnectorData, + url: &Url, + ip_version: IpVersion, + error: &anyhow::Error, +) where + H: ManualConnectorHost, +{ + data.events.emit(CoreEvent::ManualConnectError { + url: url.clone(), + ip_version, + error: format!("{error:#?}"), + }); +} + +async fn with_timeout_budget( + stage: &'static str, + started_at: Instant, + total_timeout: Duration, + future: F, +) -> anyhow::Result +where + F: Future>, +{ + let remaining = total_timeout + .checked_sub(started_at.elapsed()) + .filter(|remaining| !remaining.is_zero()) + .ok_or_else(|| anyhow::anyhow!("{stage} timeout after {:?}", started_at.elapsed()))?; + crate::foundation::time::timeout(remaining, future) + .await + .map_err(|_| anyhow::anyhow!("{stage} timeout after {remaining:?}"))? +} + +#[cfg(test)] +mod tests { + use std::sync::Mutex; + + use crate::socket::udp::UdpSessionProtocol; + + use super::*; + + #[test] + fn idn_normalization_covers_connector_schemes_and_url_round_trips() { + let cases = [ + ("example.com", "example.com"), + ("test.org:8080/path", "test.org:8080/path"), + ("räksmörgås.nu", "xn--rksmrgs-5wao1o.nu"), + ("中文.测试", "xn--fiq228c.xn--0zwm56d"), + ("räksmörgås.nu:8080", "xn--rksmrgs-5wao1o.nu:8080"), + ("例子.测试/path", "xn--fsqu00a.xn--0zwm56d/path"), + ("中文.测试:9000/api", "xn--fiq228c.xn--0zwm56d:9000/api"), + ("räksmörgås.nu:8080/path", "xn--rksmrgs-5wao1o.nu:8080/path"), + ( + "中文.测试:8000/用户/管理", + "xn--fiq228c.xn--0zwm56d:8000/%E7%94%A8%E6%88%B7/%E7%AE%A1%E7%90%86", + ), + ("[2001:db8::1]:8080", "[2001:db8::1]:8080"), + ("[2001:db8::1]/path", "[2001:db8::1]/path"), + ( + "[2001:db8::1]/路径/资源", + "[2001:db8::1]/%E8%B7%AF%E5%BE%84/%E8%B5%84%E6%BA%90", + ), + ]; + let schemes = ["tcp", "udp", "ws", "wss", "wg", "quic", "http", "https"]; + + for (host_part, expected_host_part) in cases { + for scheme in schemes { + for round_trip in [false, true] { + let input = Url::parse(&format!("{scheme}://{host_part}")).unwrap(); + let input = if round_trip { + input.to_string().parse().unwrap() + } else { + input + }; + let actual = convert_idn_to_ascii(input.clone()).unwrap().to_string(); + let mut expected = format!("{scheme}://{expected_host_part}"); + if input.is_special() + && actual.ends_with('/') + && !expected_host_part.ends_with('/') + { + expected.push('/'); + } + assert_eq!(actual, expected, "scheme={scheme}, input={host_part}"); + } + } + } + } + + struct StaticDnsResolver { + ips: Vec, + queries: Mutex>, + } + + struct ChainedEndpointResolver; + + struct CyclingEndpointResolver; + + #[async_trait] + impl ManualEndpointResolver for ChainedEndpointResolver { + async fn resolve_endpoint(&self, url: &Url) -> anyhow::Result { + match url.scheme() { + "http" => Ok("txt://discovery.example".parse().unwrap()), + "txt" => Ok("tcp://peer.example:12000".parse().unwrap()), + scheme => anyhow::bail!("unexpected endpoint scheme: {scheme}"), + } + } + } + + #[async_trait] + impl ManualEndpointResolver for CyclingEndpointResolver { + async fn resolve_endpoint(&self, url: &Url) -> anyhow::Result { + Ok(url.clone()) + } + } + + #[async_trait] + impl DnsResolver for StaticDnsResolver { + async fn resolve(&self, query: DnsQuery) -> anyhow::Result> { + self.queries.lock().unwrap().push(query); + Ok(self.ips.clone()) + } + } + + #[test] + fn manual_transport_maps_ip_protocols_to_their_socket_boundary() { + let cases = [ + ( + "tcp://127.0.0.1:1", + ManualTransport::Tcp(TcpSocketPurpose::ManualConnect), + ), + ( + "ws://127.0.0.1:1", + ManualTransport::Tcp(TcpSocketPurpose::ManualConnect), + ), + ( + "wss://127.0.0.1:1", + ManualTransport::Tcp(TcpSocketPurpose::ManualConnect), + ), + ( + "faketcp://127.0.0.1:1", + ManualTransport::Tcp(TcpSocketPurpose::FakeTcp), + ), + ( + "udp://127.0.0.1:1", + ManualTransport::Udp(UdpSessionMode::EasyTierMux), + ), + ( + "wg://127.0.0.1:1", + ManualTransport::Udp(UdpSessionMode::Classified(UdpSessionProtocol::WireGuard)), + ), + ( + "quic://127.0.0.1:1", + ManualTransport::Udp(UdpSessionMode::Classified(UdpSessionProtocol::Quic)), + ), + ("ring://local", ManualTransport::ByteStream), + ("unix:///tmp/easytier.sock", ManualTransport::ByteStream), + ]; + + for (url, expected) in cases { + assert_eq!( + ManualTransport::from_url(&url.parse().unwrap()).unwrap(), + expected + ); + } + assert!(ManualTransport::from_url(&"http://127.0.0.1:1".parse().unwrap()).is_err()); + assert!(validate_manual_url(&"http://127.0.0.1:1".parse().unwrap()).is_ok()); + } + + #[tokio::test] + async fn manual_endpoint_resolution_preserves_the_discovery_chain() { + let endpoint = resolve_manual_endpoint( + &ChainedEndpointResolver, + "http://discovery.example".parse().unwrap(), + ) + .await + .unwrap(); + + assert_eq!(endpoint.url.as_str(), "tcp://peer.example:12000"); + assert_eq!(endpoint.tunnel_prefixes, ["http", "txt"]); + } + + #[tokio::test] + async fn manual_endpoint_resolution_rejects_cycles() { + let error = resolve_manual_endpoint( + &CyclingEndpointResolver, + "http://discovery.example".parse().unwrap(), + ) + .await + .unwrap_err(); + + assert!(error.to_string().contains("cycle detected")); + } + + #[tokio::test] + async fn txt_discovery_does_not_require_address_records() { + let resolver = StaticDnsResolver { + ips: Vec::new(), + queries: Mutex::new(Vec::new()), + }; + let versions = resolve_reconnect_ip_versions( + &"txt://discovery.example".parse().unwrap(), + Duration::from_secs(1), + SocketContext::default(), + &resolver, + ) + .await + .unwrap(); + + assert_eq!(versions, [IpVersion::Both]); + assert!(resolver.queries.lock().unwrap().is_empty()); + } + + #[tokio::test] + async fn external_byte_stream_does_not_require_address_records() { + let resolver = StaticDnsResolver { + ips: Vec::new(), + queries: Mutex::new(Vec::new()), + }; + let versions = resolve_reconnect_ip_versions( + &"ring://local".parse().unwrap(), + Duration::from_secs(1), + SocketContext::default(), + &resolver, + ) + .await + .unwrap(); + + assert_eq!(versions, [IpVersion::Both]); + assert!(resolver.queries.lock().unwrap().is_empty()); + } + + #[test] + fn external_protocol_and_discovery_timeouts_are_explicit() { + struct Protocol; + + #[async_trait] + impl ClientProtocolUpgrader<()> for Protocol { + fn supports_scheme(&self, scheme: &str) -> bool { + scheme == "external" + } + + fn connect_timeout(&self, scheme: &str) -> Option { + (scheme == "external").then_some(Duration::from_secs(20)) + } + + async fn upgrade_client( + &self, + _connected: ConnectedTransport<()>, + _requested_url: Url, + ) -> anyhow::Result> { + unreachable!() + } + } + + let options = ManualConnectorOptions::default(); + assert_eq!( + options.connect_timeout(&"external://127.0.0.1".parse().unwrap(), &Protocol), + Duration::from_secs(20) + ); + assert_eq!( + options.connect_timeout(&"txt://example.com".parse().unwrap(), &Protocol), + Duration::from_secs(20) + ); + assert_eq!( + options.connect_timeout(&"tcp://127.0.0.1".parse().unwrap(), &Protocol), + Duration::from_secs(2) + ); + } + + #[tokio::test] + async fn resolver_receives_instance_socket_context_and_filters_family() { + let resolver = StaticDnsResolver { + ips: vec![IpAddr::from([127, 0, 0, 1]), Ipv6Addr::LOCALHOST.into()], + queries: Mutex::new(Vec::new()), + }; + let url: Url = "udp://example.com:12000".parse().unwrap(); + + let addrs = resolve_url_addrs( + &url, + crate::connectivity::protocol::protocol_default_port("tcp").unwrap(), + SocketContext::default() + .with_ip_version(IpVersion::V6) + .with_socket_mark(Some(7)), + &resolver, + ) + .await + .unwrap(); + + assert_eq!(addrs, vec!["[::1]:12000".parse().unwrap()]); + assert_eq!( + resolver.queries.lock().unwrap().as_slice(), + &[DnsQuery::new( + "example.com", + SocketContext { + ip_version: IpVersion::V6, + socket_mark: Some(7), + netns: None, + } + )] + ); + } + + #[tokio::test] + async fn timeout_budget_reports_the_active_stage() { + let error = + with_timeout_budget("connect", Instant::now(), Duration::from_millis(1), async { + crate::foundation::time::sleep(Duration::from_millis(20)).await; + Ok::<(), anyhow::Error>(()) + }) + .await + .unwrap_err(); + + assert!(error.to_string().contains("connect timeout after")); + } + + #[test] + fn interrupted_reconnects_return_to_the_pending_set_unless_removed() { + let connectors = DashSet::new(); + let reconnecting = DashSet::new(); + let removed = DashSet::new(); + let retained: Url = "tcp://127.0.0.1:11010".parse().unwrap(); + let deleted: Url = "udp://127.0.0.1:11010".parse().unwrap(); + reconnecting.insert(retained.clone()); + reconnecting.insert(deleted.clone()); + removed.insert(deleted.clone()); + + restore_interrupted_connectors(&connectors, &reconnecting, &removed); + restore_interrupted_connectors(&connectors, &reconnecting, &removed); + + assert!(connectors.contains(&retained)); + assert!(!connectors.contains(&deleted)); + assert!(reconnecting.is_empty()); + assert!(removed.is_empty()); + } +} diff --git a/easytier-core/src/connectivity/mod.rs b/easytier-core/src/connectivity/mod.rs new file mode 100644 index 00000000..4fe9b27c --- /dev/null +++ b/easytier-core/src/connectivity/mod.rs @@ -0,0 +1,37 @@ +//! Portable connection orchestration. + +use std::fmt::Debug; + +use url::Url; + +pub mod composite; +pub mod direct; +pub mod hole_punch; +// Kept public: the host-driven adapter chain is WASI-only production code +// (cfg(target_os = "wasi")), so crate-private visibility would surface +// dead-code warnings on host builds for code that is live on WASI. +pub mod connector_host; +pub mod manual; +pub mod protocol; +pub mod stun; +pub mod transport; + +/// Supplies the URLs of the instance's currently running listeners. +/// +/// The listener layer's running-listener registry implements this seam. +/// Connectors use it to avoid dialing addresses that would hairpin back +/// into one of their own listeners, so connectivity depends on this narrow +/// query rather than on the listener module's concrete registry type. +pub trait LocalListenerUrls: Debug + Send + Sync + 'static { + fn local_listener_urls(&self) -> Vec; +} + +/// Empty [`LocalListenerUrls`] for connectors that track no listeners. +#[derive(Debug, Default)] +pub struct NoLocalListeners; + +impl LocalListenerUrls for NoLocalListeners { + fn local_listener_urls(&self) -> Vec { + Vec::new() + } +} diff --git a/easytier-core/src/connectivity/protocol/mod.rs b/easytier-core/src/connectivity/protocol/mod.rs new file mode 100644 index 00000000..f436e544 --- /dev/null +++ b/easytier-core/src/connectivity/protocol/mod.rs @@ -0,0 +1,811 @@ +use std::{marker::PhantomData, num::NonZeroUsize, sync::Arc, time::Duration}; + +use async_trait::async_trait; +use tokio::sync::{OwnedSemaphorePermit, Semaphore}; +use url::Url; + +use crate::{ + socket::{tcp::VirtualTcpSocket, udp::UdpSession}, + tunnel::Tunnel, +}; + +use super::transport::{ConnectedTransport, UdpSessionMode}; + +pub mod raw; + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub(crate) enum ProtocolTransport { + Tcp, + FakeTcp, + Udp(UdpSessionMode), +} + +pub(crate) const fn protocol_transport(scheme: &str) -> Option { + match scheme.as_bytes() { + b"tcp" | b"ws" | b"wss" => Some(ProtocolTransport::Tcp), + b"faketcp" => Some(ProtocolTransport::FakeTcp), + b"udp" => Some(ProtocolTransport::Udp(UdpSessionMode::EasyTierMux)), + b"wg" => Some(ProtocolTransport::Udp(UdpSessionMode::Classified( + crate::socket::udp::UdpSessionProtocol::WireGuard, + ))), + b"quic" => Some(ProtocolTransport::Udp(UdpSessionMode::Classified( + crate::socket::udp::UdpSessionProtocol::Quic, + ))), + _ => None, + } +} + +pub(crate) const fn protocol_uses_udp(scheme: &str) -> bool { + matches!(protocol_transport(scheme), Some(ProtocolTransport::Udp(_))) +} + +/// Returns the listener-port offset used when expanding a single base port +/// into EasyTier's protocol-specific listener set. +pub const fn protocol_port_offset(scheme: &str) -> Option { + match scheme.as_bytes() { + b"tcp" | b"udp" => Some(0), + b"wg" | b"ws" => Some(1), + b"quic" | b"wss" => Some(2), + b"faketcp" => Some(3), + _ => None, + } +} + +/// Returns the default port for a concrete EasyTier IP protocol. +pub const fn protocol_default_port(scheme: &str) -> Option { + match scheme.as_bytes() { + b"ws" => Some(80), + b"wss" => Some(443), + _ => match protocol_port_offset(scheme) { + Some(offset) => Some(11010 + offset), + None => None, + }, + } +} + +#[async_trait] +pub trait ClientProtocolUpgrader: Send + Sync + 'static { + fn supports_scheme(&self, scheme: &str) -> bool; + + fn connect_timeout(&self, _scheme: &str) -> Option { + None + } + + async fn upgrade_client( + &self, + connected: ConnectedTransport, + requested_url: Url, + ) -> anyhow::Result>; +} + +#[async_trait] +pub trait ServerTunnelAcceptor: Send + 'static { + async fn accept(&mut self) -> anyhow::Result>; +} + +pub enum ServerProtocolUpgrade { + Tunnel(Box), + Acceptor(Box), +} + +pub struct ServerProtocolAdmission { + active_session: OwnedSemaphorePermit, + handshake_slots: Arc, +} + +impl ServerProtocolAdmission { + pub fn into_parts(self) -> (OwnedSemaphorePermit, Arc) { + (self.active_session, self.handshake_slots) + } +} + +pub struct ServerProtocolAdmissionController { + active_sessions: Arc, + handshake_slots: Arc, +} + +impl ServerProtocolAdmissionController { + pub fn new(max_active_sessions: usize, max_in_flight_handshakes: usize) -> Self { + Self { + active_sessions: Arc::new(Semaphore::new(max_active_sessions)), + handshake_slots: Arc::new(Semaphore::new(max_in_flight_handshakes)), + } + } + + pub fn try_admit(&self) -> Option { + Some(ServerProtocolAdmission { + active_session: self.active_sessions.clone().try_acquire_owned().ok()?, + handshake_slots: self.handshake_slots.clone(), + }) + } + + pub fn quic() -> Self { + Self::new(1024, 128) + } +} + +#[async_trait] +pub trait ServerProtocolUpgrader: Send + Sync + 'static { + fn supports_scheme(&self, scheme: &str) -> bool; + + fn max_pending_tcp_upgrades(&self, _scheme: &str) -> Option { + None + } + + async fn upgrade_tcp( + &self, + socket: TcpSocket, + local_url: Url, + ) -> anyhow::Result; + + async fn upgrade_udp( + &self, + session: UdpSession, + local_url: Url, + admission: Option, + ) -> anyhow::Result; + + async fn upgrade_byte_stream( + &self, + socket: TcpSocket, + local_url: Url, + remote_url: Option, + ) -> anyhow::Result; +} + +#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)] +pub struct CoreClientProtocolConfig { + pub unix: bool, + pub faketcp: bool, +} + +/// Owns portable client protocol dispatch and delegates only protocol engines +/// that are not yet available in core. +pub struct CoreClientProtocolUpgrader { + config: CoreClientProtocolConfig, + external: Option>>, +} + +impl CoreClientProtocolUpgrader { + pub fn new(config: CoreClientProtocolConfig) -> Self { + Self { + config, + external: None, + } + } + + pub fn with_external( + config: CoreClientProtocolConfig, + external: Arc>, + ) -> Self { + Self { + config, + external: Some(external), + } + } +} + +#[async_trait] +impl ClientProtocolUpgrader for CoreClientProtocolUpgrader +where + TcpSocket: VirtualTcpSocket, +{ + fn supports_scheme(&self, scheme: &str) -> bool { + match scheme { + "tcp" | "udp" | "ring" => true, + "unix" => self.config.unix, + "faketcp" => self.config.faketcp, + _ => self + .external + .as_ref() + .is_some_and(|external| external.supports_scheme(scheme)), + } + } + + fn connect_timeout(&self, scheme: &str) -> Option { + self.external + .as_ref() + .filter(|external| external.supports_scheme(scheme)) + .and_then(|external| external.connect_timeout(scheme)) + } + + async fn upgrade_client( + &self, + connected: ConnectedTransport, + requested_url: Url, + ) -> anyhow::Result> { + match requested_url.scheme() { + "tcp" => match connected { + ConnectedTransport::Tcp(socket) => { + Ok(raw::upgrade_connected_tcp(socket, requested_url)?) + } + ConnectedTransport::Udp(_) | ConnectedTransport::ByteStream(_) => { + anyhow::bail!("TCP protocol requires a TCP transport") + } + }, + "udp" => match connected { + ConnectedTransport::Udp(session) => { + Ok(raw::upgrade_connected_udp(session, requested_url)?) + } + ConnectedTransport::Tcp(_) | ConnectedTransport::ByteStream(_) => { + anyhow::bail!("UDP protocol requires a UDP session") + } + }, + "ring" => upgrade_byte_stream(connected), + "unix" if self.config.unix => upgrade_byte_stream(connected), + "faketcp" if self.config.faketcp => match connected { + ConnectedTransport::Tcp(socket) => { + Ok(raw::upgrade_connected_tcp(socket, requested_url)?) + } + ConnectedTransport::Udp(_) | ConnectedTransport::ByteStream(_) => { + anyhow::bail!("FakeTCP protocol requires a TCP transport") + } + }, + "unix" | "faketcp" => anyhow::bail!( + "unsupported client protocol upgrader: {}", + requested_url.scheme() + ), + scheme => { + let Some(external) = &self.external else { + anyhow::bail!("unsupported client protocol upgrader: {scheme}"); + }; + if !external.supports_scheme(scheme) { + anyhow::bail!("unsupported client protocol upgrader: {scheme}"); + } + external.upgrade_client(connected, requested_url).await + } + } + } +} + +fn upgrade_byte_stream( + connected: ConnectedTransport, +) -> anyhow::Result> +where + TcpSocket: VirtualTcpSocket, +{ + match connected { + ConnectedTransport::ByteStream(stream) => Ok(raw::upgrade_connected_byte_stream(stream)?), + ConnectedTransport::Tcp(_) | ConnectedTransport::Udp(_) => { + anyhow::bail!("external protocol requires a host-created byte stream") + } + } +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)] +pub struct CoreServerProtocolConfig { + pub unix: bool, + pub faketcp: bool, +} + +/// Owns portable server protocol dispatch and delegates only protocol engines +/// that are not available in core. +pub struct CoreServerProtocolUpgrader { + config: CoreServerProtocolConfig, + external: Option>>, + tcp_socket: PhantomData TcpSocket>, +} + +impl CoreServerProtocolUpgrader { + pub fn new(config: CoreServerProtocolConfig) -> Self { + Self { + config, + external: None, + tcp_socket: PhantomData, + } + } + + pub fn with_external( + config: CoreServerProtocolConfig, + external: Arc>, + ) -> Self { + Self { + config, + external: Some(external), + tcp_socket: PhantomData, + } + } + + fn supports_core_scheme(&self, scheme: &str) -> Option { + match scheme { + "tcp" | "udp" | "ring" => Some(true), + "unix" => Some(self.config.unix), + "faketcp" => Some(self.config.faketcp), + _ => None, + } + } + + fn external(&self, scheme: &str) -> anyhow::Result<&dyn ServerProtocolUpgrader> { + let external = self + .external + .as_deref() + .ok_or_else(|| anyhow::anyhow!("unsupported server protocol upgrader: {scheme}"))?; + if !external.supports_scheme(scheme) { + anyhow::bail!("unsupported server protocol upgrader: {scheme}"); + } + Ok(external) + } +} + +#[async_trait] +impl ServerProtocolUpgrader for CoreServerProtocolUpgrader +where + TcpSocket: VirtualTcpSocket, +{ + fn supports_scheme(&self, scheme: &str) -> bool { + self.supports_core_scheme(scheme).unwrap_or_else(|| { + self.external + .as_ref() + .is_some_and(|external| external.supports_scheme(scheme)) + }) + } + + fn max_pending_tcp_upgrades(&self, scheme: &str) -> Option { + self.external + .as_ref() + .filter(|external| external.supports_scheme(scheme)) + .and_then(|external| external.max_pending_tcp_upgrades(scheme)) + } + + async fn upgrade_tcp( + &self, + socket: TcpSocket, + local_url: Url, + ) -> anyhow::Result { + match local_url.scheme() { + "tcp" | "faketcp" => Ok(ServerProtocolUpgrade::Tunnel( + upgrade_accepted_tcp(socket, local_url, self.config).await?, + )), + "udp" | "wg" | "quic" => { + anyhow::bail!("{} protocol requires a UDP session", local_url.scheme()) + } + "ring" | "unix" => { + anyhow::bail!("{} protocol requires a byte stream", local_url.scheme()) + } + scheme => self.external(scheme)?.upgrade_tcp(socket, local_url).await, + } + } + + async fn upgrade_udp( + &self, + session: UdpSession, + local_url: Url, + admission: Option, + ) -> anyhow::Result { + match local_url.scheme() { + "udp" => Ok(ServerProtocolUpgrade::Tunnel(upgrade_accepted_udp( + session, &local_url, + )?)), + "tcp" | "faketcp" => { + anyhow::bail!("{} protocol requires a TCP transport", local_url.scheme()) + } + "ring" | "unix" => { + anyhow::bail!("{} protocol requires a byte stream", local_url.scheme()) + } + scheme => { + self.external(scheme)? + .upgrade_udp(session, local_url, admission) + .await + } + } + } + + async fn upgrade_byte_stream( + &self, + socket: TcpSocket, + local_url: Url, + remote_url: Option, + ) -> anyhow::Result { + match local_url.scheme() { + "ring" => Ok(ServerProtocolUpgrade::Tunnel( + raw::upgrade_accepted_byte_stream(socket, local_url, remote_url)?, + )), + "unix" if self.config.unix => Ok(ServerProtocolUpgrade::Tunnel( + raw::upgrade_accepted_byte_stream(socket, local_url, remote_url)?, + )), + "tcp" | "faketcp" => { + anyhow::bail!("{} protocol requires a TCP transport", local_url.scheme()) + } + "udp" | "wg" | "quic" => { + anyhow::bail!("{} protocol requires a UDP session", local_url.scheme()) + } + "unix" => anyhow::bail!("unsupported server protocol upgrader: unix"), + scheme => { + self.external(scheme)? + .upgrade_byte_stream(socket, local_url, remote_url) + .await + } + } + } +} + +pub(crate) async fn upgrade_accepted_tcp( + socket: TcpSocket, + local_url: Url, + config: CoreServerProtocolConfig, +) -> anyhow::Result> +where + TcpSocket: VirtualTcpSocket, +{ + match local_url.scheme() { + "tcp" => Ok(raw::upgrade_accepted_tcp_with_local_url(socket, local_url)?), + "faketcp" if config.faketcp => { + Ok(raw::upgrade_accepted_tcp_with_local_url(socket, local_url)?) + } + scheme => anyhow::bail!("unsupported TCP listener protocol: {scheme}"), + } +} + +pub(crate) fn upgrade_accepted_udp( + session: UdpSession, + local_url: &Url, +) -> anyhow::Result> { + match local_url.scheme() { + "udp" => Ok(raw::upgrade_accepted_udp_with_local_url( + session, + local_url.clone(), + )?), + scheme => anyhow::bail!("unsupported UDP listener protocol: {scheme}"), + } +} + +#[cfg(test)] +mod tests { + use std::{ + io, + net::SocketAddr, + pin::Pin, + sync::Arc, + task::{Context, Poll}, + }; + + use tokio::io::{AsyncRead, AsyncWrite, ReadBuf}; + + use super::*; + use crate::socket::udp::{UdpSessionKind, VirtualUdpSocket}; + + #[test] + fn protocol_port_metadata_is_authoritative_for_all_ip_protocols() { + let cases = [ + ("tcp", 0, 11010), + ("udp", 0, 11010), + ("wg", 1, 11011), + ("quic", 2, 11012), + ("ws", 1, 80), + ("wss", 2, 443), + ("faketcp", 3, 11013), + ]; + + for (scheme, offset, port) in cases { + assert_eq!(protocol_port_offset(scheme), Some(offset)); + assert_eq!(protocol_default_port(scheme), Some(port)); + } + assert_eq!(protocol_default_port("ring"), None); + } + + struct MockTcpSocket; + + impl AsyncRead for MockTcpSocket { + fn poll_read( + self: Pin<&mut Self>, + _cx: &mut Context<'_>, + _buf: &mut ReadBuf<'_>, + ) -> Poll> { + Poll::Pending + } + } + + impl AsyncWrite for MockTcpSocket { + fn poll_write( + self: Pin<&mut Self>, + _cx: &mut Context<'_>, + _buf: &[u8], + ) -> Poll> { + Poll::Pending + } + + fn poll_flush(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll> { + Poll::Ready(Ok(())) + } + + fn poll_shutdown( + self: Pin<&mut Self>, + _cx: &mut Context<'_>, + ) -> Poll> { + Poll::Ready(Ok(())) + } + } + + impl VirtualTcpSocket for MockTcpSocket { + fn local_addr(&self) -> io::Result { + Ok("127.0.0.1:1000".parse().unwrap()) + } + + fn peer_addr(&self) -> io::Result { + Ok("127.0.0.1:2000".parse().unwrap()) + } + } + + struct MockUdpSocket { + local_addr: SocketAddr, + } + + #[async_trait] + impl VirtualUdpSocket for MockUdpSocket { + fn local_addr(&self) -> io::Result { + Ok(self.local_addr) + } + + async fn send_to(&self, data: &[u8], _addr: SocketAddr) -> io::Result { + Ok(data.len()) + } + + async fn recv_from(&self, _buf: &mut [u8]) -> io::Result<(usize, SocketAddr)> { + std::future::pending().await + } + } + + struct MockExternalUpgrader; + + #[async_trait] + impl ClientProtocolUpgrader for MockExternalUpgrader { + fn supports_scheme(&self, _scheme: &str) -> bool { + true + } + + async fn upgrade_client( + &self, + _connected: ConnectedTransport, + _requested_url: Url, + ) -> anyhow::Result> { + anyhow::bail!("external protocol invoked") + } + } + + #[async_trait] + impl ServerProtocolUpgrader for MockExternalUpgrader { + fn supports_scheme(&self, scheme: &str) -> bool { + matches!(scheme, "external" | "ring" | "unix") + } + + async fn upgrade_tcp( + &self, + _socket: MockTcpSocket, + _local_url: Url, + ) -> anyhow::Result { + anyhow::bail!("external server protocol invoked") + } + + async fn upgrade_udp( + &self, + _session: UdpSession, + _local_url: Url, + _admission: Option, + ) -> anyhow::Result { + anyhow::bail!("external server UDP protocol invoked") + } + + async fn upgrade_byte_stream( + &self, + _socket: MockTcpSocket, + _local_url: Url, + _remote_url: Option, + ) -> anyhow::Result { + anyhow::bail!("external server byte-stream protocol invoked") + } + } + + #[tokio::test] + async fn default_core_upgrader_rejects_external_and_mismatched_protocols() { + let upgrader = + CoreClientProtocolUpgrader::::new(CoreClientProtocolConfig::default()); + + assert!(upgrader.supports_scheme("ring")); + assert!(!upgrader.supports_scheme("unix")); + + let unsupported = upgrader + .upgrade_client( + ConnectedTransport::Tcp(MockTcpSocket), + "ws://127.0.0.1:2000".parse().unwrap(), + ) + .await; + assert!(unsupported.is_err()); + + let mismatched = upgrader + .upgrade_client( + ConnectedTransport::Tcp(MockTcpSocket), + "udp://127.0.0.1:2000".parse().unwrap(), + ) + .await; + assert!(mismatched.is_err()); + } + + #[tokio::test] + async fn core_upgrader_owns_builtin_capabilities_and_delegates_external_protocols() { + let upgrader = CoreClientProtocolUpgrader::with_external( + CoreClientProtocolConfig { + unix: false, + faketcp: false, + }, + Arc::new(MockExternalUpgrader), + ); + + assert!(upgrader.supports_scheme("tcp")); + assert!(upgrader.supports_scheme("ring")); + assert!(upgrader.supports_scheme("ws")); + assert!(!upgrader.supports_scheme("unix")); + assert!(!upgrader.supports_scheme("faketcp")); + + let external = upgrader + .upgrade_client( + ConnectedTransport::Tcp(MockTcpSocket), + "ws://127.0.0.1:2000".parse().unwrap(), + ) + .await + .unwrap_err(); + assert_eq!(external.to_string(), "external protocol invoked"); + + let disabled_builtin = upgrader + .upgrade_client( + ConnectedTransport::Tcp(MockTcpSocket), + "faketcp://127.0.0.1:2000".parse().unwrap(), + ) + .await + .unwrap_err(); + assert!( + disabled_builtin + .to_string() + .contains("unsupported client protocol upgrader") + ); + + let mismatched_udp = upgrader + .upgrade_client( + ConnectedTransport::Tcp(MockTcpSocket), + "udp://127.0.0.1:2000".parse().unwrap(), + ) + .await + .unwrap_err(); + assert_eq!( + mismatched_udp.to_string(), + "UDP protocol requires a UDP session" + ); + + let mismatched_tcp = upgrader + .upgrade_client( + ConnectedTransport::ByteStream(super::super::transport::ConnectedByteStream::new( + MockTcpSocket, + None, + "ring://remote".parse().unwrap(), + None, + )), + "tcp://127.0.0.1:2000".parse().unwrap(), + ) + .await + .unwrap_err(); + assert_eq!( + mismatched_tcp.to_string(), + "TCP protocol requires a TCP transport" + ); + } + + #[tokio::test] + async fn core_server_dispatches_raw_tcp_and_enforces_host_capabilities() { + let tunnel = upgrade_accepted_tcp( + MockTcpSocket, + "tcp://0.0.0.0:2000".parse().unwrap(), + CoreServerProtocolConfig::default(), + ) + .await + .unwrap(); + assert_eq!( + tunnel.info().unwrap().local_addr.unwrap().url, + "tcp://0.0.0.0:2000" + ); + + let disabled = upgrade_accepted_tcp( + MockTcpSocket, + "ws://0.0.0.0:2000".parse().unwrap(), + CoreServerProtocolConfig::default(), + ) + .await + .unwrap_err(); + assert_eq!( + disabled.to_string(), + "unsupported TCP listener protocol: ws" + ); + } + + #[tokio::test] + async fn core_server_raw_udp_preserves_explicit_listener_url() { + let local_url: Url = "udp://listener.example:1000/path?bind_device=eth0" + .parse() + .unwrap(); + let session = UdpSession::identity_standalone( + Arc::new(MockUdpSocket { + local_addr: "127.0.0.1:1000".parse().unwrap(), + }), + "127.0.0.1:2000".parse().unwrap(), + UdpSessionKind::EasyTierMux, + ) + .unwrap(); + + let tunnel = upgrade_accepted_udp(session, &local_url).unwrap(); + + assert_eq!( + tunnel.info().unwrap().local_addr.unwrap().url, + local_url.as_str() + ); + } + + #[tokio::test] + async fn core_server_upgrader_owns_builtin_dispatch_and_delegates_external_protocols() { + let upgrader = CoreServerProtocolUpgrader::with_external( + CoreServerProtocolConfig::default(), + Arc::new(MockExternalUpgrader), + ); + + assert!(upgrader.supports_scheme("tcp")); + assert!(upgrader.supports_scheme("udp")); + assert!(!upgrader.supports_scheme("ws")); + assert!(!upgrader.supports_scheme("quic")); + assert!(upgrader.supports_scheme("ring")); + assert!(!upgrader.supports_scheme("unix")); + assert!(upgrader.supports_scheme("external")); + + let external = upgrader + .upgrade_tcp(MockTcpSocket, "external://0.0.0.0:2000".parse().unwrap()) + .await + .err() + .unwrap(); + assert_eq!(external.to_string(), "external server protocol invoked"); + + let mismatched = upgrader + .upgrade_tcp(MockTcpSocket, "quic://0.0.0.0:2000".parse().unwrap()) + .await + .err() + .unwrap(); + assert_eq!( + mismatched.to_string(), + "quic protocol requires a UDP session" + ); + + let wrong_ring_transport = upgrader + .upgrade_tcp(MockTcpSocket, "ring://local".parse().unwrap()) + .await + .err() + .unwrap(); + assert_eq!( + wrong_ring_transport.to_string(), + "ring protocol requires a byte stream" + ); + + let disabled_unix = upgrader + .upgrade_byte_stream( + MockTcpSocket, + "unix:///tmp/easytier.sock".parse().unwrap(), + None, + ) + .await + .err() + .unwrap(); + assert_eq!( + disabled_unix.to_string(), + "unsupported server protocol upgrader: unix" + ); + } + + #[test] + fn server_protocol_admission_is_scoped_to_its_controller() { + let controller = ServerProtocolAdmissionController::new(1, 2); + let admission = controller.try_admit().unwrap(); + assert!(controller.try_admit().is_none()); + + let (active_session, handshake_slots) = admission.into_parts(); + assert_eq!(handshake_slots.available_permits(), 2); + drop(active_session); + assert!(controller.try_admit().is_some()); + + let other = ServerProtocolAdmissionController::new(1, 1); + assert!(other.try_admit().is_some()); + } +} diff --git a/easytier-core/src/connectivity/protocol/raw.rs b/easytier-core/src/connectivity/protocol/raw.rs new file mode 100644 index 00000000..8c05ab8e --- /dev/null +++ b/easytier-core/src/connectivity/protocol/raw.rs @@ -0,0 +1,867 @@ +use std::{fmt, net::SocketAddr, sync::Arc}; + +use async_trait::async_trait; +use rand::seq::SliceRandom as _; +use url::Url; + +use crate::{ + connectivity::{ + manual::resolve_url_addrs, + transport::{self, ConnectedByteStream, ConnectedUdpSession, UdpSessionMode}, + }, + host::dns::DnsResolver, + proto::common::TunnelInfo, + socket::{ + IpVersion, ListenerConnectionCounter, SocketListener, + tcp::{ + TcpBindOptions, TcpListenOptions, TcpSocketListener, TcpSocketPurpose, + VirtualTcpListenerFactory, VirtualTcpSocket, VirtualTcpSocketFactory, + }, + udp::{ + UdpBindOptions, UdpSession, UdpSessionAcceptKind, UdpSessionListenRequest, + UdpSessionSocket, UdpSessionSocketListener, VirtualUdpSocketFactory, + }, + }, + tunnel::{Tunnel, TunnelError, tcp::TcpTunnelUpgrader, udp::UdpTunnelUpgrader}, +}; + +use super::protocol_default_port; + +const BYTE_STREAM_MAX_PACKET_SIZE: usize = 4096; +const TCP_DEFAULT_PORT: u16 = protocol_default_port("tcp").expect("tcp must have a default port"); +const UDP_DEFAULT_PORT: u16 = protocol_default_port("udp").expect("udp must have a default port"); + +#[async_trait] +#[auto_impl::auto_impl(Box, Arc)] +pub trait TunnelDialer: Send + Sync + 'static { + async fn connect(&self) -> anyhow::Result>; + + fn remote_url(&self) -> Url; +} + +/// Core-owned raw TCP Tunnel connector over an injected socket factory. +pub struct TcpTunnelDialer +where + F: VirtualTcpSocketFactory, +{ + remote_url: Url, + factory: Arc, + dns: Arc, + ip_version: IpVersion, + bind: TcpBindOptions, +} + +impl TcpTunnelDialer +where + F: VirtualTcpSocketFactory, +{ + pub fn new(remote_url: Url, factory: Arc, dns: Arc) -> Self { + Self { + remote_url, + factory, + dns, + ip_version: IpVersion::Both, + bind: TcpBindOptions::default(), + } + } + + pub fn with_ip_version(mut self, ip_version: IpVersion) -> Self { + self.ip_version = ip_version; + self + } + + pub fn with_bind(mut self, bind: TcpBindOptions) -> Self { + self.bind = bind; + self + } +} + +#[async_trait] +impl TunnelDialer for TcpTunnelDialer +where + F: VirtualTcpSocketFactory, +{ + async fn connect(&self) -> anyhow::Result> { + if self.remote_url.scheme() != "tcp" { + anyhow::bail!("raw TCP dialer requires tcp URL: {}", self.remote_url); + } + let remote_addr = resolve_url_addrs( + &self.remote_url, + TCP_DEFAULT_PORT, + self.bind.context.clone().with_ip_version(self.ip_version), + self.dns.as_ref(), + ) + .await? + .choose(&mut rand::thread_rng()) + .copied() + .ok_or(TunnelError::NoDnsRecordFound(self.ip_version))?; + let socket = transport::connect_tcp( + self.factory.clone(), + remote_addr, + Vec::new(), + self.bind.clone(), + TcpSocketPurpose::ManualConnect, + ) + .await?; + Ok(upgrade_connected_tcp(socket, self.remote_url.clone())?) + } + + fn remote_url(&self) -> Url { + self.remote_url.clone() + } +} + +/// Core-owned raw TCP Tunnel listener over an injected listener factory. +pub struct TcpTunnelListener +where + F: VirtualTcpListenerFactory, +{ + inner: TcpSocketListener, +} + +impl TcpTunnelListener +where + F: VirtualTcpListenerFactory, +{ + pub fn new(local_addr: SocketAddr, factory: Arc) -> Self { + let bind = TcpBindOptions::default() + .with_local_addr(Some(local_addr)) + .with_only_v6(true); + Self::new_with_bind(local_addr, bind, factory) + } + + pub fn new_with_bind(local_addr: SocketAddr, bind: TcpBindOptions, factory: Arc) -> Self { + Self { + inner: TcpSocketListener::new_with_options( + socket_url("tcp", local_addr), + TcpListenOptions::manual_connect(local_addr).with_bind(bind), + factory, + ), + } + } +} + +impl fmt::Debug for TcpTunnelListener +where + F: VirtualTcpListenerFactory, +{ + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter + .debug_struct("TcpTunnelListener") + .field("inner", &self.inner) + .finish() + } +} + +#[async_trait] +impl SocketListener for TcpTunnelListener +where + F: VirtualTcpListenerFactory, +{ + type Accepted = Box; + + async fn listen(&mut self) -> anyhow::Result<()> { + self.inner.listen().await + } + + async fn accept(&mut self) -> anyhow::Result { + let local_url = self.inner.local_url(); + let socket = self.inner.accept().await?; + Ok(upgrade_accepted_tcp_with_local_url(socket, local_url)?) + } + + fn local_url(&self) -> Url { + self.inner.local_url() + } +} + +/// Core-owned raw UDP Tunnel connector over an injected socket factory. +pub struct UdpTunnelDialer +where + H: VirtualUdpSocketFactory, +{ + remote_url: Url, + host: Arc, + dns: Arc, + ip_version: IpVersion, + bind_addrs: Vec, + bind: UdpBindOptions, +} + +impl UdpTunnelDialer +where + H: VirtualUdpSocketFactory, +{ + pub fn new(remote_url: Url, host: Arc, dns: Arc) -> Self { + Self { + remote_url, + host, + dns, + ip_version: IpVersion::Both, + bind_addrs: Vec::new(), + bind: UdpBindOptions::direct_connect(), + } + } + + pub fn with_ip_version(mut self, ip_version: IpVersion) -> Self { + self.ip_version = ip_version; + self + } + + pub fn with_bind_addrs(mut self, bind_addrs: Vec) -> Self { + self.bind_addrs = bind_addrs; + self + } + + pub fn with_bind(mut self, bind: UdpBindOptions) -> Self { + self.bind = bind; + self + } +} + +#[async_trait] +impl TunnelDialer for UdpTunnelDialer +where + H: VirtualUdpSocketFactory, +{ + async fn connect(&self) -> anyhow::Result> { + if self.remote_url.scheme() != "udp" { + anyhow::bail!("raw UDP dialer requires udp URL: {}", self.remote_url); + } + let remote_addr = resolve_url_addrs( + &self.remote_url, + UDP_DEFAULT_PORT, + self.bind.context.clone().with_ip_version(self.ip_version), + self.dns.as_ref(), + ) + .await? + .choose(&mut rand::thread_rng()) + .copied() + .ok_or(TunnelError::NoDnsRecordFound(self.ip_version))?; + let bind_addrs = udp_bind_addrs_for_remote(remote_addr, &self.bind_addrs); + let connected = transport::connect_udp( + self.host.clone(), + remote_addr, + bind_addrs, + self.bind.clone(), + UdpSessionMode::EasyTierMux, + ) + .await?; + Ok(upgrade_connected_udp(connected, self.remote_url.clone())?) + } + + fn remote_url(&self) -> Url { + self.remote_url.clone() + } +} + +/// Core-owned raw UDP Tunnel listener over an injected socket factory. +pub struct UdpTunnelListener +where + H: VirtualUdpSocketFactory, +{ + inner: UdpSessionSocketListener, +} + +impl UdpTunnelListener +where + H: VirtualUdpSocketFactory, +{ + pub fn new(local_url: Url, local_addr: SocketAddr, host: Arc) -> Self { + Self { + inner: UdpSessionSocketListener::new(local_url, local_addr, host), + } + } + + pub fn new_with_request( + local_url: Url, + request: UdpSessionListenRequest, + host: Arc, + ) -> Self { + Self { + inner: UdpSessionSocketListener::new_with_request( + local_url, + request, + UdpSessionAcceptKind::EasyTierMux, + host, + ), + } + } +} + +impl fmt::Debug for UdpTunnelListener +where + H: VirtualUdpSocketFactory, +{ + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter + .debug_struct("UdpTunnelListener") + .field("inner", &self.inner) + .finish() + } +} + +#[async_trait] +impl SocketListener for UdpTunnelListener +where + H: VirtualUdpSocketFactory, +{ + type Accepted = Box; + + async fn listen(&mut self) -> anyhow::Result<()> { + let local_url = self.inner.local_url(); + if local_url.scheme() != "udp" { + anyhow::bail!("raw UDP listener requires udp URL: {local_url}"); + } + self.inner.listen().await + } + + async fn accept(&mut self) -> anyhow::Result { + let local_url = self.inner.local_url(); + Ok(upgrade_accepted_udp_with_local_url( + self.inner.accept().await?, + local_url, + )?) + } + + fn local_url(&self) -> Url { + self.inner.local_url() + } + + fn connection_counter(&self) -> Arc { + self.inner.connection_counter() + } +} + +pub(crate) fn upgrade_connected_byte_stream( + connected: ConnectedByteStream, +) -> Result, TunnelError> +where + S: VirtualTcpSocket, +{ + let (socket, local_url, remote_url, resolved_remote_url) = connected.into_parts(); + let info = TunnelInfo { + tunnel_type: remote_url.scheme().to_owned(), + local_addr: local_url.map(Into::into), + remote_addr: Some(remote_url.clone().into()), + resolved_remote_addr: Some(resolved_remote_url.unwrap_or(remote_url).into()), + }; + TcpTunnelUpgrader::new(info) + .with_max_packet_size(BYTE_STREAM_MAX_PACKET_SIZE) + .upgrade(socket) +} + +pub(crate) fn upgrade_connected_tcp( + socket: S, + requested_remote_addr: Url, +) -> Result, TunnelError> +where + S: VirtualTcpSocket, +{ + let local_addr = socket.local_addr()?; + let resolved_remote_addr = socket.peer_addr()?; + let scheme = requested_remote_addr.scheme().to_owned(); + let tunnel_type = tcp_tunnel_type(&socket, &scheme)?; + let info = connected_tunnel_info( + &scheme, + &tunnel_type, + local_addr, + resolved_remote_addr, + requested_remote_addr, + ); + TcpTunnelUpgrader::new(info).upgrade(socket) +} + +pub(crate) fn upgrade_connected_udp( + connected: ConnectedUdpSession, + requested_remote_addr: Url, +) -> Result, TunnelError> { + let (session, layer) = connected.into_parts(); + let info = connected_tunnel_info( + "udp", + "udp", + session.local_addr()?, + session.peer_addr()?, + requested_remote_addr, + ); + UdpTunnelUpgrader::with_keep_alive(info, layer).upgrade(session) +} + +pub(crate) fn upgrade_accepted_tcp_with_local_url( + socket: S, + local_url: Url, +) -> Result, TunnelError> +where + S: VirtualTcpSocket, +{ + let remote_addr = socket.peer_addr()?; + let scheme = local_url.scheme().to_owned(); + let remote_url = socket_url(&scheme, remote_addr); + let info = TunnelInfo { + tunnel_type: tcp_tunnel_type(&socket, &scheme)?, + local_addr: Some(local_url.into()), + remote_addr: Some(remote_url.clone().into()), + resolved_remote_addr: Some(remote_url.into()), + }; + TcpTunnelUpgrader::new(info).upgrade(socket) +} + +pub(crate) fn upgrade_accepted_byte_stream( + socket: S, + local_url: Url, + remote_url: Option, +) -> Result, TunnelError> +where + S: VirtualTcpSocket, +{ + let info = TunnelInfo { + tunnel_type: local_url.scheme().to_owned(), + local_addr: Some(local_url.into()), + remote_addr: remote_url.clone().map(Into::into), + resolved_remote_addr: remote_url.map(Into::into), + }; + TcpTunnelUpgrader::new(info) + .with_max_packet_size(BYTE_STREAM_MAX_PACKET_SIZE) + .upgrade(socket) +} + +pub(crate) fn upgrade_accepted_udp_with_local_url( + session: UdpSession, + local_url: Url, +) -> Result, TunnelError> { + if local_url.scheme() != "udp" { + return Err(TunnelError::InvalidProtocol(format!( + "raw UDP listener requires udp URL: {local_url}" + ))); + } + let remote_url = socket_url("udp", session.peer_addr()?); + let info = TunnelInfo { + tunnel_type: "udp".to_owned(), + local_addr: Some(local_url.into()), + remote_addr: Some(remote_url.clone().into()), + resolved_remote_addr: Some(remote_url.into()), + }; + UdpTunnelUpgrader::new(info).upgrade(session) +} + +fn udp_bind_addrs_for_remote( + remote_addr: SocketAddr, + configured: &[SocketAddr], +) -> Vec { + if remote_addr.is_ipv6() { + Vec::new() + } else { + configured.to_vec() + } +} + +fn connected_tunnel_info( + scheme: &str, + tunnel_type: &str, + local_addr: SocketAddr, + resolved_remote_addr: SocketAddr, + requested_remote_addr: Url, +) -> TunnelInfo { + TunnelInfo { + tunnel_type: tunnel_type.to_owned(), + local_addr: Some(socket_url(scheme, local_addr).into()), + remote_addr: Some(requested_remote_addr.into()), + resolved_remote_addr: Some(socket_url(scheme, resolved_remote_addr).into()), + } +} + +fn tcp_tunnel_type(socket: &impl VirtualTcpSocket, scheme: &str) -> Result { + match socket.transport_label() { + Some(label) => Ok(label.to_owned()), + None if scheme == "faketcp" => Err(TunnelError::InternalError( + "FakeTCP upgrader received a socket without a FakeTCP transport label".to_owned(), + )), + None => Ok(scheme.to_owned()), + } +} + +fn socket_url(scheme: &str, addr: SocketAddr) -> Url { + let mut url = + Url::parse(&format!("{scheme}://0.0.0.0")).expect("static transport URL should be valid"); + url.set_ip_host(addr.ip()) + .expect("socket IP should be a valid URL host"); + url.set_port(Some(addr.port())) + .expect("transport URL should accept a port"); + url +} + +#[cfg(test)] +pub(crate) mod tests { + use std::{ + io, + pin::Pin, + task::{Context, Poll}, + }; + + use futures::{SinkExt, StreamExt}; + use tokio::io::{AsyncRead, AsyncWrite, DuplexStream, ReadBuf}; + + use crate::{ + packet::ZCPacket, + socket::udp::{UdpSessionKind, VirtualUdpSocket}, + }; + + use super::*; + + pub(crate) fn upgrade_accepted_tcp(socket: S) -> Result, TunnelError> + where + S: VirtualTcpSocket, + { + let local_addr = socket.local_addr()?; + upgrade_accepted_tcp_with_local_url(socket, socket_url("tcp", local_addr)) + } + + pub(crate) fn upgrade_accepted_udp( + session: UdpSession, + ) -> Result, TunnelError> { + let local_url = socket_url("udp", session.local_addr()?); + upgrade_accepted_udp_with_local_url(session, local_url) + } + + struct MockTcpSocket { + stream: DuplexStream, + local_addr: SocketAddr, + peer_addr: SocketAddr, + transport_label: Option<&'static str>, + } + + impl MockTcpSocket { + fn new(local_addr: SocketAddr, peer_addr: SocketAddr) -> Self { + let (stream, _) = tokio::io::duplex(64); + Self::from_stream(stream, local_addr, peer_addr) + } + + fn from_stream( + stream: DuplexStream, + local_addr: SocketAddr, + peer_addr: SocketAddr, + ) -> Self { + Self { + stream, + local_addr, + peer_addr, + transport_label: None, + } + } + + fn with_transport_label(mut self, transport_label: &'static str) -> Self { + self.transport_label = Some(transport_label); + self + } + } + + impl AsyncRead for MockTcpSocket { + fn poll_read( + mut self: Pin<&mut Self>, + cx: &mut Context<'_>, + buf: &mut ReadBuf<'_>, + ) -> Poll> { + Pin::new(&mut self.stream).poll_read(cx, buf) + } + } + + impl AsyncWrite for MockTcpSocket { + fn poll_write( + mut self: Pin<&mut Self>, + cx: &mut Context<'_>, + buf: &[u8], + ) -> Poll> { + Pin::new(&mut self.stream).poll_write(cx, buf) + } + + fn poll_flush(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + Pin::new(&mut self.stream).poll_flush(cx) + } + + fn poll_shutdown(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + Pin::new(&mut self.stream).poll_shutdown(cx) + } + } + + impl VirtualTcpSocket for MockTcpSocket { + fn local_addr(&self) -> io::Result { + Ok(self.local_addr) + } + + fn peer_addr(&self) -> io::Result { + Ok(self.peer_addr) + } + + fn transport_label(&self) -> Option<&str> { + self.transport_label + } + } + + struct MockUdpSocket { + local_addr: SocketAddr, + } + + #[async_trait] + impl VirtualUdpSocket for MockUdpSocket { + fn local_addr(&self) -> io::Result { + Ok(self.local_addr) + } + + async fn send_to(&self, data: &[u8], _addr: SocketAddr) -> io::Result { + Ok(data.len()) + } + + async fn recv_from(&self, _buf: &mut [u8]) -> io::Result<(usize, SocketAddr)> { + std::future::pending().await + } + } + + #[test] + fn raw_upgrader_preserves_requested_and_resolved_addresses() { + let local_addr: SocketAddr = "127.0.0.1:1000".parse().unwrap(); + let peer_addr: SocketAddr = "127.0.0.1:2000".parse().unwrap(); + let requested_url: Url = "tcp://example.com:2000".parse().unwrap(); + + let connected = upgrade_connected_tcp( + MockTcpSocket::new(local_addr, peer_addr), + requested_url.clone(), + ) + .unwrap(); + let connected_info = connected.info().unwrap(); + assert_eq!( + connected_info.remote_addr.unwrap().url, + requested_url.as_str() + ); + let resolved: Url = connected_info.resolved_remote_addr.unwrap().into(); + assert_eq!(resolved.host_str(), Some("127.0.0.1")); + assert_eq!(resolved.port(), Some(2000)); + + let accepted = upgrade_accepted_tcp(MockTcpSocket::new(local_addr, peer_addr)).unwrap(); + let accepted_info = accepted.info().unwrap(); + assert_eq!( + accepted_info.remote_addr, + accepted_info.resolved_remote_addr + ); + + let requested_local_url: Url = "tcp://0.0.0.0:1000".parse().unwrap(); + let accepted = upgrade_accepted_tcp_with_local_url( + MockTcpSocket::new(local_addr, peer_addr), + requested_local_url.clone(), + ) + .unwrap(); + assert_eq!( + accepted.info().unwrap().local_addr.unwrap().url, + requested_local_url.as_str() + ); + } + + #[test] + fn raw_tcp_upgrader_preserves_host_transport_label() { + let local_addr: SocketAddr = "192.0.2.1:10000".parse().unwrap(); + let peer_addr: SocketAddr = "192.0.2.2:11013".parse().unwrap(); + let requested_url: Url = "faketcp://peer.example:11013".parse().unwrap(); + + let connected = upgrade_connected_tcp( + MockTcpSocket::new(local_addr, peer_addr).with_transport_label("faketcp_test-driver"), + requested_url.clone(), + ) + .unwrap(); + let connected_info = connected.info().unwrap(); + assert_eq!(connected_info.tunnel_type, "faketcp_test-driver"); + assert_eq!( + connected_info.resolved_remote_addr.unwrap().url, + "faketcp://192.0.2.2:11013" + ); + + let accepted = upgrade_accepted_tcp_with_local_url( + MockTcpSocket::new(local_addr, peer_addr).with_transport_label("faketcp_test-driver"), + "faketcp://0.0.0.0:11013".parse().unwrap(), + ) + .unwrap(); + let accepted_info = accepted.info().unwrap(); + assert_eq!(accepted_info.tunnel_type, "faketcp_test-driver"); + assert_eq!( + accepted_info.remote_addr, + accepted_info.resolved_remote_addr + ); + } + + #[test] + fn faketcp_upgrader_rejects_socket_without_host_transport_label() { + let local_addr: SocketAddr = "192.0.2.1:10000".parse().unwrap(); + let peer_addr: SocketAddr = "192.0.2.2:11013".parse().unwrap(); + + let connected_error = upgrade_connected_tcp( + MockTcpSocket::new(local_addr, peer_addr), + "faketcp://peer.example:11013".parse().unwrap(), + ) + .unwrap_err(); + assert!(matches!(connected_error, TunnelError::InternalError(_))); + + let accepted_error = upgrade_accepted_tcp_with_local_url( + MockTcpSocket::new(local_addr, peer_addr), + "faketcp://0.0.0.0:11013".parse().unwrap(), + ) + .unwrap_err(); + assert!(matches!(accepted_error, TunnelError::InternalError(_))); + } + + #[test] + fn raw_udp_uses_default_bind_for_ipv6_remote() { + let configured = vec!["192.0.2.1:0".parse().unwrap()]; + + assert_eq!( + udp_bind_addrs_for_remote("198.51.100.1:11010".parse().unwrap(), &configured), + configured + ); + assert!( + udp_bind_addrs_for_remote("[2001:db8::1]:11010".parse().unwrap(), &configured) + .is_empty() + ); + } + + #[tokio::test] + async fn accepted_udp_uses_explicit_listener_url() { + let local_addr = "127.0.0.1:1000".parse().unwrap(); + let peer_addr = "127.0.0.1:2000".parse().unwrap(); + let local_url: Url = "udp://listener.example:1000?bind_device=eth0" + .parse() + .unwrap(); + let session = UdpSession::identity_standalone( + Arc::new(MockUdpSocket { local_addr }), + peer_addr, + UdpSessionKind::EasyTierMux, + ) + .unwrap(); + + let tunnel = upgrade_accepted_udp_with_local_url(session, local_url.clone()).unwrap(); + let info = tunnel.info().unwrap(); + + assert_eq!(info.local_addr.unwrap().url, local_url.as_str()); + assert_eq!(info.remote_addr, info.resolved_remote_addr); + } + + #[tokio::test] + async fn accepted_udp_rejects_non_udp_listener_url() { + let local_addr = "127.0.0.1:1000".parse().unwrap(); + let session = UdpSession::identity_standalone( + Arc::new(MockUdpSocket { local_addr }), + "127.0.0.1:2000".parse().unwrap(), + UdpSessionKind::EasyTierMux, + ) + .unwrap(); + + let error = + upgrade_accepted_udp_with_local_url(session, "quic://127.0.0.1:1000".parse().unwrap()) + .unwrap_err(); + + assert!(matches!(error, TunnelError::InvalidProtocol(_))); + } + + #[test] + fn byte_stream_upgrader_uses_host_endpoint_metadata() { + let local_url: Url = "ring://local".parse().unwrap(); + let remote_url: Url = "ring://remote".parse().unwrap(); + let tunnel = upgrade_connected_byte_stream(ConnectedByteStream::new( + MockTcpSocket::new( + "127.0.0.1:1000".parse().unwrap(), + "127.0.0.1:2000".parse().unwrap(), + ), + Some(local_url.clone()), + remote_url.clone(), + None, + )) + .unwrap(); + let info = tunnel.info().unwrap(); + + assert_eq!(info.tunnel_type, "ring"); + assert_eq!(info.local_addr.unwrap().url, local_url.as_str()); + assert_eq!(info.remote_addr.unwrap().url, remote_url.as_str()); + assert_eq!(info.resolved_remote_addr.unwrap().url, remote_url.as_str()); + } + + #[test] + fn accepted_byte_stream_uses_explicit_endpoint_metadata() { + let local_url: Url = "ring://local".parse().unwrap(); + let remote_url: Url = "ring://remote".parse().unwrap(); + let tunnel = upgrade_accepted_byte_stream( + MockTcpSocket::new( + "127.0.0.1:1000".parse().unwrap(), + "127.0.0.1:2000".parse().unwrap(), + ), + local_url.clone(), + Some(remote_url.clone()), + ) + .unwrap(); + let info = tunnel.info().unwrap(); + + assert_eq!(info.tunnel_type, "ring"); + assert_eq!(info.local_addr.unwrap().url, local_url.as_str()); + assert_eq!(info.remote_addr.unwrap().url, remote_url.as_str()); + assert_eq!(info.resolved_remote_addr.unwrap().url, remote_url.as_str()); + } + + #[test] + fn accepted_byte_stream_allows_unnamed_remote_endpoint() { + let tunnel = upgrade_accepted_byte_stream( + MockTcpSocket::new( + "127.0.0.1:1000".parse().unwrap(), + "127.0.0.1:2000".parse().unwrap(), + ), + "unix:///tmp/easytier.sock".parse().unwrap(), + None, + ) + .unwrap(); + let info = tunnel.info().unwrap(); + + assert_eq!(info.tunnel_type, "unix"); + assert!(info.remote_addr.is_none()); + assert!(info.resolved_remote_addr.is_none()); + } + + #[tokio::test] + async fn byte_stream_upgrader_preserves_legacy_unix_packet_limit() { + let (client_stream, server_stream) = tokio::io::duplex(8192); + let client = upgrade_connected_byte_stream(ConnectedByteStream::new( + MockTcpSocket::from_stream( + client_stream, + "127.0.0.1:1000".parse().unwrap(), + "127.0.0.1:2000".parse().unwrap(), + ), + None, + "unix:///tmp/easytier.sock".parse().unwrap(), + None, + )) + .unwrap(); + let server = upgrade_accepted_byte_stream( + MockTcpSocket::from_stream( + server_stream, + "127.0.0.1:2000".parse().unwrap(), + "127.0.0.1:1000".parse().unwrap(), + ), + "unix:///tmp/easytier.sock".parse().unwrap(), + Some("unix://anonymous/peer".parse().unwrap()), + ) + .unwrap(); + + let (mut client_stream, mut client_sink) = client.split(); + let (mut server_stream, mut server_sink) = server.split(); + let payload = vec![0x5a; 3000]; + client_sink + .send(ZCPacket::new_with_payload(&payload)) + .await + .unwrap(); + + let packet = server_stream.next().await.unwrap().unwrap(); + assert_eq!(packet.payload(), payload); + + let response_payload = vec![0xa5; 3000]; + server_sink + .send(ZCPacket::new_with_payload(&response_payload)) + .await + .unwrap(); + + let packet = client_stream.next().await.unwrap().unwrap(); + assert_eq!(packet.payload(), response_payload); + } +} diff --git a/easytier-core/src/connectivity/stun/client.rs b/easytier-core/src/connectivity/stun/client.rs new file mode 100644 index 00000000..5f35e3f4 --- /dev/null +++ b/easytier-core/src/connectivity/stun/client.rs @@ -0,0 +1,1033 @@ +//! Portable STUN clients and NAT classification. + +use std::{ + collections::BTreeSet, + net::{IpAddr, Ipv4Addr, Ipv6Addr, SocketAddr}, + sync::Arc, + time::Duration, +}; + +use anyhow::Context as _; +use bytecodec::{DecodeExt as _, EncodeExt as _}; +use quanta::Instant; +use rand::seq::IteratorRandom as _; +use stun_codec::{Message, MessageClass, MessageDecoder, MessageEncoder}; +use tokio::{ + io::{AsyncRead, AsyncReadExt as _, AsyncWriteExt as _}, + sync::{Mutex, broadcast}, + task::JoinSet, +}; +use tracing::{Instrument as _, Level}; + +use crate::{ + host::dns::{DnsQuery, DnsRecordResolver, DnsResolver}, + proto::common::NatType, + socket::{ + IpVersion, SocketContext, + tcp::{TcpBindOptions, TcpConnectOptions, VirtualTcpSocket, VirtualTcpSocketFactory}, + udp::{UdpBindOptions, VirtualUdpSocket, VirtualUdpSocketFactory}, + }, +}; + +use crate::packet::stun::{Attribute, ChangeRequest, tid_to_u32, u32_to_tid}; +use stun_codec::rfc5389::methods::BINDING; + +pub trait StunSocketRuntime: VirtualUdpSocketFactory + VirtualTcpSocketFactory {} + +impl StunSocketRuntime for T where T: VirtualUdpSocketFactory + VirtualTcpSocketFactory {} + +pub trait StunDnsRuntime: DnsResolver + DnsRecordResolver {} + +impl StunDnsRuntime for T where T: DnsResolver + DnsRecordResolver {} + +pub(super) struct HostResolverIter { + dns: Arc, + context: SocketContext, + hostnames: Vec, + ips: Vec, + max_ip_per_domain: u32, + use_ipv6: bool, +} + +impl HostResolverIter +where + D: StunDnsRuntime + ?Sized, +{ + pub(super) fn new( + dns: Arc, + context: SocketContext, + hostnames: Vec, + max_ip_per_domain: u32, + use_ipv6: bool, + ) -> Self { + Self { + dns, + context, + hostnames, + ips: Vec::new(), + max_ip_per_domain, + use_ipv6, + } + } + + async fn get_txt_record(&self, domain_name: &str) -> anyhow::Result> { + let txt_data = self + .dns + .resolve_txt(DnsQuery::new(domain_name, self.context.clone())) + .await?; + Ok(txt_data.split_whitespace().map(str::to_owned).collect()) + } + + pub(super) async fn next(&mut self) -> Option { + loop { + if let Some(addr) = self.ips.pop() { + return Some(addr); + } + + if self.hostnames.is_empty() { + return None; + } + + let endpoint = self.hostnames.remove(0); + if let Some(domain_name) = endpoint.strip_prefix("txt:") { + match self.get_txt_record(domain_name).await { + Ok(hosts) => { + tracing::info!( + ?domain_name, + ?hosts, + "get txt record success when resolve stun server" + ); + self.hostnames.splice(0..0, hosts); + } + Err(error) => { + tracing::warn!( + ?domain_name, + ?error, + "get txt record failed when resolve stun server" + ); + } + } + continue; + } + + if let Some(addr) = explicit_socket_addr(&endpoint) { + if addr.is_ipv6() == self.use_ipv6 { + return Some(addr); + } + continue; + } + + let Some((host, port)) = host_and_port(&endpoint) else { + tracing::warn!(?endpoint, "invalid stun server endpoint"); + continue; + }; + match self + .dns + .resolve(DnsQuery::new(host.clone(), self.context.clone())) + .await + { + Ok(ips) => { + self.ips = ips + .into_iter() + .filter(|ip| ip.is_ipv6() == self.use_ipv6) + .map(|ip| SocketAddr::new(ip, port)) + .choose_multiple(&mut rand::thread_rng(), self.max_ip_per_domain as usize); + } + Err(error) => { + tracing::warn!(?host, ?error, "lookup host for stun failed"); + } + } + } + } +} + +fn explicit_socket_addr(endpoint: &str) -> Option { + if let Ok(addr) = endpoint.parse() { + return Some(addr); + } + endpoint + .parse::() + .ok() + .map(|ip| SocketAddr::new(ip, 3478)) +} + +fn host_and_port(endpoint: &str) -> Option<(String, u16)> { + match endpoint.rsplit_once(':') { + Some((host, port)) if !host.is_empty() => { + Some((host.to_owned(), port.parse::().ok()?)) + } + _ => Some((endpoint.to_owned(), 3478)), + } +} + +#[derive(Debug, Clone)] +struct StunPacket { + data: Vec, + addr: SocketAddr, +} + +type StunPacketReceiver = broadcast::Receiver; + +#[derive(Debug, Clone, Copy)] +pub(super) struct BindRequestResponse { + pub(super) local_addr: SocketAddr, + pub(super) stun_server_addr: SocketAddr, + pub(super) recv_from_addr: SocketAddr, + pub(super) mapped_socket_addr: Option, + #[allow(dead_code)] + changed_socket_addr: Option, + #[allow(dead_code)] + change_ip: bool, + #[allow(dead_code)] + change_port: bool, + real_ip_changed: bool, + real_port_changed: bool, + #[allow(dead_code)] + latency_us: u32, +} + +#[derive(Debug, Clone)] +struct StunClient { + stun_server: SocketAddr, + resp_timeout: Duration, + req_repeat: u32, + socket: Arc, + stun_packet_receiver: Arc>, +} + +impl StunClient +where + S: VirtualUdpSocket, +{ + fn new( + stun_server: SocketAddr, + socket: Arc, + stun_packet_receiver: StunPacketReceiver, + ) -> Self { + Self { + stun_server, + resp_timeout: Duration::from_millis(3000), + req_repeat: 2, + socket, + stun_packet_receiver: Arc::new(Mutex::new(stun_packet_receiver)), + } + } + + async fn wait_stun_response( + &self, + tids: &[u32], + stun_host: &SocketAddr, + ) -> anyhow::Result<(Message, SocketAddr)> { + let mut now = tokio::time::Instant::now(); + let deadline = now + self.resp_timeout; + + while now < deadline { + let mut receiver = self.stun_packet_receiver.lock().await; + let packet = tokio::time::timeout(deadline - now, receiver.recv()).await??; + now = tokio::time::Instant::now(); + + if packet.data.len() < 20 { + continue; + } + + let mut decoder = MessageDecoder::::new(); + let Ok(message) = decoder + .decode_from_bytes(&packet.data) + .with_context(|| format!("decode stun message from {}", packet.addr))? + else { + continue; + }; + + tracing::trace!( + data = ?packet.data, + ?tids, + remote_addr = ?packet.addr, + ?stun_host, + "recv stun response: {message:#?}" + ); + + if message.class() != MessageClass::SuccessResponse + || message.method() != BINDING + || !tids.contains(&tid_to_u32(&message.transaction_id())) + { + continue; + } + + return Ok((message, packet.addr)); + } + + anyhow::bail!("timed out waiting for STUN response") + } + + fn extract_mapped_addr(message: &Message) -> Option { + message.attributes().find_map(|attribute| match attribute { + Attribute::MappedAddress(addr) => Some(addr.address()), + Attribute::XorMappedAddress(addr) => Some(addr.address()), + _ => None, + }) + } + + fn extract_changed_addr(message: &Message) -> Option { + message.attributes().find_map(|attribute| match attribute { + Attribute::OtherAddress(addr) => Some(addr.address()), + Attribute::ChangedAddress(addr) => Some(addr.address()), + _ => None, + }) + } + + #[tracing::instrument(ret, level = Level::TRACE, skip(self))] + async fn bind_request( + self, + change_ip: bool, + change_port: bool, + ) -> anyhow::Result { + let stun_host = self.stun_server; + let mut tids = Vec::new(); + for _ in 0..self.req_repeat { + let tid = rand::random::(); + let mut message = + Message::::new(MessageClass::Request, BINDING, u32_to_tid(tid)); + message.add_attribute(ChangeRequest::new(change_ip, change_port)); + let bytes = MessageEncoder::new() + .encode_into_bytes(message.clone()) + .with_context(|| "encode stun message")?; + tids.push(tid); + tracing::trace!(?message, ?bytes, tid, "send stun request"); + self.socket.send_to(&bytes, stun_host).await?; + } + + let now = Instant::now(); + let (message, recv_addr) = self.wait_stun_response(&tids, &stun_host).await?; + let changed_socket_addr = Self::extract_changed_addr(&message); + let response = BindRequestResponse { + local_addr: self.socket.local_addr()?, + stun_server_addr: stun_host, + recv_from_addr: recv_addr, + mapped_socket_addr: Self::extract_mapped_addr(&message), + changed_socket_addr, + change_ip, + change_port, + real_ip_changed: stun_host.ip() != recv_addr.ip(), + real_port_changed: stun_host.port() != recv_addr.port(), + latency_us: now.elapsed().as_micros() as u32, + }; + + tracing::trace!( + ?stun_host, + ?recv_addr, + ?changed_socket_addr, + "finish stun bind request" + ); + Ok(response) + } +} + +struct StunClientBuilder +where + S: VirtualUdpSocket, +{ + socket: Arc, + tasks: JoinSet<()>, + stun_packet_sender: broadcast::Sender, +} + +impl StunClientBuilder +where + S: VirtualUdpSocket, +{ + fn new(socket: Arc) -> Self { + let (stun_packet_sender, _) = broadcast::channel(1024); + let mut tasks = JoinSet::new(); + let listener_socket = socket.clone(); + let sender = stun_packet_sender.clone(); + tasks.spawn( + async move { + let mut buf = [0; 1620]; + tracing::trace!("start stun packet listener"); + loop { + let Ok((len, addr)) = listener_socket.recv_from(&mut buf).await else { + tracing::error!("udp recv_from error"); + break; + }; + let data = buf[..len].to_vec(); + tracing::trace!(?addr, ?data, "recv udp stun packet"); + let _ = sender.send(StunPacket { data, addr }); + } + } + .instrument(tracing::info_span!("stun_packet_listener")), + ); + Self { + socket, + tasks, + stun_packet_sender, + } + } + + fn new_stun_client(&self, stun_server: SocketAddr) -> StunClient { + StunClient::new( + stun_server, + self.socket.clone(), + self.stun_packet_sender.subscribe(), + ) + } + + async fn stop(&mut self) { + self.tasks.shutdown().await; + } +} + +impl Drop for StunClientBuilder +where + S: VirtualUdpSocket, +{ + fn drop(&mut self) { + self.tasks.abort_all(); + } +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum StunTransport { + Udp, + Tcp, +} + +#[derive(Debug, Clone)] +pub struct StunNatTypeDetectResult { + transport: StunTransport, + source_addr: SocketAddr, + pub(super) stun_resps: Vec, + pub(super) extra_bind_test: Option, +} + +impl StunNatTypeDetectResult { + fn new( + transport: StunTransport, + source_addr: SocketAddr, + stun_resps: Vec, + ) -> Self { + Self { + transport, + source_addr, + stun_resps, + extra_bind_test: None, + } + } + + fn has_ip_changed_resp(&self) -> bool { + self.stun_resps.iter().any(|resp| resp.real_ip_changed) + } + + fn has_port_changed_resp(&self) -> bool { + self.stun_resps.iter().any(|resp| resp.real_port_changed) + } + + fn is_open_internet(&self) -> bool { + self.stun_resps + .iter() + .any(|resp| resp.mapped_socket_addr == Some(self.source_addr)) + } + + fn is_no_pat(&self) -> bool { + self.stun_resps.iter().any(|resp| { + resp.mapped_socket_addr.map(|addr| addr.port()) == Some(self.source_addr.port()) + }) + } + + fn stun_server_count(&self) -> usize { + self.stun_resps + .iter() + .map(|resp| resp.recv_from_addr) + .collect::>() + .len() + } + + fn is_cone(&self) -> bool { + self.stun_resps + .iter() + .filter_map(|resp| resp.mapped_socket_addr) + .collect::>() + .len() + == 1 + } + + fn nat_type_udp(&self) -> NatType { + if self.stun_server_count() < 2 { + return NatType::Unknown; + } + + if self.is_cone() { + if self.has_ip_changed_resp() { + if self.is_open_internet() { + NatType::OpenInternet + } else if self.is_no_pat() { + NatType::NoPat + } else { + NatType::FullCone + } + } else if self.has_port_changed_resp() { + NatType::Restricted + } else { + NatType::PortRestricted + } + } else if !self.stun_resps.is_empty() { + if self.public_ips().len() != 1 + || self.usable_stun_resp_count() <= 1 + || self.max_port() - self.min_port() > 15 + { + NatType::Symmetric + } else if let Some(extra_bind_mapped) = self + .extra_bind_test + .as_ref() + .and_then(|extra| extra.mapped_socket_addr) + { + let extra_port = extra_bind_mapped.port(); + let max_port_diff = extra_port.saturating_sub(self.max_port()); + let min_port_diff = self.min_port().saturating_sub(extra_port); + if max_port_diff != 0 && max_port_diff < 100 { + NatType::SymmetricEasyInc + } else if min_port_diff != 0 && min_port_diff < 100 { + NatType::SymmetricEasyDec + } else { + NatType::Symmetric + } + } else { + NatType::Symmetric + } + } else { + NatType::Unknown + } + } + + fn nat_type_tcp(&self) -> NatType { + if self.is_open_internet() { + return NatType::OpenInternet; + } + if self.stun_server_count() < 2 || self.stun_resps.is_empty() { + return NatType::Unknown; + } + if self.is_cone() { + if self.is_no_pat() { + NatType::NoPat + } else { + NatType::FullCone + } + } else { + NatType::Symmetric + } + } + + pub fn nat_type(&self) -> NatType { + match self.transport { + StunTransport::Udp => self.nat_type_udp(), + StunTransport::Tcp => self.nat_type_tcp(), + } + } + + pub fn public_ips(&self) -> Vec { + self.stun_resps + .iter() + .filter_map(|resp| resp.mapped_socket_addr.map(|addr| addr.ip())) + .collect::>() + .into_iter() + .collect() + } + + pub fn collect_available_stun_server(&self) -> Vec { + let mut servers = Vec::new(); + for response in &self.stun_resps { + if !servers.contains(&response.stun_server_addr) { + servers.push(response.stun_server_addr); + } + } + servers + } + + pub fn local_addr(&self) -> SocketAddr { + self.source_addr + } + + pub fn extend_result(&mut self, other: Self) { + self.stun_resps.extend(other.stun_resps); + } + + pub fn min_port(&self) -> u16 { + self.stun_resps + .iter() + .filter_map(|response| response.mapped_socket_addr.map(|addr| addr.port())) + .min() + .unwrap_or(0) + } + + pub fn max_port(&self) -> u16 { + self.stun_resps + .iter() + .filter_map(|response| response.mapped_socket_addr.map(|addr| addr.port())) + .max() + .unwrap_or(u16::MAX) + } + + pub fn usable_stun_resp_count(&self) -> usize { + self.stun_resps + .iter() + .filter(|response| response.mapped_socket_addr.is_some()) + .count() + } +} + +pub struct UdpNatTypeDetector { + runtime: Arc, + dns: Arc, + socket_context: SocketContext, + stun_server_hosts: Vec, + max_ip_per_domain: u32, +} + +impl UdpNatTypeDetector +where + R: StunSocketRuntime, + D: StunDnsRuntime + ?Sized, +{ + pub fn new( + runtime: Arc, + dns: Arc, + socket_context: SocketContext, + stun_server_hosts: Vec, + max_ip_per_domain: u32, + ) -> Self { + Self { + runtime, + dns, + socket_context, + stun_server_hosts, + max_ip_per_domain, + } + } + + pub(super) async fn get_extra_bind_result( + &self, + source_port: u16, + stun_server: SocketAddr, + ) -> anyhow::Result { + let socket = self + .runtime + .bind_udp(stun_udp_bind_options( + self.socket_context.clone(), + IpVersion::V4, + SocketAddr::new(Ipv4Addr::UNSPECIFIED.into(), source_port), + )) + .await?; + udp_bind_request(socket, stun_server).await + } + + pub async fn detect_nat_type( + &self, + source_port: u16, + ) -> anyhow::Result { + let socket = self + .runtime + .bind_udp(stun_udp_bind_options( + self.socket_context.clone(), + IpVersion::V4, + SocketAddr::new(Ipv4Addr::UNSPECIFIED.into(), source_port), + )) + .await?; + self.detect_nat_type_with_socket(socket).await + } + + #[tracing::instrument(skip(self, socket))] + pub async fn detect_nat_type_with_socket( + &self, + socket: Arc<::Socket>, + ) -> anyhow::Result { + let mut resolver = HostResolverIter::new( + self.dns.clone(), + self.socket_context.clone().with_ip_version(IpVersion::V4), + self.stun_server_hosts.clone(), + self.max_ip_per_domain, + false, + ); + let mut stun_servers = Vec::new(); + while let Some(addr) = resolver.next().await { + stun_servers.push(addr); + } + + let client_builder = StunClientBuilder::new(socket.clone()); + let mut tasks = JoinSet::new(); + for stun_server in stun_servers { + tasks.spawn( + client_builder + .new_stun_client(stun_server) + .bind_request(false, false), + ); + tasks.spawn( + client_builder + .new_stun_client(stun_server) + .bind_request(false, true), + ); + tasks.spawn( + client_builder + .new_stun_client(stun_server) + .bind_request(true, true), + ); + } + + let mut responses = Vec::new(); + while let Some(response) = tasks.join_next().await { + if let Ok(Ok(response)) = response { + responses.push(response); + } + } + Ok(StunNatTypeDetectResult::new( + StunTransport::Udp, + socket.local_addr()?, + responses, + )) + } +} + +pub(super) async fn udp_bind_request( + socket: Arc, + stun_server: SocketAddr, +) -> anyhow::Result +where + S: VirtualUdpSocket, +{ + let mut clients = StunClientBuilder::new(socket); + let result = clients + .new_stun_client(stun_server) + .bind_request(false, false) + .await; + clients.stop().await; + result +} + +pub(super) fn stun_udp_bind_options( + context: SocketContext, + ip_version: IpVersion, + local_addr: SocketAddr, +) -> UdpBindOptions { + UdpBindOptions::stun_probe() + .with_context(context.with_ip_version(ip_version)) + .with_local_addr(Some(local_addr)) +} + +struct TcpStunClient { + runtime: Arc, + socket_context: SocketContext, + stun_server: SocketAddr, + conn_timeout: Duration, + io_timeout: Duration, + source_port: u16, +} + +impl TcpStunClient +where + R: StunSocketRuntime, +{ + fn new( + runtime: Arc, + socket_context: SocketContext, + stun_server: SocketAddr, + source_port: u16, + ) -> Self { + Self { + runtime, + socket_context, + stun_server, + conn_timeout: Duration::from_millis(1500), + io_timeout: Duration::from_millis(3000), + source_port, + } + } + + fn extract_mapped_addr(message: &Message) -> Option { + message.attributes().find_map(|attribute| match attribute { + Attribute::MappedAddress(addr) => Some(addr.address()), + Attribute::XorMappedAddress(addr) => Some(addr.address()), + _ => None, + }) + } + + fn message_size_from_header(header: &[u8; 20]) -> anyhow::Result { + if (header[0] & 0b1100_0000) != 0 { + anyhow::bail!("invalid stun message type") + } + let message_len = u16::from_be_bytes([header[2], header[3]]) as usize; + if !message_len.is_multiple_of(4) { + anyhow::bail!("invalid stun message length") + } + let total = 20usize + .checked_add(message_len) + .context("invalid stun message size")?; + if total > 4096 { + anyhow::bail!("stun message too large") + } + Ok(total) + } + + async fn tcp_read_stun_message( + stream: &mut S, + timeout: Duration, + ) -> anyhow::Result> + where + S: AsyncRead + Unpin, + { + let mut header = [0u8; 20]; + tokio::time::timeout(timeout, stream.read_exact(&mut header)).await??; + let total_size = Self::message_size_from_header(&header)?; + let mut buf = vec![0u8; total_size]; + buf[..20].copy_from_slice(&header); + if total_size > 20 { + tokio::time::timeout(timeout, stream.read_exact(&mut buf[20..])).await??; + } + + let mut decoder = MessageDecoder::::new(); + let Ok(message) = decoder + .decode_from_bytes(&buf) + .with_context(|| "decode tcp stun message")? + else { + anyhow::bail!("invalid stun message") + }; + Ok(message) + } + + async fn connect(&self) -> anyhow::Result<::Socket> { + let (bind_addr, ip_version) = match self.stun_server { + SocketAddr::V4(_) => ( + SocketAddr::new(Ipv4Addr::UNSPECIFIED.into(), self.source_port), + IpVersion::V4, + ), + SocketAddr::V6(_) => ( + SocketAddr::new(Ipv6Addr::UNSPECIFIED.into(), self.source_port), + IpVersion::V6, + ), + }; + let bind = TcpBindOptions::default() + .with_context(self.socket_context.clone().with_ip_version(ip_version)) + .with_local_addr(Some(bind_addr)) + .with_reuse_addr(true) + .with_reuse_port(true) + .with_only_v6(bind_addr.is_ipv6()); + tokio::time::timeout( + self.conn_timeout, + self.runtime.connect_tcp( + TcpConnectOptions::stun_probe(self.stun_server, bind_addr).with_bind(bind), + ), + ) + .await? + } + + #[tracing::instrument(ret, level = Level::TRACE, skip(self))] + async fn bind_request(self) -> anyhow::Result { + let mut stream = self.connect().await?; + let local_addr = stream.local_addr()?; + let stun_host = self.stun_server; + let tid = rand::random::(); + let message = Message::::new(MessageClass::Request, BINDING, u32_to_tid(tid)); + let bytes = MessageEncoder::new() + .encode_into_bytes(message) + .with_context(|| "encode tcp stun message")?; + tokio::time::timeout(self.io_timeout, stream.write_all(&bytes)).await??; + + let now = Instant::now(); + let message = Self::tcp_read_stun_message(&mut stream, self.io_timeout).await?; + if message.class() != MessageClass::SuccessResponse + || message.method() != BINDING + || tid_to_u32(&message.transaction_id()) != tid + { + anyhow::bail!("unexpected stun response") + } + + Ok(BindRequestResponse { + local_addr, + stun_server_addr: stun_host, + recv_from_addr: stun_host, + mapped_socket_addr: Self::extract_mapped_addr(&message), + changed_socket_addr: None, + change_ip: false, + change_port: false, + real_ip_changed: false, + real_port_changed: false, + latency_us: now.elapsed().as_micros() as u32, + }) + } +} + +pub(super) async fn tcp_bind_request( + runtime: Arc, + socket_context: SocketContext, + stun_server: SocketAddr, + source_port: u16, +) -> anyhow::Result +where + R: StunSocketRuntime, +{ + TcpStunClient::new(runtime, socket_context, stun_server, source_port) + .bind_request() + .await +} + +pub struct TcpNatTypeDetector { + runtime: Arc, + dns: Arc, + socket_context: SocketContext, + stun_server_hosts: Vec, + max_ip_per_domain: u32, +} + +impl TcpNatTypeDetector +where + R: StunSocketRuntime, + D: StunDnsRuntime + ?Sized, +{ + pub fn new( + runtime: Arc, + dns: Arc, + socket_context: SocketContext, + stun_server_hosts: Vec, + max_ip_per_domain: u32, + ) -> Self { + Self { + runtime, + dns, + socket_context, + stun_server_hosts, + max_ip_per_domain, + } + } + + #[tracing::instrument(skip(self))] + pub async fn detect_nat_type( + &self, + source_port: u16, + ) -> anyhow::Result { + let mut resolver = HostResolverIter::new( + self.dns.clone(), + self.socket_context.clone().with_ip_version(IpVersion::V4), + self.stun_server_hosts.clone(), + self.max_ip_per_domain, + false, + ); + let mut stun_servers = Vec::new(); + while let Some(addr) = resolver.next().await { + stun_servers.push(addr); + } + + let mut responses = Vec::new(); + let mut source_addr = None; + let mut selected_source_port = (source_port != 0).then_some(source_port); + for server in stun_servers { + let response = TcpStunClient::new( + self.runtime.clone(), + self.socket_context.clone(), + server, + selected_source_port.unwrap_or(0), + ) + .bind_request() + .await; + if let Ok(response) = response { + if selected_source_port.is_none() { + selected_source_port = Some(response.local_addr.port()); + } + source_addr.get_or_insert(response.local_addr); + responses.push(response); + if responses.len() >= 3 { + break; + } + } + } + + let source_addr = source_addr.context("no TCP STUN response")?; + Ok(StunNatTypeDetectResult::new( + StunTransport::Tcp, + source_addr, + responses, + )) + } +} + +#[cfg(test)] +mod tests { + use async_trait::async_trait; + + use crate::host::dns::{DnsQuery, DnsRecordResolver, DnsResolver, DnsSrvRecord}; + + use super::*; + + struct EmptyDns; + + #[async_trait] + impl DnsResolver for EmptyDns { + async fn resolve(&self, _query: DnsQuery) -> anyhow::Result> { + Ok(Vec::new()) + } + } + + #[async_trait] + impl DnsRecordResolver for EmptyDns { + async fn resolve_txt(&self, _query: DnsQuery) -> anyhow::Result { + Ok(String::new()) + } + + async fn resolve_srv(&self, _query: DnsQuery) -> anyhow::Result> { + Ok(Vec::new()) + } + } + + #[test] + fn explicit_endpoints_keep_ports_and_default_bare_ips() { + assert_eq!( + explicit_socket_addr("127.0.0.1:5555"), + Some("127.0.0.1:5555".parse().unwrap()) + ); + assert_eq!( + explicit_socket_addr("[::1]:5555"), + Some("[::1]:5555".parse().unwrap()) + ); + assert_eq!( + explicit_socket_addr("2001:db8::1"), + Some("[2001:db8::1]:3478".parse().unwrap()) + ); + } + + #[tokio::test] + async fn invalid_endpoint_does_not_hide_later_servers() { + let mut resolver = HostResolverIter::new( + Arc::new(EmptyDns), + SocketContext::default(), + vec![ + "bad.example:not-a-port".to_owned(), + "127.0.0.1:3478".to_owned(), + ], + 1, + false, + ); + + assert_eq!( + resolver.next().await, + Some("127.0.0.1:3478".parse().unwrap()) + ); + } + + #[test] + fn stun_udp_bind_request_preserves_context_and_family() { + let context = SocketContext::default() + .with_socket_mark(Some(0)) + .with_netns(Some(crate::socket::NetNamespace::new("instance-a"))); + let local_addr = "0.0.0.0:0".parse().unwrap(); + + let options = stun_udp_bind_options(context, IpVersion::V4, local_addr); + + assert_eq!(options.context.ip_version, IpVersion::V4); + assert_eq!(options.context.socket_mark, Some(0)); + assert_eq!( + options.context.netns.as_ref().map(|netns| netns.token()), + Some("instance-a") + ); + assert_eq!(options.local_addr, Some(local_addr)); + assert_eq!( + options.purpose, + crate::socket::udp::UdpSocketPurpose::StunProbe + ); + } +} diff --git a/easytier-core/src/connectivity/stun/collector.rs b/easytier-core/src/connectivity/stun/collector.rs new file mode 100644 index 00000000..b76b9971 --- /dev/null +++ b/easytier-core/src/connectivity/stun/collector.rs @@ -0,0 +1,768 @@ +//! Per-instance STUN state and background detection lifecycle. + +use std::{ + collections::BTreeSet, + net::{IpAddr, Ipv4Addr, Ipv6Addr, SocketAddr}, + sync::{ + Arc, Mutex, RwLock, + atomic::{AtomicBool, Ordering}, + }, + time::{Duration, SystemTime, UNIX_EPOCH}, +}; + +use async_trait::async_trait; +use rand::seq::IteratorRandom as _; +use serde::{Deserialize, Serialize}; +use tokio::task::JoinSet; + +use crate::{ + config::{ + DEFAULT_TCP_STUN_SERVERS, DEFAULT_UDP_STUN_SERVERS, DEFAULT_UDP_V6_STUN_SERVERS, + default_stun_servers, + }, + proto::common::{NatType, StunInfo}, + socket::{ + IpVersion, SocketContext, + udp::{VirtualUdpSocket, VirtualUdpSocketFactory}, + }, +}; + +use super::client::{ + HostResolverIter, StunDnsRuntime, StunNatTypeDetectResult, StunSocketRuntime, + TcpNatTypeDetector, UdpNatTypeDetector, stun_udp_bind_options, tcp_bind_request, + udp_bind_request, +}; + +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +pub struct StunServerConfig { + pub udp_servers: Vec, + pub tcp_servers: Vec, + pub udp_v6_servers: Vec, +} + +impl Default for StunServerConfig { + fn default() -> Self { + Self { + udp_servers: default_stun_servers(DEFAULT_UDP_STUN_SERVERS), + tcp_servers: default_stun_servers(DEFAULT_TCP_STUN_SERVERS), + udp_v6_servers: default_stun_servers(DEFAULT_UDP_V6_STUN_SERVERS), + } + } +} + +#[async_trait] +#[auto_impl::auto_impl(&, Arc, Box)] +pub trait StunInfoProvider: Send + Sync { + fn get_stun_info(&self) -> StunInfo; + + async fn get_udp_port_mapping(&self, local_port: u16) -> anyhow::Result; + + async fn get_tcp_port_mapping(&self, local_port: u16) -> anyhow::Result; + + fn update_stun_info(&self); +} + +#[async_trait] +#[auto_impl::auto_impl(&, Arc, Box)] +pub trait StunSocketMapper: StunInfoProvider + Send + Sync +where + S: VirtualUdpSocket, +{ + async fn get_udp_port_mapping_with_socket(&self, socket: Arc) -> anyhow::Result; +} + +pub struct StunInfoCollector +where + R: StunSocketRuntime, + D: StunDnsRuntime, +{ + runtime: Arc, + dns: Arc, + udp_socket_context: SocketContext, + tcp_socket_context: SocketContext, + stun_servers: Arc>>, + tcp_stun_servers: Arc>>, + stun_servers_v6: Arc>>, + udp_nat_test_result: Arc>>, + tcp_nat_test_result: Arc>>, + public_ipv6: Arc>>, + nat_test_result_time: Arc>, + redetect_notify: Arc, + tasks: Mutex>, + started: AtomicBool, +} + +impl StunInfoCollector +where + R: StunSocketRuntime, + D: StunDnsRuntime + ?Sized, +{ + pub fn new( + runtime: Arc, + dns: Arc, + socket_context: SocketContext, + udp_stun_servers: Vec, + tcp_stun_servers: Vec, + stun_servers_v6: Vec, + ) -> Self { + Self::new_with_socket_contexts( + runtime, + dns, + socket_context.clone(), + socket_context, + udp_stun_servers, + tcp_stun_servers, + stun_servers_v6, + ) + } + + pub fn new_with_socket_contexts( + runtime: Arc, + dns: Arc, + udp_socket_context: SocketContext, + tcp_socket_context: SocketContext, + udp_stun_servers: Vec, + tcp_stun_servers: Vec, + stun_servers_v6: Vec, + ) -> Self { + Self { + runtime, + dns, + udp_socket_context, + tcp_socket_context, + stun_servers: Arc::new(RwLock::new(udp_stun_servers)), + tcp_stun_servers: Arc::new(RwLock::new(tcp_stun_servers)), + stun_servers_v6: Arc::new(RwLock::new(stun_servers_v6)), + udp_nat_test_result: Arc::new(RwLock::new(None)), + tcp_nat_test_result: Arc::new(RwLock::new(None)), + public_ipv6: Arc::new(RwLock::new(None)), + nat_test_result_time: Arc::new(RwLock::new(unix_timestamp())), + redetect_notify: Arc::new(tokio::sync::Notify::new()), + tasks: Mutex::new(JoinSet::new()), + started: AtomicBool::new(false), + } + } + + pub fn new_with_default_servers( + runtime: Arc, + dns: Arc, + socket_context: SocketContext, + ) -> Self { + Self::new( + runtime, + dns, + socket_context, + Self::get_default_servers(), + Self::get_default_tcp_servers(), + Self::get_default_servers_v6(), + ) + } + + pub fn new_with_default_servers_and_socket_contexts( + runtime: Arc, + dns: Arc, + udp_socket_context: SocketContext, + tcp_socket_context: SocketContext, + ) -> Self { + Self::new_with_socket_contexts( + runtime, + dns, + udp_socket_context, + tcp_socket_context, + Self::get_default_servers(), + Self::get_default_tcp_servers(), + Self::get_default_servers_v6(), + ) + } + + pub fn set_stun_servers(&self, stun_servers: Vec) { + *self.stun_servers.write().unwrap() = stun_servers; + } + + pub fn set_stun_servers_v6(&self, stun_servers_v6: Vec) { + *self.stun_servers_v6.write().unwrap() = stun_servers_v6; + } + + pub fn set_tcp_stun_servers(&self, stun_servers: Vec) { + *self.tcp_stun_servers.write().unwrap() = stun_servers; + } + + pub fn get_default_servers() -> Vec { + StunServerConfig::default().udp_servers + } + + pub fn get_default_tcp_servers() -> Vec { + StunServerConfig::default().tcp_servers + } + + pub fn get_default_servers_v6() -> Vec { + StunServerConfig::default().udp_v6_servers + } + + async fn get_public_ipv6( + runtime: Arc, + dns: Arc, + socket_context: SocketContext, + servers: &[String], + ) -> Option { + let mut resolver = HostResolverIter::new( + dns, + socket_context.clone().with_ip_version(IpVersion::V6), + servers.to_vec(), + 10, + true, + ); + while let Some(server) = resolver.next().await { + let socket = runtime + .bind_udp(stun_udp_bind_options( + socket_context.clone(), + IpVersion::V6, + SocketAddr::new(Ipv6Addr::UNSPECIFIED.into(), 0), + )) + .await + .ok()?; + let response = udp_bind_request(socket, server).await; + tracing::debug!(?response, "finish ipv6 udp nat type detect"); + if let Ok(Some(IpAddr::V6(ip))) = + response.map(|response| response.mapped_socket_addr.map(|addr| addr.ip())) + { + return Some(ip); + } + } + None + } + + fn start_stun_routine(&self) { + if self.started.swap(true, Ordering::AcqRel) { + return; + } + + let runtime = self.runtime.clone(); + let dns = self.dns.clone(); + let socket_context = self.udp_socket_context.clone(); + let stun_servers = self.stun_servers.clone(); + let udp_nat_test_result = self.udp_nat_test_result.clone(); + let nat_test_time = self.nat_test_result_time.clone(); + let redetect_notify = self.redetect_notify.clone(); + self.tasks.lock().unwrap().spawn(async move { + loop { + let servers = sampled_servers(&stun_servers.read().unwrap()); + let detector = UdpNatTypeDetector::new( + runtime.clone(), + dns.clone(), + socket_context.clone(), + servers, + 1, + ); + let mut result = detector.detect_nat_type(0).await; + tracing::debug!(?result, "finish udp nat type detect"); + + let nat_type = result + .as_ref() + .map(StunNatTypeDetectResult::nat_type) + .unwrap_or(NatType::Unknown); + if nat_type == NatType::Symmetric { + let old_result = result.as_mut().unwrap(); + tracing::debug!(?old_result, "start get extra bind result"); + for server in old_result.collect_available_stun_server() { + let extra = detector.get_extra_bind_result(0, server).await; + tracing::debug!(?extra, "finish udp nat type detect with another port"); + if let Ok(response) = extra { + old_result.extra_bind_test = Some(response); + break; + } + } + } + + let mut sleep_sec = 10; + if let Ok(result) = result { + *nat_test_time.write().unwrap() = unix_timestamp(); + let completed_extra_test = result.extra_bind_test.is_some(); + *udp_nat_test_result.write().unwrap() = Some(result); + if nat_type != NatType::Unknown + && (nat_type != NatType::Symmetric || completed_extra_test) + { + sleep_sec = 600; + } + } + + tokio::select! { + _ = redetect_notify.notified() => {} + _ = tokio::time::sleep(Duration::from_secs(sleep_sec)) => {} + } + } + }); + + let runtime = self.runtime.clone(); + let dns = self.dns.clone(); + let socket_context = self.tcp_socket_context.clone(); + let tcp_stun_servers = self.tcp_stun_servers.clone(); + let tcp_nat_test_result = self.tcp_nat_test_result.clone(); + let nat_test_time = self.nat_test_result_time.clone(); + let redetect_notify = self.redetect_notify.clone(); + self.tasks.lock().unwrap().spawn(async move { + loop { + let servers = sampled_servers(&tcp_stun_servers.read().unwrap()); + let detector = TcpNatTypeDetector::new( + runtime.clone(), + dns.clone(), + socket_context.clone(), + servers, + 1, + ); + let result = detector.detect_nat_type(0).await; + tracing::debug!(?result, "finish tcp nat type detect"); + + let mut sleep_sec = 10; + if let Ok(result) = result { + *nat_test_time.write().unwrap() = unix_timestamp(); + let nat_type = result.nat_type(); + *tcp_nat_test_result.write().unwrap() = Some(result); + if nat_type != NatType::Unknown { + sleep_sec = 600; + } + } + + tokio::select! { + _ = redetect_notify.notified() => {} + _ = tokio::time::sleep(Duration::from_secs(sleep_sec)) => {} + } + } + }); + + let runtime = self.runtime.clone(); + let dns = self.dns.clone(); + let socket_context = self.udp_socket_context.clone(); + let stun_servers_v6 = self.stun_servers_v6.clone(); + let public_ipv6 = self.public_ipv6.clone(); + let redetect_notify = self.redetect_notify.clone(); + self.tasks.lock().unwrap().spawn(async move { + loop { + let servers = stun_servers_v6.read().unwrap().clone(); + if let Some(ip) = Self::get_public_ipv6( + runtime.clone(), + dns.clone(), + socket_context.clone(), + &servers, + ) + .await + { + *public_ipv6.write().unwrap() = Some(ip); + } + + let sleep_sec = if public_ipv6.read().unwrap().is_none() { + 60 + } else { + 360 + }; + tokio::select! { + _ = redetect_notify.notified() => {} + _ = tokio::time::sleep(Duration::from_secs(sleep_sec)) => {} + } + } + }); + } +} + +#[async_trait] +impl StunInfoProvider for StunInfoCollector +where + R: StunSocketRuntime, + D: StunDnsRuntime + ?Sized, +{ + fn get_stun_info(&self) -> StunInfo { + self.start_stun_routine(); + let udp_result = self.udp_nat_test_result.read().unwrap().clone(); + let tcp_result = self.tcp_nat_test_result.read().unwrap().clone(); + if udp_result.is_none() && tcp_result.is_none() { + return StunInfo::default(); + } + + let mut public_ip = BTreeSet::::new(); + if let Some(result) = &udp_result { + public_ip.extend(result.public_ips().into_iter().map(|ip| ip.to_string())); + } + if let Some(result) = &tcp_result { + public_ip.extend(result.public_ips().into_iter().map(|ip| ip.to_string())); + } + if let Some(ip) = *self.public_ipv6.read().unwrap() { + public_ip.insert(ip.to_string()); + } + + StunInfo { + udp_nat_type: udp_result + .as_ref() + .map(|result| result.nat_type() as i32) + .unwrap_or(NatType::Unknown as i32), + tcp_nat_type: tcp_result + .as_ref() + .map(|result| result.nat_type() as i32) + .unwrap_or(NatType::Unknown as i32), + last_update_time: *self.nat_test_result_time.read().unwrap(), + public_ip: public_ip.into_iter().collect(), + min_port: udp_result + .as_ref() + .map(|result| result.min_port() as u32) + .or_else(|| tcp_result.as_ref().map(|result| result.min_port() as u32)) + .unwrap_or(0), + max_port: udp_result + .as_ref() + .map(|result| result.max_port() as u32) + .or_else(|| tcp_result.as_ref().map(|result| result.max_port() as u32)) + .unwrap_or(0), + } + } + + async fn get_udp_port_mapping(&self, local_port: u16) -> anyhow::Result { + let socket = self + .runtime + .bind_udp(stun_udp_bind_options( + self.udp_socket_context.clone(), + IpVersion::V4, + SocketAddr::new(Ipv4Addr::UNSPECIFIED.into(), local_port), + )) + .await?; + StunSocketMapper::get_udp_port_mapping_with_socket(self, socket).await + } + + async fn get_tcp_port_mapping(&self, local_port: u16) -> anyhow::Result { + self.start_stun_routine(); + let mut servers = self + .tcp_nat_test_result + .read() + .unwrap() + .clone() + .map(|result| result.collect_available_stun_server()) + .unwrap_or_default(); + if servers.is_empty() { + let mut resolver = HostResolverIter::new( + self.dns.clone(), + self.tcp_socket_context + .clone() + .with_ip_version(IpVersion::V4), + self.tcp_stun_servers.read().unwrap().clone(), + 2, + false, + ); + while let Some(addr) = resolver.next().await { + servers.push(addr); + if servers.len() >= 2 { + break; + } + } + } + + for server in servers { + match tcp_bind_request( + self.runtime.clone(), + self.tcp_socket_context.clone(), + server, + local_port, + ) + .await + { + Ok(response) => { + if let Some(mapped_addr) = response.mapped_socket_addr { + return Ok(mapped_addr); + } + } + Err(error) => tracing::warn!(?server, ?error, "tcp stun bind request failed"), + } + } + anyhow::bail!("no TCP STUN mapping found") + } + + fn update_stun_info(&self) { + self.redetect_notify.notify_waiters(); + } +} + +#[async_trait] +impl StunSocketMapper<::Socket> for StunInfoCollector +where + R: StunSocketRuntime, + D: StunDnsRuntime + ?Sized, +{ + async fn get_udp_port_mapping_with_socket( + &self, + socket: Arc<::Socket>, + ) -> anyhow::Result { + self.start_stun_routine(); + let mut servers = self + .udp_nat_test_result + .read() + .unwrap() + .clone() + .map(|result| result.collect_available_stun_server()) + .unwrap_or_default(); + if servers.is_empty() { + let mut resolver = HostResolverIter::new( + self.dns.clone(), + self.udp_socket_context + .clone() + .with_ip_version(IpVersion::V4), + self.stun_servers.read().unwrap().clone(), + 2, + false, + ); + while let Some(addr) = resolver.next().await { + servers.push(addr); + if servers.len() >= 2 { + break; + } + } + } + + for server in servers { + match udp_bind_request(socket.clone(), server).await { + Ok(response) => { + if let Some(mapped_addr) = response.mapped_socket_addr { + return Ok(mapped_addr); + } + } + Err(error) => tracing::warn!(?server, ?error, "stun bind request failed"), + } + } + anyhow::bail!("no UDP STUN mapping found") + } +} + +fn sampled_servers(servers: &[String]) -> Vec { + servers + .iter() + .take(2) + .chain(servers.iter().skip(2).choose(&mut rand::thread_rng())) + .cloned() + .collect() +} + +fn unix_timestamp() -> i64 { + SystemTime::now() + .duration_since(UNIX_EPOCH) + .unwrap_or_default() + .as_secs() as i64 +} + +#[cfg(test)] +mod tests { + use std::{ + collections::VecDeque, + io, + pin::Pin, + task::{Context, Poll}, + }; + + use bytecodec::{DecodeExt as _, EncodeExt as _}; + use tokio::io::{AsyncRead, AsyncWrite, ReadBuf}; + + use crate::host::dns::{DnsQuery, DnsRecordResolver, DnsResolver, DnsSrvRecord}; + use crate::socket::{ + NetNamespace, + tcp::{TcpConnectOptions, VirtualTcpSocket, VirtualTcpSocketFactory}, + udp::{UdpBindOptions, UdpSocketPurpose}, + }; + + use crate::packet::stun::Attribute; + use stun_codec::rfc5389::{attributes::XorMappedAddress, methods::BINDING}; + use stun_codec::{Message, MessageClass, MessageDecoder, MessageEncoder}; + + use super::*; + + struct MockUdpSocket { + local_addr: SocketAddr, + mapped_addr: SocketAddr, + responses: Mutex, SocketAddr)>>, + response_ready: tokio::sync::Notify, + } + + #[async_trait] + impl VirtualUdpSocket for MockUdpSocket { + fn local_addr(&self) -> io::Result { + Ok(self.local_addr) + } + + async fn send_to(&self, data: &[u8], addr: SocketAddr) -> io::Result { + let request = MessageDecoder::::new() + .decode_from_bytes(data) + .map_err(|error| io::Error::other(format!("{error:?}")))? + .map_err(|error| io::Error::other(format!("{error:?}")))?; + let mut response = Message::::new( + MessageClass::SuccessResponse, + BINDING, + request.transaction_id(), + ); + response.add_attribute(Attribute::XorMappedAddress(XorMappedAddress::new( + self.mapped_addr, + ))); + let bytes = MessageEncoder::new() + .encode_into_bytes(response) + .map_err(io::Error::other)?; + self.responses.lock().unwrap().push_back((bytes, addr)); + self.response_ready.notify_one(); + Ok(data.len()) + } + + async fn recv_from(&self, buf: &mut [u8]) -> io::Result<(usize, SocketAddr)> { + loop { + if let Some((bytes, addr)) = self.responses.lock().unwrap().pop_front() { + buf[..bytes.len()].copy_from_slice(&bytes); + return Ok((bytes.len(), addr)); + } + self.response_ready.notified().await; + } + } + } + + struct MockTcpSocket(tokio::io::DuplexStream); + + impl AsyncRead for MockTcpSocket { + fn poll_read( + mut self: Pin<&mut Self>, + cx: &mut Context<'_>, + buf: &mut ReadBuf<'_>, + ) -> Poll> { + Pin::new(&mut self.0).poll_read(cx, buf) + } + } + + impl AsyncWrite for MockTcpSocket { + fn poll_write( + mut self: Pin<&mut Self>, + cx: &mut Context<'_>, + buf: &[u8], + ) -> Poll> { + Pin::new(&mut self.0).poll_write(cx, buf) + } + + fn poll_flush(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + Pin::new(&mut self.0).poll_flush(cx) + } + + fn poll_shutdown(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + Pin::new(&mut self.0).poll_shutdown(cx) + } + } + + impl VirtualTcpSocket for MockTcpSocket { + fn local_addr(&self) -> io::Result { + Ok("127.0.0.1:40000".parse().unwrap()) + } + + fn peer_addr(&self) -> io::Result { + Ok("127.0.0.1:3478".parse().unwrap()) + } + } + + #[derive(Default)] + struct MockRuntime { + udp_binds: Mutex>, + } + + #[async_trait] + impl VirtualUdpSocketFactory for MockRuntime { + type Socket = MockUdpSocket; + + async fn bind_udp(&self, options: UdpBindOptions) -> anyhow::Result> { + self.udp_binds.lock().unwrap().push(options); + Ok(Arc::new(MockUdpSocket { + local_addr: "0.0.0.0:40000".parse().unwrap(), + mapped_addr: "198.51.100.10:40123".parse().unwrap(), + responses: Mutex::new(VecDeque::new()), + response_ready: tokio::sync::Notify::new(), + })) + } + } + + #[async_trait] + impl VirtualTcpSocketFactory for MockRuntime { + type Socket = MockTcpSocket; + + async fn connect_tcp(&self, _options: TcpConnectOptions) -> anyhow::Result { + anyhow::bail!("TCP is not used by this test") + } + } + + struct MockDns; + + #[async_trait] + impl DnsResolver for MockDns { + async fn resolve(&self, _query: DnsQuery) -> anyhow::Result> { + Ok(vec!["192.0.2.1".parse().unwrap()]) + } + } + + #[async_trait] + impl DnsRecordResolver for MockDns { + async fn resolve_txt(&self, _query: DnsQuery) -> anyhow::Result { + Ok(String::new()) + } + + async fn resolve_srv(&self, _query: DnsQuery) -> anyhow::Result> { + Ok(Vec::new()) + } + } + + #[test] + fn sampled_servers_keep_first_two_and_at_most_one_extra() { + let servers = ["a", "b", "c", "d"] + .into_iter() + .map(str::to_owned) + .collect::>(); + let sampled = sampled_servers(&servers); + assert_eq!(&sampled[..2], &["a", "b"]); + assert_eq!(sampled.len(), 3); + assert!(matches!(sampled[2].as_str(), "c" | "d")); + } + + #[test] + fn collector_keeps_udp_and_tcp_socket_contexts_separate() { + let udp_context = SocketContext::default() + .with_socket_mark(Some(11)) + .with_netns(Some(NetNamespace::new("udp-instance"))); + let tcp_context = SocketContext::default() + .with_socket_mark(Some(22)) + .with_netns(Some(NetNamespace::new("tcp-instance"))); + let collector = StunInfoCollector::new_with_socket_contexts( + Arc::new(MockRuntime::default()), + Arc::new(MockDns), + udp_context.clone(), + tcp_context.clone(), + Vec::new(), + Vec::new(), + Vec::new(), + ); + + assert_eq!(collector.udp_socket_context, udp_context); + assert_eq!(collector.tcp_socket_context, tcp_context); + } + + #[tokio::test] + async fn udp_mapping_uses_portable_runtime_and_instance_context() { + let runtime = Arc::new(MockRuntime::default()); + let context = SocketContext::default() + .with_socket_mark(Some(0)) + .with_netns(Some(NetNamespace::new("instance-a"))); + let collector = StunInfoCollector::new( + runtime.clone(), + Arc::new(MockDns), + context, + vec!["stun.example".to_owned()], + Vec::new(), + Vec::new(), + ); + + let mapped = collector.get_udp_port_mapping(0).await.unwrap(); + + assert_eq!(mapped, "198.51.100.10:40123".parse().unwrap()); + let binds = runtime.udp_binds.lock().unwrap(); + assert!(!binds.is_empty()); + assert_eq!(binds[0].purpose, UdpSocketPurpose::StunProbe); + assert_eq!(binds[0].local_addr, Some("0.0.0.0:0".parse().unwrap())); + assert_eq!(binds[0].context.ip_version, IpVersion::V4); + assert_eq!(binds[0].context.socket_mark, Some(0)); + assert_eq!( + binds[0].context.netns.as_ref().map(|netns| netns.token()), + Some("instance-a") + ); + } +} diff --git a/easytier-core/src/connectivity/stun/mod.rs b/easytier-core/src/connectivity/stun/mod.rs new file mode 100644 index 00000000..1c349b46 --- /dev/null +++ b/easytier-core/src/connectivity/stun/mod.rs @@ -0,0 +1,9 @@ +mod client; +mod collector; +mod responder; + +pub use client::{ + StunDnsRuntime, StunNatTypeDetectResult, StunSocketRuntime, TcpNatTypeDetector, + UdpNatTypeDetector, +}; +pub use collector::{StunInfoCollector, StunInfoProvider, StunServerConfig, StunSocketMapper}; diff --git a/easytier-core/src/connectivity/stun/responder.rs b/easytier-core/src/connectivity/stun/responder.rs new file mode 100644 index 00000000..42cc9402 --- /dev/null +++ b/easytier-core/src/connectivity/stun/responder.rs @@ -0,0 +1,247 @@ +//! STUN binding-response support over UDP sockets. + +use std::net::SocketAddr; +use std::sync::Arc; + +use anyhow::Context as _; +use bytecodec::{DecodeExt as _, EncodeExt as _}; +use stun_codec::rfc5389::attributes::XorMappedAddress; +use stun_codec::rfc5389::methods::BINDING; +use stun_codec::{Message, MessageClass, MessageDecoder, MessageEncoder}; + +use crate::packet::stun::{Attribute, ChangeRequest, tid_to_u32, u32_to_tid}; +use crate::socket::udp::{UdpBindOptions, VirtualUdpSocket, VirtualUdpSocketFactory}; + +#[derive(Debug, Clone, PartialEq, Eq)] +pub enum StunResponseSendSource { + SameSocket, + NewSocket, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct StunResponse { + pub bytes: Vec, + pub send_source: StunResponseSendSource, +} + +pub fn build_stun_response(addr: SocketAddr, req_buf: &[u8]) -> anyhow::Result { + let mut decoder = MessageDecoder::::new(); + let req_msg = decoder + .decode_from_bytes(req_buf) + .map_err(|e| anyhow::anyhow!("stun decode error: {:?}", e))? + .map_err(|e| anyhow::anyhow!("stun decode broken message error: {:?}", e))?; + + let tid = req_msg.transaction_id(); + // we only respond easytier stun req, whose tid has 0xdeadbeef prefix + if tid.as_bytes()[0..4] != [0xde, 0xad, 0xbe, 0xef] { + anyhow::bail!("stun req tid not from easytier"); + } + + let mut resp_msg = Message::::new( + MessageClass::SuccessResponse, + BINDING, + // we discard the prefix, make sure our implementation is not compatible with other stun client + u32_to_tid(tid_to_u32(&tid)), + ); + resp_msg.add_attribute(Attribute::XorMappedAddress(XorMappedAddress::new(addr))); + + let mut encoder = MessageEncoder::new(); + let bytes = encoder + .encode_into_bytes(resp_msg.clone()) + .map_err(|e| anyhow::anyhow!("stun encode error: {:?}", e))?; + + let change_req = req_msg + .get_attribute::() + .map(|r| r.ip() || r.port()) + .unwrap_or(false); + + Ok(StunResponse { + bytes, + send_source: if change_req { + StunResponseSendSource::NewSocket + } else { + StunResponseSendSource::SameSocket + }, + }) +} + +async fn respond_stun_packet( + socket: Arc, + factory: &F, + addr: SocketAddr, + req_buf: &[u8], +) -> anyhow::Result<()> +where + S: VirtualUdpSocket, + F: VirtualUdpSocketFactory + ?Sized, +{ + let response = build_stun_response(addr, req_buf)?; + match response.send_source { + StunResponseSendSource::SameSocket => { + socket + .send_to(&response.bytes, addr) + .await + .with_context(|| "send stun response error")?; + } + StunResponseSendSource::NewSocket => { + let bind_addr = if addr.is_ipv4() { + "0.0.0.0:0".parse().unwrap() + } else { + "[::]:0".parse().unwrap() + }; + let socket = factory + .bind_udp( + UdpBindOptions::hole_punch_control() + .with_context(socket.socket_context()) + .with_local_addr(Some(bind_addr)), + ) + .await?; + socket.send_to(&response.bytes, addr).await?; + } + } + + tracing::debug!(?addr, "udp respond stun packet done"); + Ok(()) +} + +#[async_trait::async_trait] +impl crate::socket::udp::UdpSessionStunResponder for F +where + S: VirtualUdpSocket, + F: VirtualUdpSocketFactory, +{ + async fn respond_stun( + &self, + socket: Arc, + datagram: &[u8], + remote_addr: SocketAddr, + ) -> std::io::Result<()> { + respond_stun_packet(socket, self, remote_addr, datagram) + .await + .map_err(|error| std::io::Error::other(error.to_string())) + } +} + +#[cfg(test)] +mod tests { + use std::{io, sync::Mutex}; + + use async_trait::async_trait; + use stun_codec::TransactionId; + + use super::*; + + #[derive(Debug, Default)] + struct MockSocket { + sent: Mutex, SocketAddr)>>, + } + + impl MockSocket { + fn sent(&self) -> Vec<(Vec, SocketAddr)> { + self.sent.lock().unwrap().clone() + } + } + + #[async_trait] + impl VirtualUdpSocket for MockSocket { + fn local_addr(&self) -> io::Result { + Ok("127.0.0.1:0".parse().unwrap()) + } + + async fn send_to(&self, data: &[u8], addr: SocketAddr) -> io::Result { + self.sent.lock().unwrap().push((data.to_vec(), addr)); + Ok(data.len()) + } + + async fn recv_from(&self, _buf: &mut [u8]) -> io::Result<(usize, SocketAddr)> { + std::future::pending().await + } + } + + #[derive(Debug, Default)] + struct MockFactory { + bind_options: Mutex>, + sockets: Mutex>>, + } + + impl MockFactory { + fn bind_options(&self) -> Vec { + self.bind_options.lock().unwrap().clone() + } + + fn sockets(&self) -> Vec> { + self.sockets.lock().unwrap().clone() + } + } + + #[async_trait] + impl VirtualUdpSocketFactory for MockFactory { + type Socket = MockSocket; + + async fn bind_udp(&self, options: UdpBindOptions) -> anyhow::Result> { + self.bind_options.lock().unwrap().push(options); + let socket = Arc::new(MockSocket::default()); + self.sockets.lock().unwrap().push(socket.clone()); + Ok(socket) + } + } + + fn stun_request(change_ip: bool, change_port: bool) -> Vec { + let mut request = Message::::new(MessageClass::Request, BINDING, u32_to_tid(7)); + if change_ip || change_port { + request.add_attribute(Attribute::ChangeRequest(ChangeRequest::new( + change_ip, + change_port, + ))); + } + MessageEncoder::new().encode_into_bytes(request).unwrap() + } + + #[test] + fn build_stun_response_rejects_non_easytier_tid() { + let request = + Message::::new(MessageClass::Request, BINDING, TransactionId::new([0; 12])); + let mut encoder = MessageEncoder::new(); + let request = encoder.encode_into_bytes(request).unwrap(); + + assert!(build_stun_response("127.0.0.1:1234".parse().unwrap(), &request).is_err()); + } + + #[test] + fn build_stun_response_detects_change_request() { + let request = stun_request(true, false); + + let response = build_stun_response("127.0.0.1:1234".parse().unwrap(), &request).unwrap(); + + assert_eq!(response.send_source, StunResponseSendSource::NewSocket); + assert!(!response.bytes.is_empty()); + } + + #[tokio::test] + async fn respond_stun_packet_uses_ipv6_unspecified_socket_for_ipv6_change_request() { + let listener_socket = Arc::new(MockSocket::default()); + let factory = MockFactory::default(); + let remote_addr = "[::1]:1234".parse().unwrap(); + + respond_stun_packet( + listener_socket.clone(), + &factory, + remote_addr, + &stun_request(true, false), + ) + .await + .unwrap(); + + assert!(listener_socket.sent().is_empty()); + assert_eq!( + factory.bind_options(), + vec![ + UdpBindOptions::hole_punch_control() + .with_local_addr(Some("[::]:0".parse().unwrap())) + ] + ); + let sockets = factory.sockets(); + assert_eq!(sockets.len(), 1); + assert_eq!(sockets[0].sent()[0].1, remote_addr); + } +} diff --git a/easytier-core/src/connectivity/transport/mod.rs b/easytier-core/src/connectivity/transport/mod.rs new file mode 100644 index 00000000..6bf41c25 --- /dev/null +++ b/easytier-core/src/connectivity/transport/mod.rs @@ -0,0 +1,70 @@ +use std::future::Future; + +use futures::{StreamExt, stream::FuturesUnordered}; +use url::Url; + +mod tcp; +mod udp; + +pub(crate) use tcp::connect_tcp; +pub use udp::{ConnectedUdpSession, UdpSessionMode, connect_udp}; + +/// A host-created non-IP byte stream with host-provided endpoint metadata. +/// +/// Unix and in-process transports use the same stream framing boundary as TCP, +/// but their endpoints cannot be represented by `SocketAddr`. +pub struct ConnectedByteStream { + socket: S, + local_url: Option, + remote_url: Url, + resolved_remote_url: Option, +} + +impl ConnectedByteStream { + pub fn new( + socket: S, + local_url: Option, + remote_url: Url, + resolved_remote_url: Option, + ) -> Self { + Self { + socket, + local_url, + remote_url, + resolved_remote_url, + } + } + + pub fn into_parts(self) -> (S, Option, Url, Option) { + ( + self.socket, + self.local_url, + self.remote_url, + self.resolved_remote_url, + ) + } +} + +/// A transport endpoint established by a connectivity strategy. +/// +/// Manual, direct, and hole-punch strategies stop at this boundary. Protocol +/// code consumes the endpoint and upgrades it into an EasyTier tunnel. +pub enum ConnectedTransport { + Tcp(TcpSocket), + Udp(ConnectedUdpSession), + ByteStream(ConnectedByteStream), +} + +async fn first_success(mut futures: FuturesUnordered) -> anyhow::Result +where + F: Future> + Send, +{ + let mut last_error = None; + while let Some(result) = futures.next().await { + match result { + Ok(value) => return Ok(value), + Err(error) => last_error = Some(error), + } + } + Err(last_error.unwrap_or_else(|| anyhow::anyhow!("no transport candidates"))) +} diff --git a/easytier-core/src/connectivity/transport/tcp.rs b/easytier-core/src/connectivity/transport/tcp.rs new file mode 100644 index 00000000..6edc3707 --- /dev/null +++ b/easytier-core/src/connectivity/transport/tcp.rs @@ -0,0 +1,53 @@ +use std::{net::SocketAddr, sync::Arc}; + +use futures::stream::FuturesUnordered; + +use crate::socket::{ + IpVersion, + tcp::{TcpBindOptions, TcpConnectOptions, TcpSocketPurpose, VirtualTcpSocketFactory}, +}; + +use super::first_success; + +pub async fn connect_tcp( + host: Arc, + remote_addr: SocketAddr, + bind_addrs: Vec, + default_bind: TcpBindOptions, + purpose: TcpSocketPurpose, +) -> anyhow::Result +where + H: VirtualTcpSocketFactory, +{ + let ip_version = if remote_addr.is_ipv4() { + IpVersion::V4 + } else { + IpVersion::V6 + }; + let default_bind = default_bind.with_ip_version(ip_version); + let futures = FuturesUnordered::new(); + if bind_addrs.is_empty() { + futures.push( + host.connect_tcp( + TcpConnectOptions::direct_connect(remote_addr) + .with_purpose(purpose) + .with_bind(default_bind), + ), + ); + } else { + for bind_addr in bind_addrs { + let bind = default_bind + .clone() + .with_local_addr(Some(bind_addr)) + .with_only_v6(true); + futures.push( + host.connect_tcp( + TcpConnectOptions::direct_connect(remote_addr) + .with_purpose(purpose) + .with_bind(bind), + ), + ); + } + } + first_success(futures).await +} diff --git a/easytier-core/src/connectivity/transport/udp.rs b/easytier-core/src/connectivity/transport/udp.rs new file mode 100644 index 00000000..8a48d7a5 --- /dev/null +++ b/easytier-core/src/connectivity/transport/udp.rs @@ -0,0 +1,102 @@ +use std::{net::SocketAddr, sync::Arc}; + +use futures::stream::FuturesUnordered; + +use crate::socket::{ + IpVersion, + udp::{ + UdpBindOptions, UdpSession, UdpSessionLayer, UdpSessionProtocol, VirtualUdpSocketFactory, + }, +}; + +use super::first_success; + +pub struct ConnectedUdpSession { + session: UdpSession, + keep_alive: Box, +} + +impl ConnectedUdpSession { + pub fn new(session: UdpSession, keep_alive: T) -> Self + where + T: Send + Sync + 'static, + { + Self { + session, + keep_alive: Box::new(keep_alive), + } + } + + pub fn session(&self) -> &UdpSession { + &self.session + } + + pub fn into_parts(self) -> (UdpSession, Box) { + (self.session, self.keep_alive) + } +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum UdpSessionMode { + EasyTierMux, + Classified(UdpSessionProtocol), +} + +pub async fn connect_udp( + host: Arc, + remote_addr: SocketAddr, + bind_addrs: Vec, + default_bind: UdpBindOptions, + mode: UdpSessionMode, +) -> anyhow::Result +where + H: VirtualUdpSocketFactory, +{ + let ip_version = if remote_addr.is_ipv4() { + IpVersion::V4 + } else { + IpVersion::V6 + }; + let default_bind = default_bind.with_ip_version(ip_version); + let futures = FuturesUnordered::new(); + if bind_addrs.is_empty() { + let local_addr = if remote_addr.is_ipv4() { + "0.0.0.0:0".parse().expect("static IPv4 bind address") + } else { + "[::]:0".parse().expect("static IPv6 bind address") + }; + let bind = default_bind + .with_local_addr(Some(local_addr)) + .with_only_v6(true); + futures.push(bind_and_connect(host.clone(), bind, remote_addr, mode)); + } else { + for bind_addr in bind_addrs { + let bind = default_bind + .clone() + .with_local_addr(Some(bind_addr)) + .with_only_v6(true); + futures.push(bind_and_connect(host.clone(), bind, remote_addr, mode)); + } + } + first_success(futures).await +} + +async fn bind_and_connect( + host: Arc, + bind: UdpBindOptions, + remote_addr: SocketAddr, + mode: UdpSessionMode, +) -> anyhow::Result +where + H: VirtualUdpSocketFactory, +{ + let socket = host.bind_udp(bind).await?; + let layer = Arc::new(UdpSessionLayer::new_with_stun_responder(socket, host)); + let session = match mode { + UdpSessionMode::EasyTierMux => layer.connect(remote_addr).await?, + UdpSessionMode::Classified(protocol) => { + layer.open_classified_session(protocol, remote_addr)? + } + }; + Ok(ConnectedUdpSession::new(session, layer)) +} diff --git a/easytier-core/src/events.rs b/easytier-core/src/events.rs new file mode 100644 index 00000000..2fb8705f --- /dev/null +++ b/easytier-core/src/events.rs @@ -0,0 +1,104 @@ +use std::sync::Arc; + +use cidr::{Ipv4Cidr, Ipv6Inet}; +use url::Url; + +use crate::{ + config::{PeerId, gateway::PortForwardConfig}, + socket::{IpVersion, ListenerConnectionCounter}, +}; + +/// Notifications emitted by the portable core runtime to its host. +#[derive(Debug, Clone)] +pub enum CoreEvent { + PeerAdded(PeerId), + PeerRemoved(PeerId), + PeerConnAdded(easytier_proto::core_peer::peer::PeerConnInfo), + PeerConnRemoved(easytier_proto::core_peer::peer::PeerConnInfo), + CredentialChanged, + + ManualConnecting { + url: Url, + }, + ManualConnectError { + url: Url, + ip_version: IpVersion, + error: String, + }, + + ListenerPlanFailed { + url: Url, + error: String, + }, + ListenerAdded { + url: Url, + connection_counter: Arc, + }, + ListenerRemoved { + url: Url, + }, + ListenerAddFailed { + url: Url, + error: String, + retry_count: usize, + will_retry: bool, + }, + ListenerAcceptFailed { + url: Url, + error: String, + }, + ListenerSocketAccepted { + url: Url, + }, + ListenerAcceptedSocketHandleFailed { + url: Url, + error: String, + }, + + TunnelAccepted { + local_url: String, + remote_url: String, + }, + TunnelAdmissionFailed { + local_url: String, + remote_url: String, + error: String, + }, + UdpPortMappingEstablished { + local_listener: Url, + mapped_listener: Url, + backend: String, + }, + + ProxyCidrsUpdated { + added: Vec, + removed: Vec, + }, + PublicIpv6LeaseChanged { + old: Option, + new: Option, + }, + PublicIpv6RoutesChanged { + added: Vec, + removed: Vec, + }, + + VpnPortalStarted(String), + VpnPortalClientConnected { + portal: String, + client: String, + }, + VpnPortalClientDisconnected { + portal: String, + client: String, + }, + GatewayPortForwardAdded(PortForwardConfig), +} + +pub trait CoreEventSink: Send + Sync + 'static { + fn emit(&self, event: CoreEvent); +} + +impl CoreEventSink for () { + fn emit(&self, _event: CoreEvent) {} +} diff --git a/easytier-core/src/foundation/mod.rs b/easytier-core/src/foundation/mod.rs new file mode 100644 index 00000000..967a6ad8 --- /dev/null +++ b/easytier-core/src/foundation/mod.rs @@ -0,0 +1,11 @@ +//! Infrastructure Modules with no domain dependency. +//! +//! Everything in `foundation` may be used by any layer, and nothing here may +//! depend on a domain Module. See `CONTEXT.md` "Module layers". + +#[cfg(any(feature = "proxy-smoltcp-stack", test))] +pub(crate) mod operation_broker; +pub mod stats; +pub(crate) mod task; +pub(crate) mod time; +pub(crate) mod token_bucket; diff --git a/easytier-core/src/foundation/operation_broker.rs b/easytier-core/src/foundation/operation_broker.rs new file mode 100644 index 00000000..84f0e796 --- /dev/null +++ b/easytier-core/src/foundation/operation_broker.rs @@ -0,0 +1,471 @@ +//! Domain-neutral lifecycle and completion storage for externally submitted +//! asynchronous operations. + +use std::{ + collections::{HashMap, VecDeque}, + hash::Hash, +}; + +use tokio_util::sync::CancellationToken; + +#[derive(Clone, Copy, Debug, Eq, Hash, Ord, PartialEq, PartialOrd)] +#[repr(transparent)] +pub(crate) struct OperationId(u64); + +impl OperationId { + pub(crate) fn from_raw(value: u64) -> Option { + (value != 0).then_some(Self(value)) + } + + pub(crate) fn get(self) -> u64 { + self.0 + } +} + +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +pub(crate) enum AdmissionError { + AtCapacity, + IdExhausted, +} + +pub(crate) struct Admission { + pub(crate) id: OperationId, + pub(crate) cancellation: CancellationToken, +} + +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +pub(crate) enum AccessError { + Missing, + NotDrained, +} + +pub(crate) struct Completion { + pub(crate) operation_id: OperationId, + pub(crate) kind: K, + pub(crate) status: S, +} + +pub(crate) struct ReleasedOperation { + pub(crate) metadata: M, + pub(crate) outcome: Option, +} + +pub(crate) struct TakenOperation { + pub(crate) metadata: M, + pub(crate) value: T, +} + +enum OperationState { + Pending, + Queued(O), + Drained(O), + Discarding, +} + +struct Operation { + kind: K, + metadata: Option, + cancellation: CancellationToken, + state: OperationState, +} + +pub(crate) struct OperationBroker { + max_operations: usize, + next_operation_id: u64, + operations: HashMap>, + completions: VecDeque, + wake_generation: u64, +} + +impl OperationBroker +where + K: Copy, +{ + pub(crate) fn new(max_operations: usize) -> Self { + Self { + max_operations, + next_operation_id: 1, + operations: HashMap::new(), + completions: VecDeque::new(), + wake_generation: 0, + } + } + + pub(crate) fn admit(&mut self, kind: K, metadata: M) -> Result { + if self.operations.len() >= self.max_operations { + return Err(AdmissionError::AtCapacity); + } + let id = + OperationId::from_raw(self.next_operation_id).ok_or(AdmissionError::IdExhausted)?; + self.next_operation_id = self + .next_operation_id + .checked_add(1) + .ok_or(AdmissionError::IdExhausted)?; + let cancellation = CancellationToken::new(); + let replaced = self.operations.insert( + id, + Operation { + kind, + metadata: Some(metadata), + cancellation: cancellation.clone(), + state: OperationState::Pending, + }, + ); + debug_assert!(replaced.is_none()); + Ok(Admission { id, cancellation }) + } + + #[cfg(test)] + pub(crate) fn len(&self) -> usize { + self.operations.len() + } + + pub(crate) fn complete_with( + &mut self, + operation_id: OperationId, + complete: impl FnOnce(K, &mut M) -> O, + ) -> bool { + self.resolve_pending_with(operation_id, false, complete) + } + + pub(crate) fn cancel_with( + &mut self, + operation_id: OperationId, + complete: impl FnOnce(K, &mut M) -> O, + ) -> bool { + self.resolve_pending_with(operation_id, true, complete) + } + + fn resolve_pending_with( + &mut self, + operation_id: OperationId, + cancel: bool, + complete: impl FnOnce(K, &mut M) -> O, + ) -> bool { + let Some(mut operation) = self.operations.remove(&operation_id) else { + return false; + }; + match operation.state { + OperationState::Pending => { + if cancel { + operation.cancellation.cancel(); + } + let outcome = complete( + operation.kind, + operation + .metadata + .as_mut() + .expect("pending operation metadata is present"), + ); + let notify = self.completions.is_empty(); + operation.state = OperationState::Queued(outcome); + self.operations.insert(operation_id, operation); + self.completions.push_back(operation_id); + notify + } + OperationState::Discarding => { + if cancel { + self.operations.insert(operation_id, operation); + } + false + } + OperationState::Queued(_) | OperationState::Drained(_) => { + self.operations.insert(operation_id, operation); + false + } + } + } + + pub(crate) fn free(&mut self, operation_id: OperationId) -> Option> { + let mut operation = self.operations.remove(&operation_id)?; + let state = std::mem::replace(&mut operation.state, OperationState::Discarding); + match state { + OperationState::Pending => { + operation.cancellation.cancel(); + let released = ReleasedOperation { + metadata: operation + .metadata + .take() + .expect("pending operation metadata is present"), + outcome: None, + }; + self.operations.insert(operation_id, operation); + Some(released) + } + OperationState::Queued(outcome) => { + self.completions + .retain(|completion| *completion != operation_id); + Some(ReleasedOperation { + metadata: operation + .metadata + .take() + .expect("queued operation metadata is present"), + outcome: Some(outcome), + }) + } + OperationState::Drained(outcome) => Some(ReleasedOperation { + metadata: operation + .metadata + .take() + .expect("drained operation metadata is present"), + outcome: Some(outcome), + }), + OperationState::Discarding => { + operation.state = OperationState::Discarding; + self.operations.insert(operation_id, operation); + None + } + } + } + + pub(crate) fn drain( + &mut self, + max_count: usize, + mut status: impl FnMut(&O) -> S, + ) -> Vec> { + let mut completions = Vec::with_capacity(max_count.min(self.completions.len())); + while completions.len() < max_count { + let Some(operation_id) = self.completions.pop_front() else { + break; + }; + let Some(operation) = self.operations.get_mut(&operation_id) else { + continue; + }; + let old_state = std::mem::replace(&mut operation.state, OperationState::Discarding); + match old_state { + OperationState::Queued(outcome) => { + completions.push(Completion { + operation_id, + kind: operation.kind, + status: status(&outcome), + }); + operation.state = OperationState::Drained(outcome); + } + other => operation.state = other, + } + } + completions + } + + pub(crate) fn has_completions(&self) -> bool { + !self.completions.is_empty() + } + + pub(crate) fn with_drained( + &self, + operation_id: OperationId, + inspect: impl FnOnce(K, &M, &O) -> T, + ) -> Result { + let operation = self + .operations + .get(&operation_id) + .ok_or(AccessError::Missing)?; + match &operation.state { + OperationState::Drained(outcome) => Ok(inspect( + operation.kind, + operation + .metadata + .as_ref() + .expect("drained operation metadata is present"), + outcome, + )), + OperationState::Pending | OperationState::Queued(_) | OperationState::Discarding => { + Err(AccessError::NotDrained) + } + } + } + + pub(crate) fn take_with( + &mut self, + operation_id: OperationId, + take: impl FnOnce(&O) -> Option, + ) -> Result>, AccessError> { + let Some(mut operation) = self.operations.remove(&operation_id) else { + return Err(AccessError::Missing); + }; + let old_state = std::mem::replace(&mut operation.state, OperationState::Discarding); + let outcome = match old_state { + OperationState::Drained(outcome) => outcome, + other => { + operation.state = other; + self.operations.insert(operation_id, operation); + return Err(AccessError::NotDrained); + } + }; + let Some(value) = take(&outcome) else { + operation.state = OperationState::Drained(outcome); + self.operations.insert(operation_id, operation); + return Ok(None); + }; + Ok(Some(TakenOperation { + metadata: operation + .metadata + .take() + .expect("drained operation metadata is present"), + value, + })) + } + + pub(crate) fn pending_ids(&self) -> Vec { + self.operations + .iter() + .filter_map(|(operation_id, operation)| { + matches!(operation.state, OperationState::Pending).then_some(*operation_id) + }) + .collect() + } + + pub(crate) fn wake_generation(&self) -> u64 { + self.wake_generation + } + + pub(crate) fn invalidate_waiters(&mut self) { + self.wake_generation = self.wake_generation.wrapping_add(1); + } + + pub(crate) fn discard_all(&mut self) { + for operation in self.operations.values() { + operation.cancellation.cancel(); + } + self.operations.clear(); + self.completions.clear(); + self.invalidate_waiters(); + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[derive(Clone, Copy, Debug, Eq, PartialEq)] + enum Kind { + Read, + Write, + } + + #[test] + fn completion_is_drained_and_taken_once() { + let mut broker = OperationBroker::new(4); + let admission = broker.admit(Kind::Read, "metadata").unwrap(); + + assert!(broker.complete_with(admission.id, |_, _| Ok::<_, u8>(7))); + let completions = broker.drain(4, |outcome| outcome.is_ok()); + assert_eq!(completions.len(), 1); + assert_eq!(completions[0].operation_id, admission.id); + assert_eq!(completions[0].kind, Kind::Read); + assert!(completions[0].status); + assert!(broker.drain(4, |_| true).is_empty()); + + assert!( + broker + .take_with(admission.id, |_| None::) + .unwrap() + .is_none() + ); + let taken = broker + .take_with(admission.id, |outcome| outcome.as_ref().ok().copied()) + .unwrap() + .unwrap(); + assert_eq!(taken.metadata, "metadata"); + assert_eq!(taken.value, 7); + assert!(matches!( + broker.take_with(admission.id, |_| Some(())), + Err(AccessError::Missing) + )); + } + + #[test] + fn free_pending_absorbs_late_completion() { + let mut broker = OperationBroker::new(4); + let admission = broker.admit(Kind::Write, 9).unwrap(); + + let released = broker.free(admission.id).unwrap(); + assert_eq!(released.metadata, 9); + assert!(released.outcome.is_none()); + assert!(admission.cancellation.is_cancelled()); + assert!(!broker.complete_with(admission.id, |_, _| 3)); + + assert_eq!(broker.len(), 0); + assert!(!broker.has_completions()); + } + + #[test] + fn cancellation_and_completion_have_one_terminal_outcome() { + let mut cancelled_first = OperationBroker::new(4); + let cancelled = cancelled_first.admit(Kind::Read, ()).unwrap(); + assert!(cancelled_first.cancel_with(cancelled.id, |_, _| Err::("cancelled"))); + assert!(!cancelled_first.complete_with(cancelled.id, |_, _| Ok(1))); + let completions = cancelled_first.drain(4, |outcome| outcome.is_ok()); + assert_eq!(completions.len(), 1); + assert!(!completions[0].status); + + let mut completed_first = OperationBroker::new(4); + let completed = completed_first.admit(Kind::Read, ()).unwrap(); + assert!(completed_first.complete_with(completed.id, |_, _| Ok::<_, &str>(1))); + assert!(!completed_first.cancel_with(completed.id, |_, _| Err("cancelled"))); + let completions = completed_first.drain(4, |outcome| outcome.is_ok()); + assert_eq!(completions.len(), 1); + assert!(completions[0].status); + } + + #[test] + fn completion_notification_is_an_empty_to_nonempty_edge() { + let mut broker = OperationBroker::new(4); + let first = broker.admit(Kind::Read, ()).unwrap(); + let second = broker.admit(Kind::Write, ()).unwrap(); + + assert!(broker.complete_with(first.id, |_, _| ())); + assert!(!broker.complete_with(second.id, |_, _| ())); + assert_eq!(broker.drain(2, |_| ()).len(), 2); + assert!(!broker.has_completions()); + + let third = broker.admit(Kind::Read, ()).unwrap(); + assert!(broker.complete_with(third.id, |_, _| ())); + } + + #[test] + fn admission_limit_counts_discarding_tombstones() { + let mut broker = OperationBroker::new(1); + let admission = broker.admit(Kind::Read, ()).unwrap(); + broker.free(admission.id); + + assert!(matches!( + broker.admit(Kind::Write, ()), + Err(AdmissionError::AtCapacity) + )); + + broker.complete_with(admission.id, |_, _| ()); + broker.admit(Kind::Write, ()).unwrap(); + } + + #[test] + fn cancellation_keeps_discarding_tombstone_until_completion() { + let mut broker: OperationBroker = OperationBroker::new(1); + let admission = broker.admit(Kind::Read, ()).unwrap(); + broker.free(admission.id); + + assert!(!broker.cancel_with(admission.id, |_, _| ())); + assert_eq!(broker.len(), 1); + assert!(matches!( + broker.admit(Kind::Write, ()), + Err(AdmissionError::AtCapacity) + )); + + assert!(!broker.complete_with(admission.id, |_, _| ())); + assert_eq!(broker.len(), 0); + broker.admit(Kind::Write, ()).unwrap(); + } + + #[test] + fn discard_invalidates_waiters_and_cancels_operations() { + let mut broker: OperationBroker = OperationBroker::new(1); + let admission = broker.admit(Kind::Read, ()).unwrap(); + let generation = broker.wake_generation(); + + broker.discard_all(); + + assert!(admission.cancellation.is_cancelled()); + assert_ne!(broker.wake_generation(), generation); + assert_eq!(broker.len(), 0); + } +} diff --git a/easytier/src/common/stats_manager.rs b/easytier-core/src/foundation/stats.rs similarity index 75% rename from easytier/src/common/stats_manager.rs rename to easytier-core/src/foundation/stats.rs index 5627e24a..5f7354a4 100644 --- a/easytier/src/common/stats_manager.rs +++ b/easytier-core/src/foundation/stats.rs @@ -3,11 +3,89 @@ use quanta::Instant; use serde::{Deserialize, Serialize}; use std::cell::UnsafeCell; use std::fmt; -use std::sync::Arc; +use std::sync::{Arc, Mutex}; use std::time::Duration; -use tokio::time::interval; use tokio_util::task::AbortOnDropHandle; +use crate::foundation::time::interval; + +#[derive(Clone, Debug, PartialEq, Eq)] +pub struct RpcMetricLabels { + pub network_name: String, + pub src_peer_id: u32, + pub dst_peer_id: u32, + pub service_name: String, + pub method_name: String, +} + +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub enum RpcMetricStatus { + Success, + Error, +} + +impl RpcMetricStatus { + pub fn as_str(self) -> &'static str { + match self { + Self::Success => "success", + Self::Error => "error", + } + } +} + +pub trait RpcMetrics: Send + Sync + 'static { + fn client_tx(&self, _labels: &RpcMetricLabels) {} + + fn client_rx(&self, _labels: &RpcMetricLabels, _duration_ms: u64) {} + + fn client_error( + &self, + _labels: &RpcMetricLabels, + _error_type: Option, + _duration_ms: u64, + ) { + } + + fn server_rx(&self, _labels: &RpcMetricLabels) {} + + fn server_tx(&self, _labels: &RpcMetricLabels, _duration_ms: u64) {} + + fn server_error( + &self, + _labels: &RpcMetricLabels, + _error_type: Option, + _duration_ms: u64, + ) { + } +} + +pub type ArcRpcMetrics = Arc; + +pub trait RpcMetricsProvider: Send + Sync + 'static { + fn into_rpc_metrics(self) -> Option; +} + +impl RpcMetricsProvider for () { + fn into_rpc_metrics(self) -> Option { + None + } +} + +impl RpcMetricsProvider for ArcRpcMetrics { + fn into_rpc_metrics(self) -> Option { + Some(self) + } +} + +impl RpcMetricsProvider for Arc +where + T: RpcMetrics, +{ + fn into_rpc_metrics(self) -> Option { + Some(self) + } +} + /// Predefined metric names for type safety #[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)] pub enum MetricName { @@ -473,13 +551,6 @@ impl MetricData { } } - fn new_with_value(initial: u64) -> Self { - Self { - counter: UnsafeCounter::new_with_value(initial), - last_updated: UnsafeCell::new(Instant::now()), - } - } - /// Update the last_updated timestamp /// # Safety /// This method is unsafe because it uses UnsafeCell. The caller must ensure @@ -600,17 +671,34 @@ impl MetricSnapshot { /// StatsManager manages global statistics with high performance counters pub struct StatsManager { counters: Arc>>, - cleanup_task: AbortOnDropHandle<()>, + cleanup_task: Mutex>>, } impl StatsManager { /// Create a new StatsManager pub fn new() -> Self { - let counters = Arc::new(DashMap::new()); + let manager = Self { + counters: Arc::new(DashMap::new()), + cleanup_task: Mutex::new(None), + }; + manager.start_cleanup_task(); + manager + } - // Start cleanup task only if we're in a tokio runtime - let counters_clone = Arc::downgrade(&counters); - let cleanup_task = tokio::spawn(async move { + pub(crate) fn start_cleanup_task(&self) { + let mut cleanup_task = self.cleanup_task.lock().unwrap(); + if cleanup_task + .as_ref() + .is_some_and(|task| !task.is_finished()) + { + return; + } + cleanup_task.take(); + let Ok(runtime) = tokio::runtime::Handle::try_current() else { + return; + }; + let counters = Arc::downgrade(&self.counters); + *cleanup_task = Some(AbortOnDropHandle::new(runtime.spawn(async move { let mut interval = interval(Duration::from_secs(60)); // Check every minute loop { interval.tick().await; @@ -619,7 +707,7 @@ impl StatsManager { continue; }; - let Some(counters) = counters_clone.upgrade() else { + let Some(counters) = counters.upgrade() else { break; }; @@ -629,11 +717,14 @@ impl StatsManager { }); counters.shrink_to_fit(); } - }); + }))); + } - Self { - counters, - cleanup_task: AbortOnDropHandle::new(cleanup_task), + pub(crate) async fn stop_cleanup_task(&self) { + let task = self.cleanup_task.lock().unwrap().take(); + if let Some(task) = task { + task.abort(); + let _ = task.await; } } @@ -650,11 +741,6 @@ impl StatsManager { CounterHandle::new(metric_data, key) } - /// Get a counter with no labels - pub fn get_simple_counter(&self, name: MetricName) -> CounterHandle { - self.get_counter(name, LabelSet::new()) - } - /// Get all metric snapshots pub fn get_all_metrics(&self) -> Vec { let mut metrics = Vec::new(); @@ -683,40 +769,11 @@ impl StatsManager { metrics } - /// Get metrics filtered by name prefix - pub fn get_metrics_by_prefix(&self, prefix: &str) -> Vec { - self.get_all_metrics() - .into_iter() - .filter(|m| m.name.to_string().starts_with(prefix)) - .collect() - } - - /// Get a specific metric by name and labels - pub fn get_metric(&self, name: MetricName, labels: &LabelSet) -> Option { - let key = MetricKey::new(name, labels.clone()); - - if let Some(metric_data) = self.counters.get(&key) { - let value = unsafe { metric_data.counter.get() }; - Some(MetricSnapshot { - name, - labels: labels.clone(), - value, - }) - } else { - None - } - } - /// Clear all metrics pub fn clear(&self) { self.counters.clear(); } - /// Get the number of tracked metrics - pub fn metric_count(&self) -> usize { - self.counters.len() - } - /// Export metrics in Prometheus format pub fn export_prometheus(&self) -> String { let metrics = self.get_all_metrics(); @@ -761,14 +818,257 @@ impl Default for StatsManager { } } +pub struct StatsRpcMetrics { + stats_manager: Arc, +} + +impl StatsRpcMetrics { + pub fn new(stats_manager: Arc) -> Self { + Self { stats_manager } + } +} + +fn rpc_base_labels(labels: &RpcMetricLabels) -> LabelSet { + LabelSet::new() + .with_label_type(LabelType::NetworkName(labels.network_name.clone())) + .with_label_type(LabelType::SrcPeerId(labels.src_peer_id)) + .with_label_type(LabelType::DstPeerId(labels.dst_peer_id)) + .with_label_type(LabelType::ServiceName(labels.service_name.clone())) + .with_label_type(LabelType::MethodName(labels.method_name.clone())) +} + +fn rpc_labels_with_status(labels: &RpcMetricLabels, status: RpcMetricStatus) -> LabelSet { + rpc_base_labels(labels).with_label_type(LabelType::Status(status.as_str().to_string())) +} + +fn record_rpc_client_tx(stats_manager: &StatsManager, labels: &RpcMetricLabels) { + stats_manager + .get_counter(MetricName::PeerRpcClientTx, rpc_base_labels(labels)) + .inc(); +} + +fn record_rpc_client_rx(stats_manager: &StatsManager, labels: &RpcMetricLabels, duration_ms: u64) { + let labels = rpc_labels_with_status(labels, RpcMetricStatus::Success); + stats_manager + .get_counter(MetricName::PeerRpcClientRx, labels.clone()) + .inc(); + stats_manager + .get_counter(MetricName::PeerRpcDuration, labels) + .add(duration_ms); +} + +fn record_rpc_client_error( + stats_manager: &StatsManager, + labels: &RpcMetricLabels, + error_type: Option, + duration_ms: u64, +) { + let mut labels = rpc_labels_with_status(labels, RpcMetricStatus::Error); + if let Some(error_type) = error_type { + labels = labels.with_label_type(LabelType::ErrorType(error_type)); + } + stats_manager + .get_counter(MetricName::PeerRpcErrors, labels.clone()) + .inc(); + stats_manager + .get_counter(MetricName::PeerRpcDuration, labels) + .add(duration_ms); +} + +fn record_rpc_server_rx(stats_manager: &StatsManager, labels: &RpcMetricLabels) { + stats_manager + .get_counter(MetricName::PeerRpcServerRx, rpc_base_labels(labels)) + .inc(); +} + +fn record_rpc_server_tx(stats_manager: &StatsManager, labels: &RpcMetricLabels, duration_ms: u64) { + let labels = rpc_labels_with_status(labels, RpcMetricStatus::Success); + stats_manager + .get_counter(MetricName::PeerRpcServerTx, labels.clone()) + .inc(); + stats_manager + .get_counter(MetricName::PeerRpcDuration, labels) + .add(duration_ms); +} + +fn record_rpc_server_error( + stats_manager: &StatsManager, + labels: &RpcMetricLabels, + duration_ms: u64, +) { + let labels = rpc_labels_with_status(labels, RpcMetricStatus::Error); + stats_manager + .get_counter(MetricName::PeerRpcErrors, labels.clone()) + .inc(); + stats_manager + .get_counter(MetricName::PeerRpcDuration, labels) + .add(duration_ms); +} + +impl RpcMetrics for StatsRpcMetrics { + fn client_tx(&self, labels: &RpcMetricLabels) { + record_rpc_client_tx(&self.stats_manager, labels); + } + + fn client_rx(&self, labels: &RpcMetricLabels, duration_ms: u64) { + record_rpc_client_rx(&self.stats_manager, labels, duration_ms); + } + + fn client_error(&self, labels: &RpcMetricLabels, error_type: Option, duration_ms: u64) { + record_rpc_client_error(&self.stats_manager, labels, error_type, duration_ms); + } + + fn server_rx(&self, labels: &RpcMetricLabels) { + record_rpc_server_rx(&self.stats_manager, labels); + } + + fn server_tx(&self, labels: &RpcMetricLabels, duration_ms: u64) { + record_rpc_server_tx(&self.stats_manager, labels, duration_ms); + } + + fn server_error( + &self, + labels: &RpcMetricLabels, + _error_type: Option, + duration_ms: u64, + ) { + record_rpc_server_error(&self.stats_manager, labels, duration_ms); + } +} + +impl RpcMetrics for StatsManager { + fn client_tx(&self, labels: &RpcMetricLabels) { + record_rpc_client_tx(self, labels); + } + + fn client_rx(&self, labels: &RpcMetricLabels, duration_ms: u64) { + record_rpc_client_rx(self, labels, duration_ms); + } + + fn client_error(&self, labels: &RpcMetricLabels, error_type: Option, duration_ms: u64) { + record_rpc_client_error(self, labels, error_type, duration_ms); + } + + fn server_rx(&self, labels: &RpcMetricLabels) { + record_rpc_server_rx(self, labels); + } + + fn server_tx(&self, labels: &RpcMetricLabels, duration_ms: u64) { + record_rpc_server_tx(self, labels, duration_ms); + } + + fn server_error( + &self, + labels: &RpcMetricLabels, + _error_type: Option, + duration_ms: u64, + ) { + record_rpc_server_error(self, labels, duration_ms); + } +} + #[cfg(test)] mod tests { use super::*; - use crate::common::stats_manager::{LabelSet, LabelType, MetricName, StatsManager}; - use crate::proto::api::instance::{ - GetPrometheusStatsRequest, GetPrometheusStatsResponse, GetStatsRequest, GetStatsResponse, - }; - use std::collections::BTreeMap; + + impl StatsManager { + pub(crate) fn cleanup_task_is_stopped(&self) -> bool { + self.cleanup_task.lock().unwrap().is_none() + } + + fn get_simple_counter(&self, name: MetricName) -> CounterHandle { + self.get_counter(name, LabelSet::new()) + } + + fn get_metrics_by_prefix(&self, prefix: &str) -> Vec { + self.get_all_metrics() + .into_iter() + .filter(|m| m.name.to_string().starts_with(prefix)) + .collect() + } + + pub(crate) fn get_metric( + &self, + name: MetricName, + labels: &LabelSet, + ) -> Option { + let key = MetricKey::new(name, labels.clone()); + + if let Some(metric_data) = self.counters.get(&key) { + let value = unsafe { metric_data.counter.get() }; + Some(MetricSnapshot { + name, + labels: labels.clone(), + value, + }) + } else { + None + } + } + + fn metric_count(&self) -> usize { + self.counters.len() + } + } + + #[test] + fn cleanup_task_can_start_after_sync_construction() { + let stats = StatsManager::new(); + assert!(stats.cleanup_task_is_stopped()); + + tokio::runtime::Builder::new_current_thread() + .enable_time() + .build() + .unwrap() + .block_on(async { + stats.start_cleanup_task(); + assert!(!stats.cleanup_task_is_stopped()); + stats.stop_cleanup_task().await; + }); + + assert!(stats.cleanup_task_is_stopped()); + } + + #[test] + fn cleanup_task_restarts_after_its_runtime_stops() { + let first_runtime = tokio::runtime::Builder::new_current_thread() + .enable_time() + .build() + .unwrap(); + let stats = first_runtime.block_on(async { + let stats = StatsManager::new(); + assert!(!stats.cleanup_task_is_stopped()); + stats + }); + drop(first_runtime); + assert!( + stats + .cleanup_task + .lock() + .unwrap() + .as_ref() + .unwrap() + .is_finished() + ); + + tokio::runtime::Builder::new_current_thread() + .enable_time() + .build() + .unwrap() + .block_on(async { + stats.start_cleanup_task(); + assert!( + !stats + .cleanup_task + .lock() + .unwrap() + .as_ref() + .unwrap() + .is_finished() + ); + stats.stop_cleanup_task().await; + }); + } #[tokio::test] async fn test_label_set() { @@ -969,69 +1269,6 @@ mod tests { assert_eq!(stats.metric_count(), 0); } - #[tokio::test] - async fn test_stats_rpc_data_structures() { - // Test GetStatsRequest - let request = GetStatsRequest { instance: None }; - assert_eq!(request, GetStatsRequest { instance: None }); - - // Test GetStatsResponse - let response = GetStatsResponse { metrics: vec![] }; - assert!(response.metrics.is_empty()); - - // Test GetPrometheusStatsRequest - let prometheus_request = GetPrometheusStatsRequest { instance: None }; - assert_eq!( - prometheus_request, - GetPrometheusStatsRequest { instance: None } - ); - - // Test GetPrometheusStatsResponse - let prometheus_response = GetPrometheusStatsResponse { - prometheus_text: "# Test metrics\n".to_string(), - }; - assert_eq!(prometheus_response.prometheus_text, "# Test metrics\n"); - } - - #[tokio::test] - async fn test_metric_snapshot_creation() { - let stats_manager = StatsManager::new(); - - // Create some test metrics - let counter1 = stats_manager.get_counter( - MetricName::PeerRpcClientTx, - LabelSet::new() - .with_label_type(LabelType::SrcPeerId(123)) - .with_label_type(LabelType::ServiceName("test_service".to_string())), - ); - counter1.add(100); - - let counter2 = stats_manager.get_counter( - MetricName::TrafficBytesTx, - LabelSet::new().with_label_type(LabelType::Protocol("tcp".to_string())), - ); - counter2.add(1024); - - // Get all metrics - let metrics = stats_manager.get_all_metrics(); - assert_eq!(metrics.len(), 2); - - // Verify the metrics can be converted to the format expected by RPC - for metric in metrics { - let mut labels = BTreeMap::new(); - for label in metric.labels.labels() { - labels.insert(label.key.clone(), label.value.clone()); - } - - // This simulates what the RPC service would do - let _metric_snapshot = crate::proto::api::instance::MetricSnapshot { - name: metric.name.to_string(), - value: metric.value, - labels, - }; - } - } - #[tokio::test] async fn test_prometheus_export_format() { let stats_manager = StatsManager::new(); diff --git a/easytier-core/src/foundation/task.rs b/easytier-core/src/foundation/task.rs new file mode 100644 index 00000000..eb950e73 --- /dev/null +++ b/easytier-core/src/foundation/task.rs @@ -0,0 +1,310 @@ +use std::{ + result::Result, + sync::{Arc, Mutex, atomic::Ordering}, + time::Duration, +}; + +use anyhow::Error; +use async_trait::async_trait; +use atomic_shim::AtomicU64; +use dashmap::DashMap; +use tokio::{ + select, + sync::Notify, + task::{JoinHandle, JoinSet}, +}; +use tokio_util::task::AbortOnDropHandle; + +pub(crate) async fn reap_joinset_background(tasks: Arc>>, origin: &'static str) +where + T: Send + 'static, +{ + let tasks = Arc::downgrade(&tasks); + loop { + crate::foundation::time::sleep(Duration::from_secs(1)).await; + let Some(tasks) = tasks.upgrade() else { + break; + }; + while tasks.lock().unwrap().try_join_next().is_some() {} + } + tracing::debug!(origin, "joinset task reaper exited"); +} + +pub struct ExternalTaskSignal { + version: AtomicU64, + notify: Notify, +} + +impl Default for ExternalTaskSignal { + fn default() -> Self { + Self::new() + } +} + +impl ExternalTaskSignal { + pub fn new() -> Self { + Self { + version: AtomicU64::new(0), + notify: Notify::new(), + } + } + + pub fn notify(&self) { + self.version.fetch_add(1, Ordering::Relaxed); + self.notify.notify_waiters(); + } + + pub fn version(&self) -> u64 { + self.version.load(Ordering::Relaxed) + } + + pub fn notified(&self) -> impl std::future::Future + '_ { + self.notify.notified() + } +} + +#[async_trait] +pub trait PeerTaskLauncher: Send + Sync + Clone + 'static { + type CollectPeerItem; + type TaskRet; + + async fn collect_peers_need_task(&self) -> Vec; + async fn launch_task( + &self, + item: Self::CollectPeerItem, + ) -> JoinHandle>; + + async fn all_task_done(&self) {} + + fn loop_interval_ms(&self) -> u64 { + 5000 + } +} + +type PeerTaskMap = DashMap< + ::CollectPeerItem, + AbortOnDropHandle::TaskRet, Error>>, +>; + +pub struct PeerTaskManager { + launcher: Launcher, + main_loop_task: Mutex>>, + peer_tasks: Arc>, + run_signal: Arc, + external_signal: Option>, +} + +impl PeerTaskManager +where + C: std::fmt::Debug + Send + Sync + Clone + core::hash::Hash + Eq + 'static, + T: Send + 'static, + L: PeerTaskLauncher + 'static, +{ + pub fn new_with_external_signal( + launcher: L, + external_signal: Option>, + ) -> Self { + Self { + launcher, + main_loop_task: Mutex::new(None), + peer_tasks: Arc::new(DashMap::new()), + run_signal: Arc::new(Notify::new()), + external_signal, + } + } + + pub fn start(&self) { + let mut task_slot = self.main_loop_task.lock().unwrap(); + if task_slot.as_ref().is_some_and(|task| !task.is_finished()) { + return; + } + let task = AbortOnDropHandle::new(tokio::spawn(Self::main_loop( + self.launcher.clone(), + self.run_signal.clone(), + self.external_signal.clone(), + self.peer_tasks.clone(), + ))); + task_slot.replace(task); + } + + pub async fn stop(&self) { + let task = self.main_loop_task.lock().unwrap().take(); + if let Some(task) = task { + task.abort(); + let _ = task.await; + } + let keys = self + .peer_tasks + .iter() + .map(|entry| entry.key().clone()) + .collect::>(); + for key in keys { + if let Some((_, task)) = self.peer_tasks.remove(&key) { + task.abort(); + let _ = task.await; + } + } + self.peer_tasks.shrink_to_fit(); + self.launcher.all_task_done().await; + } + + async fn main_loop( + launcher: L, + signal: Arc, + external_signal: Option>, + peer_task_map: Arc>>>, + ) { + let mut external_signal_version = external_signal.as_ref().map(|signal| signal.version()); + + loop { + let peers_to_connect = launcher.collect_peers_need_task().await; + + let mut to_remove = vec![]; + for item in peer_task_map.iter() { + if !peers_to_connect.contains(item.key()) || item.value().is_finished() { + to_remove.push(item.key().clone()); + } + } + + for key in to_remove { + if let Some((_, task)) = peer_task_map.remove(&key) { + task.abort(); + match task.await { + Ok(Ok(_)) => {} + Ok(Err(task_ret)) => { + tracing::error!( + target: "easytier_core::peers::peer_task", + ?task_ret, + "hole punching task failed" + ); + } + Err(e) => { + tracing::error!( + target: "easytier_core::peers::peer_task", + ?e, + "hole punching task aborted" + ); + } + } + } + peer_task_map.shrink_to_fit(); + } + + if !peers_to_connect.is_empty() { + for item in peers_to_connect { + if peer_task_map.contains_key(&item) { + continue; + } + + tracing::debug!( + target: "easytier_core::peers::peer_task", + ?item, + "launch hole punching task" + ); + peer_task_map.insert( + item.clone(), + AbortOnDropHandle::new(launcher.launch_task(item).await), + ); + } + } else if peer_task_map.is_empty() { + launcher.all_task_done().await; + } + + if let Some(external_signal) = external_signal.as_ref() { + let notified = external_signal.notified(); + tokio::pin!(notified); + let cur_version = external_signal.version(); + if external_signal_version != Some(cur_version) { + external_signal_version = Some(cur_version); + continue; + } + + select! { + _ = crate::foundation::time::sleep(std::time::Duration::from_millis( + launcher.loop_interval_ms(), + )) => {}, + _ = signal.notified() => {}, + _ = &mut notified => { + external_signal_version = Some(external_signal.version()); + } + } + } else { + select! { + _ = crate::foundation::time::sleep(std::time::Duration::from_millis( + launcher.loop_interval_ms(), + )) => {}, + _ = signal.notified() => {} + } + } + } + } +} + +#[cfg(test)] +mod tests { + use std::sync::{ + Arc, + atomic::{AtomicUsize, Ordering}, + }; + + use super::*; + + #[derive(Clone)] + struct TestLauncher { + active_tasks: Arc, + } + + struct ActiveTaskGuard(Arc); + + impl Drop for ActiveTaskGuard { + fn drop(&mut self) { + self.0.fetch_sub(1, Ordering::SeqCst); + } + } + + #[async_trait] + impl PeerTaskLauncher for TestLauncher { + type CollectPeerItem = u8; + type TaskRet = (); + + async fn collect_peers_need_task(&self) -> Vec { + vec![1] + } + + async fn launch_task(&self, _item: u8) -> JoinHandle> { + let active_tasks = self.active_tasks.clone(); + tokio::spawn(async move { + active_tasks.fetch_add(1, Ordering::SeqCst); + let _guard = ActiveTaskGuard(active_tasks); + std::future::pending::<()>().await; + Ok(()) + }) + } + } + + #[tokio::test] + async fn peer_task_manager_is_cold_and_joins_children_on_stop() { + let active_tasks = Arc::new(AtomicUsize::new(0)); + let manager = PeerTaskManager::new_with_external_signal( + TestLauncher { + active_tasks: active_tasks.clone(), + }, + None, + ); + + tokio::task::yield_now().await; + assert_eq!(active_tasks.load(Ordering::SeqCst), 0); + + manager.start(); + crate::foundation::time::timeout(std::time::Duration::from_secs(1), async { + while active_tasks.load(Ordering::SeqCst) == 0 { + tokio::task::yield_now().await; + } + }) + .await + .unwrap(); + + manager.stop().await; + assert_eq!(active_tasks.load(Ordering::SeqCst), 0); + } +} diff --git a/easytier-core/src/foundation/time.rs b/easytier-core/src/foundation/time.rs new file mode 100644 index 00000000..95a8646f --- /dev/null +++ b/easytier-core/src/foundation/time.rs @@ -0,0 +1,14 @@ +//! Tokio time facade. +//! +//! Native builds use Tokio directly. WASI builds use the deadline-tracking +//! implementation in [`crate::wasi::time`] so an external runtime can drive +//! the guest without polling. + +#[cfg(not(any(test, target_os = "wasi")))] +pub use tokio::time::{Duration, Instant, Interval, error, interval, sleep, timeout}; + +#[cfg(any(test, target_os = "wasi"))] +pub use crate::wasi::time::{Duration, Instant, Interval, error, interval, sleep, timeout}; + +#[cfg(target_os = "wasi")] +pub(crate) use crate::wasi::time::{clear_domain, enter_domain, next_deadline_millis}; diff --git a/easytier/src/common/token_bucket.rs b/easytier-core/src/foundation/token_bucket.rs similarity index 79% rename from easytier/src/common/token_bucket.rs rename to easytier-core/src/foundation/token_bucket.rs index ab1c9d4c..785be405 100644 --- a/easytier/src/common/token_bucket.rs +++ b/easytier-core/src/foundation/token_bucket.rs @@ -1,13 +1,31 @@ use atomic_shim::AtomicU64; use dashmap::DashMap; -use std::sync::atomic::Ordering; +use std::sync::atomic::{AtomicBool, Ordering}; use std::sync::{Arc, Mutex}; use std::time::{Duration, Instant}; use tokio::sync::Notify; -use tokio::time; use tokio_util::task::AbortOnDropHandle; -use crate::proto::common::LimiterConfig; +use crate::foundation::time; +use easytier_proto::common::LimiterConfig; + +#[async_trait::async_trait] +pub(crate) trait ByteLimiter: Send + Sync { + async fn consume(&self, bytes: u64); + + fn try_consume(&self, bytes: u64) -> bool; +} + +#[async_trait::async_trait] +impl ByteLimiter for () { + async fn consume(&self, _bytes: u64) {} + + fn try_consume(&self, _bytes: u64) -> bool { + true + } +} + +pub(crate) type ArcByteLimiter = Arc; /// Token Bucket rate limiter using atomic operations pub struct TokenBucket { @@ -18,6 +36,7 @@ pub struct TokenBucket { start_time: Instant, // Bucket creation time refill_notifier: Arc, + stopped: AtomicBool, } #[derive(Clone, Copy)] @@ -68,6 +87,7 @@ impl TokenBucket { refill_task: Mutex::new(None), start_time: std::time::Instant::now(), refill_notifier: Arc::new(Notify::new()), + stopped: AtomicBool::new(false), }); // Start background refill task @@ -137,6 +157,9 @@ impl TokenBucket { /// # Returns /// `true` if tokens were consumed, `false` if insufficient tokens pub fn try_consume(&self, tokens: u64) -> bool { + if self.stopped.load(Ordering::Acquire) { + return true; + } // Fast path for oversized packets if tokens > self.config.capacity { return false; @@ -163,16 +186,40 @@ impl TokenBucket { /// Consume tokens, blocking if not available pub async fn consume(&self, tokens: u64) { - while !self.try_consume(tokens) { - self.refill_notifier.notified().await; + loop { + let notified = self.refill_notifier.notified(); + if self.try_consume(tokens) { + return; + } + notified.await; } } + + async fn stop(&self) { + self.stopped.store(true, Ordering::Release); + self.refill_notifier.notify_waiters(); + let task = self.refill_task.lock().unwrap().take(); + if let Some(task) = task { + task.abort(); + let _ = task.await; + } + } +} + +#[async_trait::async_trait] +impl ByteLimiter for TokenBucket { + async fn consume(&self, bytes: u64) { + TokenBucket::consume(self, bytes).await; + } + + fn try_consume(&self, bytes: u64) -> bool { + TokenBucket::try_consume(self, bytes) + } } pub struct TokenBucketManager { buckets: Arc>>, - - retain_task: AbortOnDropHandle<()>, + retain_task: Mutex>>, } impl Default for TokenBucketManager { @@ -194,7 +241,7 @@ impl TokenBucketManager { buckets_clone.retain(|_, bucket| Arc::::strong_count(bucket) > 1); buckets_clone.shrink_to_fit(); // Sleep for a while before next retention check - tokio::time::sleep(Duration::from_secs(5)).await; + time::sleep(Duration::from_secs(5)).await; tracing::info!( "Retained buckets: {} ({} dropped)", buckets_clone.len(), @@ -205,7 +252,7 @@ impl TokenBucketManager { Self { buckets, - retain_task: AbortOnDropHandle::new(retain_task), + retain_task: Mutex::new(Some(AbortOnDropHandle::new(retain_task))), } } @@ -216,20 +263,27 @@ impl TokenBucketManager { .or_insert_with(|| TokenBucket::new_from_cfg(cfg)) .clone() } + + pub async fn stop(&self) { + let retain_task = self.retain_task.lock().unwrap().take(); + if let Some(retain_task) = retain_task { + retain_task.abort(); + let _ = retain_task.await; + } + let buckets = self + .buckets + .iter() + .map(|entry| entry.value().clone()) + .collect::>(); + for bucket in buckets { + bucket.stop().await; + } + self.buckets.clear(); + } } #[cfg(test)] mod tests { - use crate::{ - connector::udp_hole_punch::tests::create_mock_peer_manager_with_mock_stun, - peers::{ - foreign_network_manager::tests::create_mock_peer_manager_for_foreign_network, - tests::connect_peer_manager, - }, - proto::common::NatType, - tunnel::common::tests::wait_for_condition, - }; - use super::*; use tokio::time::{Duration, sleep}; @@ -258,6 +312,26 @@ mod tests { assert!(bucket.try_consume(500)); } + #[tokio::test] + async fn stop_releases_waiting_consumers() { + let bucket = TokenBucket::new(1, 1, Duration::from_secs(60)); + assert!(bucket.try_consume(1)); + let waiting = tokio::spawn({ + let bucket = bucket.clone(); + async move { bucket.consume(1).await } + }); + tokio::task::yield_now().await; + assert!(!waiting.is_finished()); + + bucket.stop().await; + + tokio::time::timeout(Duration::from_secs(1), waiting) + .await + .expect("stopped limiter should release consumers") + .unwrap(); + assert!(bucket.try_consume(u64::MAX)); + } + /// Test background refill functionality #[tokio::test] async fn test_refill() { @@ -349,56 +423,4 @@ mod tests { tokens ); } - - #[tokio::test] - async fn test_token_bucket_free() { - let pm_center1 = create_mock_peer_manager_with_mock_stun(NatType::Unknown).await; - - for i in 0..10 { - let pma_net1 = create_mock_peer_manager_for_foreign_network(&format!("net{}", i)).await; - - connect_peer_manager(pma_net1.clone(), pm_center1.clone()).await; - wait_for_condition( - || async { pma_net1.list_routes().await.len() == 1 }, - Duration::from_secs(5), - ) - .await; - println!("net{}", i); - println!( - "buckets: {}", - pm_center1 - .get_global_ctx() - .token_bucket_manager() - .buckets - .len() - ); - - drop(pma_net1); - wait_for_condition( - || async { - pm_center1 - .get_foreign_network_manager() - .list_foreign_networks() - .await - .foreign_networks - .is_empty() - }, - Duration::from_secs(5), - ) - .await; - } - - // wait token bucket empty - wait_for_condition( - || async { - pm_center1 - .get_global_ctx() - .token_bucket_manager() - .buckets - .is_empty() - }, - Duration::from_secs(10), - ) - .await; - } } diff --git a/easytier-core/src/gateway/dataplane/deadline.rs b/easytier-core/src/gateway/dataplane/deadline.rs new file mode 100644 index 00000000..a1128b9e --- /dev/null +++ b/easytier-core/src/gateway/dataplane/deadline.rs @@ -0,0 +1,52 @@ +//! One absolute deadline shared by every stage of a data-plane operation. + +use std::{future::Future, time::Duration}; + +use quanta::Instant; + +use crate::foundation::time; + +use super::{DataPlaneError, DataPlaneResult}; + +#[derive(Clone, Copy, Debug)] +pub(super) struct DataPlaneDeadline(Option); + +impl DataPlaneDeadline { + pub(super) fn from_timeout(timeout: Duration) -> Self { + Self(Instant::now().checked_add(timeout)) + } + + pub(super) fn from_optional_timeout(timeout: Option) -> Self { + match timeout { + Some(timeout) => Self::from_timeout(timeout), + None => Self(None), + } + } + + pub(super) fn remaining(self) -> DataPlaneResult> { + let Some(deadline) = self.0 else { + return Ok(None); + }; + let now = Instant::now(); + if now >= deadline { + return Err(DataPlaneError::deadline_exceeded()); + } + Ok(Some(deadline - now)) + } + + pub(super) async fn run( + self, + future: impl Future>, + ) -> DataPlaneResult + where + E: Into, + { + match self.remaining()? { + Some(remaining) => time::timeout(remaining, future) + .await + .map_err(|_| DataPlaneError::deadline_exceeded())? + .map_err(Into::into), + None => future.await.map_err(Into::into), + } + } +} diff --git a/easytier-core/src/gateway/dataplane/error.rs b/easytier-core/src/gateway/dataplane/error.rs new file mode 100644 index 00000000..d74d9eeb --- /dev/null +++ b/easytier-core/src/gateway/dataplane/error.rs @@ -0,0 +1,117 @@ +//! Stable errors returned by the public data-plane interface. + +use std::{fmt, io}; + +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +#[non_exhaustive] +#[repr(u16)] +pub enum DataPlaneErrorKind { + Cancelled = 1, + DeadlineExceeded = 2, + InstanceStopped = 3, + HandleClosed = 4, + NoOverlayRoute = 5, + PathNotReady = 6, + AddressFamilyUnsupported = 7, + AddressInUse = 8, + ConnectionRefused = 9, + NetworkChanged = 10, + ResourceLimit = 11, + Io = 12, + BufferTooSmall = 13, +} + +#[derive(Clone, Debug)] +pub struct DataPlaneError { + kind: DataPlaneErrorKind, + message: String, +} + +impl DataPlaneError { + pub(crate) fn new(kind: DataPlaneErrorKind, message: impl Into) -> Self { + Self { + kind, + message: message.into(), + } + } + + pub(crate) fn deadline_exceeded() -> Self { + Self::new( + DataPlaneErrorKind::DeadlineExceeded, + "data-plane deadline exceeded", + ) + } + + pub fn kind(&self) -> DataPlaneErrorKind { + self.kind + } + + pub fn message(&self) -> &str { + &self.message + } + + pub(crate) fn into_io_error(self) -> io::Error { + let kind = match self.kind { + DataPlaneErrorKind::Cancelled => io::ErrorKind::Interrupted, + DataPlaneErrorKind::DeadlineExceeded => io::ErrorKind::TimedOut, + DataPlaneErrorKind::AddressFamilyUnsupported => io::ErrorKind::Unsupported, + DataPlaneErrorKind::AddressInUse => io::ErrorKind::AddrInUse, + DataPlaneErrorKind::ConnectionRefused => io::ErrorKind::ConnectionRefused, + DataPlaneErrorKind::ResourceLimit => io::ErrorKind::OutOfMemory, + DataPlaneErrorKind::BufferTooSmall => io::ErrorKind::InvalidInput, + DataPlaneErrorKind::InstanceStopped + | DataPlaneErrorKind::HandleClosed + | DataPlaneErrorKind::NetworkChanged => io::ErrorKind::BrokenPipe, + DataPlaneErrorKind::NoOverlayRoute | DataPlaneErrorKind::PathNotReady => { + io::ErrorKind::NotConnected + } + DataPlaneErrorKind::Io => io::ErrorKind::Other, + }; + io::Error::new(kind, self) + } +} + +impl fmt::Display for DataPlaneError { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.write_str(&self.message) + } +} + +impl std::error::Error for DataPlaneError {} + +impl From for DataPlaneError { + fn from(error: io::Error) -> Self { + if let Some(error) = error + .get_ref() + .and_then(|error| error.downcast_ref::()) + { + return Self::new(error.kind, error.message.clone()); + } + let kind = match error.kind() { + io::ErrorKind::TimedOut => DataPlaneErrorKind::DeadlineExceeded, + io::ErrorKind::AddrInUse => DataPlaneErrorKind::AddressInUse, + io::ErrorKind::ConnectionRefused => DataPlaneErrorKind::ConnectionRefused, + io::ErrorKind::Interrupted => DataPlaneErrorKind::Cancelled, + io::ErrorKind::Unsupported => DataPlaneErrorKind::AddressFamilyUnsupported, + io::ErrorKind::NotConnected => DataPlaneErrorKind::PathNotReady, + io::ErrorKind::BrokenPipe => DataPlaneErrorKind::HandleClosed, + io::ErrorKind::OutOfMemory => DataPlaneErrorKind::ResourceLimit, + _ => DataPlaneErrorKind::Io, + }; + Self::new(kind, error.to_string()) + } +} + +impl From for DataPlaneError { + fn from(error: anyhow::Error) -> Self { + if let Some(error) = error.downcast_ref::() { + return Self::new(error.kind, error.message.clone()); + } + if let Some(error) = error.downcast_ref::() { + return Self::from(io::Error::new(error.kind(), error.to_string())); + } + Self::new(DataPlaneErrorKind::Io, format!("{error:#}")) + } +} + +pub type DataPlaneResult = Result; diff --git a/easytier-core/src/gateway/dataplane/flow.rs b/easytier-core/src/gateway/dataplane/flow.rs new file mode 100644 index 00000000..b1fd1c32 --- /dev/null +++ b/easytier-core/src/gateway/dataplane/flow.rs @@ -0,0 +1,415 @@ +//! Shared data-plane flow registration and ownership. + +use std::{ + net::{IpAddr, SocketAddr}, + sync::{ + Arc, + atomic::{AtomicUsize, Ordering}, + }, +}; + +use atomic_shim::AtomicU64; +use dashmap::{DashMap, mapref::entry::Entry}; + +#[derive(Clone, Copy, Debug, Eq, Hash, PartialEq)] +#[repr(u8)] +pub(crate) enum FlowKind { + Udp = 1, + Tcp = 2, + TcpListen = 3, +} + +#[derive(Clone, Debug, Eq, Hash, PartialEq)] +pub(crate) struct FlowKey { + pub src: SocketAddr, + pub dst: SocketAddr, + pub kind: FlowKind, +} + +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +pub(crate) struct FlowCountChange { + pub previous: usize, + pub current: usize, +} + +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +pub(crate) struct FlowInsert { + pub replaced: bool, + pub count: FlowCountChange, +} + +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +pub(crate) struct FlowRemoval { + pub removed: bool, + pub count: FlowCountChange, +} + +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +pub(crate) struct FlowRetain { + pub removed: usize, + pub count: FlowCountChange, +} + +pub(crate) struct FlowTable { + entries: DashMap>, + count: AtomicUsize, + next_registration: AtomicU64, +} + +struct RegisteredFlow { + registration: u64, + value: V, +} + +pub(crate) struct FlowLease { + table: Arc>, + entry: FlowKey, + registration: u64, + active: bool, +} + +impl FlowLease { + pub fn register(table: Arc>, entry: FlowKey, value: V) -> (Self, FlowInsert) { + let (registration, insert) = table.insert_registered(entry.clone(), value); + ( + Self { + table, + entry, + registration, + active: true, + }, + insert, + ) + } + + pub fn try_register(table: Arc>, entry: FlowKey, value: V) -> Option { + let registration = table.try_insert_registered(entry.clone(), value)?; + Some(Self { + table, + entry, + registration, + active: true, + }) + } +} + +impl Drop for FlowLease { + fn drop(&mut self) { + if self.active { + self.table + .remove_registration(&self.entry, self.registration); + } + } +} + +impl Default for FlowTable { + fn default() -> Self { + Self { + entries: DashMap::new(), + count: AtomicUsize::new(0), + next_registration: AtomicU64::new(1), + } + } +} + +impl FlowTable { + pub fn count(&self) -> usize { + self.count.load(Ordering::Relaxed) + } + + pub fn len(&self) -> usize { + self.entries.len() + } + + pub fn is_empty(&self) -> bool { + self.entries.is_empty() + } + + pub fn contains_key(&self, entry: &FlowKey) -> bool { + self.entries.contains_key(entry) + } + + pub fn contains_destination_ip(&self, destination: IpAddr) -> bool { + self.entries + .iter() + .any(|entry| entry.key().dst.ip() == destination) + } + + #[cfg(test)] + pub fn with_entry(&self, entry: &FlowKey, f: impl FnOnce(&V) -> R) -> Option { + self.entries.get(entry).map(|value| f(&value.value().value)) + } + + #[cfg(test)] + pub fn insert(&self, entry: FlowKey, value: V) -> FlowInsert { + self.insert_registered(entry, value).1 + } + + fn insert_registered(&self, entry: FlowKey, value: V) -> (u64, FlowInsert) { + let registration = self.next_registration(); + match self.entries.entry(entry) { + Entry::Occupied(mut occupied) => { + occupied.insert(RegisteredFlow { + registration, + value, + }); + let count = self.count(); + ( + registration, + FlowInsert { + replaced: true, + count: FlowCountChange { + previous: count, + current: count, + }, + }, + ) + } + Entry::Vacant(vacant) => { + // Reserve the count while holding the shard lock so retain cannot + // observe the entry before its count is accounted for. + let count = self.increment_count(); + vacant.insert(RegisteredFlow { + registration, + value, + }); + ( + registration, + FlowInsert { + replaced: false, + count, + }, + ) + } + } + } + + #[cfg(test)] + pub fn try_insert(&self, entry: FlowKey, value: V) -> bool { + self.try_insert_registered(entry, value).is_some() + } + + fn try_insert_registered(&self, entry: FlowKey, value: V) -> Option { + match self.entries.entry(entry) { + Entry::Occupied(_) => None, + Entry::Vacant(vacant) => { + let registration = self.next_registration(); + self.increment_count(); + vacant.insert(RegisteredFlow { + registration, + value, + }); + Some(registration) + } + } + } + + #[cfg(test)] + pub fn remove(&self, entry: &FlowKey) -> FlowRemoval { + let removed = self.entries.remove(entry).is_some(); + let count = if removed { + self.decrement_count_by(1) + } else { + let count = self.count(); + FlowCountChange { + previous: count, + current: count, + } + }; + FlowRemoval { removed, count } + } + + pub fn retain(&self, mut f: impl FnMut(&FlowKey, &mut V) -> bool) -> FlowRetain { + let mut removed = 0; + self.entries.retain(|entry, value| { + let keep = f(entry, &mut value.value); + if !keep { + removed += 1; + } + keep + }); + FlowRetain { + removed, + count: self.decrement_count_by(removed), + } + } + + pub fn clear(&self) -> FlowRetain { + self.retain(|_, _| false) + } + + fn increment_count(&self) -> FlowCountChange { + let previous = self + .count + .fetch_update(Ordering::Relaxed, Ordering::Relaxed, |count| { + count.checked_add(1) + }) + .unwrap_or_else(|count| count); + FlowCountChange { + previous, + current: previous.saturating_add(1), + } + } + + fn next_registration(&self) -> u64 { + self.next_registration.fetch_add(1, Ordering::Relaxed) + } + + fn remove_registration(&self, entry: &FlowKey, registration: u64) -> FlowRemoval { + let removed = self + .entries + .remove_if(entry, |_, flow| flow.registration == registration) + .is_some(); + let count = if removed { + self.decrement_count_by(1) + } else { + let count = self.count(); + FlowCountChange { + previous: count, + current: count, + } + }; + FlowRemoval { removed, count } + } + + fn decrement_count_by(&self, delta: usize) -> FlowCountChange { + if delta == 0 { + let count = self.count(); + return FlowCountChange { + previous: count, + current: count, + }; + } + + let previous = self + .count + .fetch_update(Ordering::Relaxed, Ordering::Relaxed, |count| { + Some(count.saturating_sub(delta)) + }) + .unwrap_or_else(|count| count); + FlowCountChange { + previous, + current: previous.saturating_sub(delta), + } + } +} + +#[cfg(test)] +mod tests { + use super::{FlowKey, FlowKind, FlowLease, FlowTable}; + use std::{ + net::{IpAddr, Ipv4Addr, SocketAddr}, + sync::Arc, + }; + + impl FlowLease { + fn remove(mut self) -> super::FlowRemoval { + self.active = false; + self.table + .remove_registration(&self.entry, self.registration) + } + } + + #[test] + fn entry_kind_values_preserve_native_table_identity() { + assert_eq!(FlowKind::Udp as u8, 1); + assert_eq!(FlowKind::Tcp as u8, 2); + assert_eq!(FlowKind::TcpListen as u8, 3); + } + + fn table_entry(port: u16) -> FlowKey { + FlowKey { + src: SocketAddr::new(IpAddr::V4(Ipv4Addr::new(10, 42, 0, 2)), port), + dst: SocketAddr::new(IpAddr::V4(Ipv4Addr::new(10, 42, 0, 1)), 22), + kind: FlowKind::Tcp, + } + } + + #[test] + fn flow_table_tracks_insert_replace_and_remove() { + let table = FlowTable::default(); + let entry = table_entry(40000); + + let inserted = table.insert(entry.clone(), "first"); + assert!(!inserted.replaced); + assert_eq!(inserted.count.previous, 0); + assert_eq!(inserted.count.current, 1); + assert_eq!(table.with_entry(&entry, |value| *value), Some("first")); + + let replaced = table.insert(entry.clone(), "second"); + assert!(replaced.replaced); + assert_eq!(replaced.count.previous, 1); + assert_eq!(replaced.count.current, 1); + assert_eq!(table.with_entry(&entry, |value| *value), Some("second")); + + let removed = table.remove(&entry); + assert!(removed.removed); + assert_eq!(removed.count.previous, 1); + assert_eq!(removed.count.current, 0); + + let missing = table.remove(&entry); + assert!(!missing.removed); + assert_eq!(missing.count.previous, 0); + assert_eq!(missing.count.current, 0); + } + + #[test] + fn flow_table_try_insert_and_retain_keep_count_consistent() { + let table = FlowTable::default(); + let first = table_entry(40000); + let second = table_entry(40001); + + assert!(table.try_insert(first.clone(), 1)); + assert!(!table.try_insert(first.clone(), 2)); + assert!(table.try_insert(second.clone(), 3)); + assert_eq!(table.count(), 2); + assert!(table.contains_destination_ip(first.dst.ip())); + + let retained = table.retain(|entry, _| entry == &second); + assert_eq!(retained.removed, 1); + assert_eq!(retained.count.previous, 2); + assert_eq!(retained.count.current, 1); + assert!(!table.contains_key(&first)); + assert!(table.contains_key(&second)); + + let cleared = table.clear(); + assert_eq!(cleared.removed, 1); + assert_eq!(cleared.count.current, 0); + assert!(table.is_empty()); + } + + #[test] + fn entry_guard_owns_registration_lifetime() { + let table = Arc::new(FlowTable::default()); + let entry = table_entry(40000); + + let (guard, insert) = FlowLease::register(table.clone(), entry.clone(), "first"); + assert!(!insert.replaced); + assert!(table.contains_key(&entry)); + assert!(FlowLease::try_register(table.clone(), entry.clone(), "second").is_none()); + assert_eq!(table.with_entry(&entry, |value| *value), Some("first")); + + drop(guard); + assert!(!table.contains_key(&entry)); + + let guard = FlowLease::try_register(table.clone(), entry.clone(), "third").unwrap(); + let removal = guard.remove(); + assert!(removal.removed); + assert_eq!(table.count(), 0); + } + + #[test] + fn replaced_lease_cannot_remove_new_registration() { + let table = Arc::new(FlowTable::default()); + let entry = table_entry(40000); + let (old, _) = FlowLease::register(table.clone(), entry.clone(), "old"); + let (new, replaced) = FlowLease::register(table.clone(), entry.clone(), "new"); + assert!(replaced.replaced); + + drop(old); + assert_eq!(table.with_entry(&entry, |value| *value), Some("new")); + + drop(new); + assert!(!table.contains_key(&entry)); + } +} diff --git a/easytier-core/src/gateway/dataplane/mod.rs b/easytier-core/src/gateway/dataplane/mod.rs new file mode 100644 index 00000000..2d882edf --- /dev/null +++ b/easytier-core/src/gateway/dataplane/mod.rs @@ -0,0 +1,976 @@ +//! Data-plane access built on top of the core gateway smoltcp stack. +//! +//! This module exposes TCP streams and UDP sockets (mainly for FFI callers that +//! send traffic through EasyTier without creating OS-level proxy listeners). +//! +//! Typical usage: +//! +//! ```ignore +//! let instance = CoreInstance::new(...); +//! instance.start().await?; +//! +//! let socket = instance.data_plane_udp_bind(local_port, timeout).await?; +//! socket.send_to(buf, peer_addr).await?; +//! ``` + +use std::{ + any::Any, + net::{IpAddr, Ipv4Addr, SocketAddr}, + sync::{ + Arc, Weak, + atomic::{AtomicBool, Ordering}, + }, + time::Duration, +}; + +use pnet_packet::{Packet, ip::IpNextHeaderProtocols, ipv4::Ipv4Packet, tcp::TcpPacket}; +use tokio::{ + select, + sync::{Mutex, mpsc}, + task::JoinSet, +}; + +use crate::{ + config::runtime::CoreRuntimeConfigStore, + foundation::task::reap_joinset_background, + gateway::{ + proxy::{ + traits::TcpProxyStream, + wrapped_transport::{WrappedTransportKind, WrappedTransportProxyModule}, + }, + smoltcp::{Net, UdpSocket}, + }, + packet::{PacketType, ZCPacket}, + peers::{ + PeerPacketFilter, + peer_manager::{PeerManagerCore, PipelineRegistrationGuard}, + }, + socket::{ + SocketContext, + tcp::{ + TcpBindOptions, TcpConnectOptions, TcpListenOptions, TcpSocketPurpose, + VirtualTcpListener, VirtualTcpListenerFactory, VirtualTcpSocket, + VirtualTcpSocketFactory, + }, + udp::{UdpBindOptions, VirtualUdpSocket, VirtualUdpSocketFactory}, + }, +}; + +mod deadline; +mod error; +mod flow; +mod operation; +mod packet; +mod resource; +mod route; +mod session; +mod stack; +mod tcp; +#[cfg(test)] +mod tests; +mod udp; + +use self::{ + deadline::DataPlaneDeadline, + error::DataPlaneResult, + flow::{FlowKey, FlowKind, FlowLease, FlowTable}, + packet::PeerPacketRoute, + resource::{DataPlaneConsumers, DataPlaneIoGuard, DataPlaneLease}, + route::{ + DataPlaneRoutePolicy, DataPlaneTcpRoute, DataPlaneTcpRouteInput, + DataPlaneTransportPreference, + }, + stack::SmoltcpPlane, +}; + +pub(crate) use self::resource::DataPlaneConsumerLease; +use self::tcp::DataPlaneTcpStreamRoute; +pub use self::{ + error::{DataPlaneError, DataPlaneErrorKind}, + operation::{ + DataPlaneCompletionDescriptor, DataPlaneCompletionStatus, DataPlaneOperationId, + DataPlaneOperationKind, DataPlaneOperationOutcome, DataPlaneOperationResult, + DataPlaneResourceId, + }, + session::{DataPlaneSession, DataPlaneSessionLimits}, + tcp::{DataPlaneTcpListener, DataPlaneTcpStream}, + udp::DataPlaneUdpSocket, +}; + +#[derive(Clone, Copy, Debug)] +pub(crate) struct DataPlaneTcpConnectOptions { + policy: DataPlaneRoutePolicy, + transport: DataPlaneTransportPreference, + purpose: TcpSocketPurpose, + source_hint: Option, + deadline: DataPlaneDeadline, +} + +impl DataPlaneTcpConnectOptions { + fn public(timeout: Duration) -> Self { + Self::public_with_timeout(Some(timeout)) + } + + fn public_with_timeout(timeout: Option) -> Self { + Self::public_with_deadline(DataPlaneDeadline::from_optional_timeout(timeout)) + } + + fn public_with_deadline(deadline: DataPlaneDeadline) -> Self { + Self { + policy: DataPlaneRoutePolicy::OverlayOnly, + transport: DataPlaneTransportPreference::SmoltcpOnly, + purpose: TcpSocketPurpose::DataPlane, + source_hint: None, + deadline, + } + } + + pub(crate) fn gateway( + timeout: Duration, + purpose: TcpSocketPurpose, + source_hint: SocketAddr, + ) -> Self { + Self { + policy: DataPlaneRoutePolicy::OverlayOrDirect, + transport: DataPlaneTransportPreference::PreferKcp, + purpose, + source_hint: Some(source_hint), + deadline: DataPlaneDeadline::from_timeout(timeout), + } + } +} + +pub(super) struct DataPlaneUdpIo(UdpSocket); + +impl DataPlaneUdpIo { + pub async fn send_to(&self, buf: &[u8], addr: SocketAddr) -> Result { + self.0.send_to(buf, addr).await + } + + pub async fn recv_from(&self, buf: &mut [u8]) -> Result<(usize, SocketAddr), std::io::Error> { + self.0.recv_from(buf).await + } + + pub async fn recv_from_limited( + &self, + max_len: usize, + ) -> Result<(Vec, SocketAddr, bool), std::io::Error> { + self.0.recv_from_limited(max_len).await + } +} + +pub(super) enum FlowData { + Tcp { + _reservation: Arc, + }, + // a data-plane routing entry that owns no resource. the entry_type in the + // key distinguishes a listen route from an actively outbound route. + DataPlaneRoute, + Udp, +} + +const UDP_ENTRY: FlowKind = FlowKind::Udp; +const TCP_ENTRY: FlowKind = FlowKind::Tcp; +const TCP_LISTEN_ENTRY: FlowKind = FlowKind::TcpListen; + +type FlowSet = Arc>; +type DataPlaneTcpIo = Box; + +pub(crate) struct DataPlaneRuntime +where + H: VirtualTcpSocketFactory + VirtualTcpListenerFactory + VirtualUdpSocketFactory, +{ + operation: Mutex<()>, + pub(super) runtime_started: AtomicBool, + runtime_guard: DataPlaneIoGuard, + runtime_config: CoreRuntimeConfigStore, + pub(super) peer_manager: Weak, + pub(super) transport_proxy: Option>, + pub(super) host: Arc, + pub(super) socket_context: SocketContext, + + pub(super) runtime_tasks: Arc>>, + packet_sender: mpsc::Sender, + packet_recv: Arc>>, + + net: Arc>>, + pub(super) entries: FlowSet, + + data_plane_consumers: Arc, + // Tracks whether the smoltcp `net` is ready for data-plane callers. + data_plane_net_ready: tokio::sync::watch::Sender, + pipeline_guard: Mutex>, +} + +#[async_trait::async_trait] +impl PeerPacketFilter for DataPlaneRuntime +where + H: VirtualTcpSocketFactory + VirtualTcpListenerFactory + VirtualUdpSocketFactory, +{ + async fn try_process_packet_from_peer(&self, packet: ZCPacket) -> Option { + let entry_count = self.entries.count(); + if entry_count == 0 && self.entries.is_empty() { + if tracing::enabled!(tracing::Level::TRACE) + && let Some(hdr) = packet.peer_manager_header() + && matches!( + hdr.packet_type, + x if x == PacketType::Data as u8 + || x == PacketType::DataWithKcpSrcModified as u8 + || x == PacketType::DataWithQuicSrcModified as u8 + ) + { + if let Some(ipv4) = Ipv4Packet::new(packet.payload()) { + let (tcp_src_port, tcp_dst_port, tcp_flags) = + if ipv4.get_next_level_protocol() == IpNextHeaderProtocols::Tcp { + TcpPacket::new(ipv4.payload()) + .map(|tcp| { + ( + Some(tcp.get_source()), + Some(tcp.get_destination()), + Some(tcp.get_flags()), + ) + }) + .unwrap_or((None, None, None)) + } else { + (None, None, None) + }; + tracing::trace!( + packet_type = hdr.packet_type, + from_peer_id = hdr.from_peer_id.get(), + to_peer_id = hdr.to_peer_id.get(), + ipv4_src = %ipv4.get_source(), + ipv4_dst = %ipv4.get_destination(), + next_protocol = ?ipv4.get_next_level_protocol(), + ?tcp_src_port, + ?tcp_dst_port, + ?tcp_flags, + entry_count, + "data plane fast gate passed packet from peer" + ); + } else { + tracing::trace!( + packet_type = hdr.packet_type, + from_peer_id = hdr.from_peer_id.get(), + to_peer_id = hdr.to_peer_id.get(), + entry_count, + "data plane fast gate passed non-ipv4 packet from peer" + ); + } + } + return Some(packet); + } + let route = self.entries.route_peer_packet(&packet, true); + let (entry_key, tcp_flags) = match route { + PeerPacketRoute::Pass => return Some(packet), + PeerPacketRoute::Unmatched { entry, tcp_flags } => { + tracing::trace!( + entry_key = ?entry, + ?tcp_flags, + ipv4_src = %entry.dst.ip(), + ipv4_dst = %entry.src.ip(), + entry_count = self.entries.count(), + "data plane has no flow for packet from peer" + ); + return Some(packet); + } + PeerPacketRoute::Deliver { entry, tcp_flags } => (entry, tcp_flags), + PeerPacketRoute::FragmentedUdp { source, mirror } => { + let source: IpAddr = source.into(); + tracing::trace!( + is_in_entries = mirror, + "ipv4 src = {:?}, check need send both smoltcp and kernel tun", + source + ); + if mirror { + // if the packet is fragmented, no matther what the payload is, need send it to both smoltcp and kernel tun. because + // we cannot determine the udp port of the packet. + match self.packet_sender.try_send(packet.clone()) { + Ok(()) => tracing::trace!( + ?source, + entry_count = self.entries.count(), + "data plane delivered fragmented packet from peer to smoltcp" + ), + Err(err) => tracing::trace!( + ?source, + ?err, + entry_count = self.entries.count(), + "data plane failed to deliver fragmented packet from peer to smoltcp" + ), + } + } + return Some(packet); + } + }; + + tracing::trace!( + ?entry_key, + ?tcp_flags, + ipv4_src = %entry_key.dst.ip(), + ipv4_dst = %entry_key.src.ip(), + entry_count = self.entries.count(), + "data plane found entry for packet from peer" + ); + + match self.packet_sender.try_send(packet) { + Ok(()) => tracing::trace!( + ?entry_key, + ?tcp_flags, + entry_count = self.entries.count(), + "data plane delivered packet from peer to smoltcp" + ), + Err(err) => tracing::trace!( + ?entry_key, + ?tcp_flags, + ?err, + entry_count = self.entries.count(), + "data plane failed to deliver packet from peer to smoltcp" + ), + } + + None + } +} + +impl DataPlaneRuntime +where + H: VirtualTcpSocketFactory + VirtualTcpListenerFactory + VirtualUdpSocketFactory, +{ + pub(crate) fn new( + runtime_config: CoreRuntimeConfigStore, + peer_manager: Arc, + transport_proxy: Option<&Arc>, + host: Arc, + socket_context: SocketContext, + ) -> Arc { + let (packet_sender, packet_recv) = mpsc::channel(1024); + Arc::new(Self { + operation: Mutex::new(()), + runtime_started: AtomicBool::new(false), + runtime_guard: DataPlaneIoGuard::new(), + runtime_config, + peer_manager: Arc::downgrade(&peer_manager), + transport_proxy: transport_proxy.map(Arc::downgrade), + host, + socket_context, + + runtime_tasks: Arc::new(std::sync::Mutex::new(JoinSet::new())), + packet_recv: Arc::new(Mutex::new(packet_recv)), + packet_sender, + + net: Arc::new(Mutex::new(None)), + entries: Arc::new(FlowTable::default()), + + data_plane_consumers: Arc::new(DataPlaneConsumers::new()), + data_plane_net_ready: tokio::sync::watch::channel(false).0, + pipeline_guard: Mutex::new(None), + }) + } + + fn runtime_ipv4(runtime_config: &CoreRuntimeConfigStore) -> Option { + let prefix = runtime_config + .snapshot() + .peer + .runtime + .core + .routes + .ipv4 + .clone()?; + let IpAddr::V4(address) = prefix.address else { + return None; + }; + cidr::Ipv4Inet::new(address, prefix.prefix_len).ok() + } + + pub(crate) fn is_local_virtual_ip(&self, ip: IpAddr) -> bool { + Self::runtime_ipv4(&self.runtime_config) + .is_some_and(|inet| IpAddr::V4(inet.address()) == ip) + } + + async fn run_net_update_task(self: &Arc) { + let net = self.net.clone(); + let runtime_config = self.runtime_config.clone(); + let peer_manager = self.peer_manager.clone(); + let packet_recv = self.packet_recv.clone(); + let entries = self.entries.clone(); + let data_plane_consumers = self.data_plane_consumers.clone(); + let data_plane_net_ready = self.data_plane_net_ready.clone(); + self.runtime_tasks.lock().unwrap().spawn(async move { + let mut prev_ipv4 = None; + let mut peer_changes = runtime_config.subscribe_peer_runtime_changes(); + loop { + let data_plane_active = data_plane_consumers.has_consumers(); + + if !data_plane_active { + let old_net = { + let mut net_guard = net.lock().await; + // New leases and the zero-consumer teardown decision + // are serialized while this net slot is held. A bind + // can retain this generation or wait for its + // replacement, but cannot receive a stale generation. + if data_plane_consumers.has_consumers() { + continue; + } + net_guard.take() + }; + if let Some(old_net) = &old_net { + old_net.close(DataPlaneErrorKind::HandleClosed); + } + let had_net = old_net.is_some(); + prev_ipv4 = None; + let cleared = entries.clear(); + tracing::trace!( + had_net, + data_plane_active, + removed_entries = cleared.removed, + entry_count = cleared.count.current, + entries_len = entries.len(), + "data plane waiting for consumers" + ); + let _ = data_plane_net_ready.send_replace(false); + select! { + _ = peer_changes.changed() => {} + _ = data_plane_consumers.changed() => {} + } + continue; + } + + let cur_ipv4 = Self::runtime_ipv4(&runtime_config); + if prev_ipv4 != cur_ipv4 { + let old_ipv4 = prev_ipv4; + prev_ipv4 = cur_ipv4; + + tracing::trace!( + ?old_ipv4, + ?cur_ipv4, + old_entry_count = entries.count(), + old_entries_len = entries.len(), + "data plane resetting flows for ipv4 change" + ); + let _ = data_plane_net_ready.send_replace(false); + let old_net = net.lock().await.take(); + if let Some(old_net) = old_net { + old_net.close(DataPlaneErrorKind::NetworkChanged); + } + let cleared = entries.clear(); + tracing::trace!( + ?old_ipv4, + ?cur_ipv4, + removed_entries = cleared.removed, + new_entry_count = cleared.count.current, + new_entries_len = entries.len(), + "data plane reset flows complete" + ); + + if let Some(cur_ipv4) = cur_ipv4 { + net.lock().await.replace(SmoltcpPlane::new( + cur_ipv4, + peer_manager.clone(), + packet_recv.clone(), + )); + tracing::trace!( + ?cur_ipv4, + entry_count = entries.count(), + entries_len = entries.len(), + "data plane installed smoltcp net" + ); + // Wake any data-plane callers waiting in + // `wait_data_plane_net` for the smoltcp net to appear. + let _ = data_plane_net_ready.send_replace(true); + } else { + tracing::trace!( + entry_count = entries.count(), + entries_len = entries.len(), + "data plane removed smoltcp net" + ); + } + } + + select! { + _ = peer_changes.changed() => {} + _ = data_plane_consumers.changed() => {} + } + } + }); + } + + async fn start_runtime_inner(self: &Arc) -> anyhow::Result<()> { + let Some(peer_manager) = self.peer_manager.upgrade() else { + return Err(anyhow::anyhow!("peer manager is gone")); + }; + let guard = peer_manager + .add_managed_packet_process_pipeline(Box::new(self.clone())) + .await; + self.pipeline_guard.lock().await.replace(guard); + + self.runtime_tasks + .lock() + .unwrap() + .spawn(reap_joinset_background( + self.runtime_tasks.clone(), + "data plane runtime", + )); + self.run_net_update_task().await; + + tracing::trace!("data plane peer packet pipeline registered"); + Ok(()) + } + + pub(crate) async fn start_runtime(self: &Arc) -> anyhow::Result<()> { + let _operation = self.operation.lock().await; + if self.runtime_started.load(Ordering::Acquire) { + return Ok(()); + } + if let Err(error) = self.start_runtime_inner().await { + self.stop_runtime_inner().await; + return Err(error); + } + self.runtime_started.store(true, Ordering::Release); + Ok(()) + } + + async fn shutdown_tasks(tasks: &Arc>>) { + let mut tasks = { + let mut guard = tasks.lock().unwrap(); + std::mem::replace(&mut *guard, JoinSet::new()) + }; + tasks.shutdown().await; + } + + async fn stop_runtime_inner(&self) { + self.runtime_started.store(false, Ordering::Release); + self.runtime_guard + .close(DataPlaneErrorKind::InstanceStopped); + if let Some(guard) = self.pipeline_guard.lock().await.take() { + guard.close(); + } + if let Some(net) = self.net.lock().await.take() { + net.close(DataPlaneErrorKind::InstanceStopped); + } + let _ = self.data_plane_net_ready.send_replace(false); + self.entries.clear(); + Self::shutdown_tasks(&self.runtime_tasks).await; + } + + pub(crate) async fn stop_runtime(&self) { + let _operation = self.operation.lock().await; + self.stop_runtime_inner().await; + } +} + +impl DataPlaneRuntime +where + H: VirtualTcpSocketFactory + VirtualTcpListenerFactory + VirtualUdpSocketFactory, +{ + async fn tcp_route_input( + &self, + dst_addr: SocketAddr, + options: DataPlaneTcpConnectOptions, + ) -> DataPlaneResult { + let IpAddr::V4(dst_ip) = dst_addr.ip() else { + return Err(DataPlaneError::new( + DataPlaneErrorKind::AddressFamilyUnsupported, + "the EasyTier data plane currently supports IPv4 destinations only", + )); + }; + self.runtime_guard.ensure_open()?; + let local_virtual_ip = Self::runtime_ipv4(&self.runtime_config).map(|inet| inet.address()); + let local_virtual_destination = local_virtual_ip == Some(dst_ip); + let local_endpoint = local_virtual_destination + && self.entries.contains_key(&FlowKey { + src: dst_addr, + dst: SocketAddr::new(IpAddr::V4(Ipv4Addr::UNSPECIFIED), 0), + kind: TCP_LISTEN_ENTRY, + }); + + let peer_manager = self.peer_manager.upgrade().ok_or_else(|| { + DataPlaneError::new( + DataPlaneErrorKind::InstanceStopped, + "peer manager is no longer available", + ) + })?; + let overlay_destination = if local_virtual_destination { + true + } else { + let (peers, _) = options + .deadline + .run(async { + Ok::<_, DataPlaneError>(peer_manager.get_msg_dst_peer(&dst_addr.ip()).await) + }) + .await?; + !peers.is_empty() + }; + + let prefer_kcp = options.transport == DataPlaneTransportPreference::PreferKcp; + let transport_proxy = self.transport_proxy.as_ref().and_then(Weak::upgrade); + let kcp_ready = if prefer_kcp && overlay_destination { + match &transport_proxy { + Some(proxy) => { + options + .deadline + .run(async { + Ok::<_, DataPlaneError>( + proxy.source_connect_ready(WrappedTransportKind::Kcp).await, + ) + }) + .await? + } + None => false, + } + } else { + false + }; + let kcp_allowed = if prefer_kcp && kcp_ready { + options + .deadline + .run(async { + Ok::<_, DataPlaneError>( + peer_manager.check_allow_kcp_to_dst(&dst_addr.ip()).await, + ) + }) + .await? + } else { + false + }; + + Ok(DataPlaneTcpRouteInput { + policy: options.policy, + transport: options.transport, + local_endpoint, + local_virtual_destination, + overlay_destination, + smoltcp_ready: self.net.lock().await.is_some(), + kcp_ready, + kcp_allowed, + }) + } + + async fn connect_smoltcp_tcp( + &self, + dst_addr: SocketAddr, + deadline: DataPlaneDeadline, + data_plane_ref: DataPlaneLease, + ) -> DataPlaneResult { + let (ipv4_addr, smoltcp_net, generation) = self.wait_data_plane_net(deadline).await?; + let listen_options = TcpListenOptions::port_lease("0.0.0.0:0".parse().unwrap()); + let reservation = generation + .while_open( + deadline.run( + self.host.bind_tcp( + listen_options.clone().with_bind( + listen_options + .bind + .with_context(self.socket_context.clone()), + ), + ), + ), + ) + .await?; + let local_port = reservation + .local_addr() + .map_err(DataPlaneError::from)? + .port(); + let local_addr = SocketAddr::new(IpAddr::V4(ipv4_addr.address()), local_port); + let (flow, _) = FlowLease::register( + self.entries.clone(), + FlowKey { + src: local_addr, + dst: dst_addr, + kind: TCP_ENTRY, + }, + FlowData::Tcp { + _reservation: reservation, + }, + ); + let stream = generation + .while_open(deadline.run(smoltcp_net.tcp_connect(dst_addr, local_port))) + .await?; + generation.ensure_open()?; + + Ok(DataPlaneTcpStream::new( + Box::new(stream), + local_addr, + Some(data_plane_ref), + DataPlaneTcpStreamRoute::Outbound { _flow: flow }, + generation, + )) + } + + async fn connect_host_tcp( + &self, + dst_addr: SocketAddr, + options: DataPlaneTcpConnectOptions, + ) -> DataPlaneResult { + let connect_options = TcpConnectOptions::direct_connect(dst_addr) + .with_purpose(options.purpose) + .with_bind(TcpBindOptions::default().with_context(self.socket_context.clone())); + let socket = self + .runtime_guard + .while_open(options.deadline.run(self.host.connect_tcp(connect_options))) + .await?; + let local_addr = socket.local_addr().map_err(DataPlaneError::from)?; + Ok(DataPlaneTcpStream::new( + Box::new(socket), + local_addr, + None, + DataPlaneTcpStreamRoute::External, + self.runtime_guard.clone(), + )) + } + + async fn connect_kcp_tcp( + &self, + dst_addr: SocketAddr, + options: DataPlaneTcpConnectOptions, + ) -> DataPlaneResult { + let source_addr = options.source_hint.ok_or_else(|| { + DataPlaneError::new( + DataPlaneErrorKind::PathNotReady, + "KCP gateway route requires a logical source address", + ) + })?; + let transport_proxy = self + .transport_proxy + .as_ref() + .and_then(Weak::upgrade) + .ok_or_else(|| { + DataPlaneError::new( + DataPlaneErrorKind::PathNotReady, + "KCP source transport is not ready", + ) + })?; + let stream = self + .runtime_guard + .while_open(options.deadline.run(transport_proxy.connect_source( + WrappedTransportKind::Kcp, + source_addr, + dst_addr, + ))) + .await?; + Ok(DataPlaneTcpStream::new( + stream, + source_addr, + None, + DataPlaneTcpStreamRoute::External, + self.runtime_guard.clone(), + )) + } + + pub(crate) async fn connect_tcp( + &self, + dst_addr: SocketAddr, + options: DataPlaneTcpConnectOptions, + ) -> DataPlaneResult { + let mut route_input = self.tcp_route_input(dst_addr, options).await?; + let route = match route_input.select() { + Ok(route) => route, + Err(error) if error.kind() == DataPlaneErrorKind::PathNotReady => { + let data_plane_ref = self.acquire_data_plane_ref()?; + let _ = self.wait_data_plane_net(options.deadline).await?; + route_input.smoltcp_ready = true; + let route = route_input.select()?; + return self + .connect_selected_tcp(dst_addr, options, route, Some(data_plane_ref)) + .await; + } + Err(error) => return Err(error), + }; + self.connect_selected_tcp(dst_addr, options, route, None) + .await + } + + async fn connect_selected_tcp( + &self, + dst_addr: SocketAddr, + options: DataPlaneTcpConnectOptions, + route: DataPlaneTcpRoute, + data_plane_ref: Option, + ) -> DataPlaneResult { + tracing::debug!(?dst_addr, ?route, "selected data-plane TCP route"); + match route { + DataPlaneTcpRoute::LocalEndpoint | DataPlaneTcpRoute::Smoltcp => { + let data_plane_ref = match data_plane_ref { + Some(lease) => lease, + None => self.acquire_data_plane_ref()?, + }; + self.connect_smoltcp_tcp(dst_addr, options.deadline, data_plane_ref) + .await + } + DataPlaneTcpRoute::LocalHost => { + self.connect_host_tcp( + SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), dst_addr.port()), + options, + ) + .await + } + DataPlaneTcpRoute::Kcp => self.connect_kcp_tcp(dst_addr, options).await, + DataPlaneTcpRoute::Direct => self.connect_host_tcp(dst_addr, options).await, + } + } + + fn acquire_data_plane_ref(&self) -> DataPlaneResult { + if !self.runtime_started.load(Ordering::Acquire) || self.runtime_guard.is_closed() { + return Err(DataPlaneError::new( + DataPlaneErrorKind::InstanceStopped, + "data-plane runtime is not running", + )); + } + let lease = DataPlaneLease::acquire(self.data_plane_consumers.clone()); + if !self.runtime_started.load(Ordering::Acquire) || self.runtime_guard.is_closed() { + drop(lease); + return Err(DataPlaneError::new( + DataPlaneErrorKind::InstanceStopped, + "data-plane runtime stopped while acquiring a resource", + )); + } + Ok(lease) + } + + pub(crate) fn acquire_consumer_lease(&self) -> DataPlaneResult { + self.acquire_data_plane_ref() + .map(DataPlaneConsumerLease::new) + } + + async fn wait_data_plane_net( + &self, + deadline: DataPlaneDeadline, + ) -> DataPlaneResult<(cidr::Ipv4Inet, Arc, DataPlaneIoGuard)> { + let mut ready = self.data_plane_net_ready.subscribe(); + loop { + if let Some(net) = self + .net + .lock() + .await + .as_ref() + .map(|plane| (plane.ipv4_addr, plane.net.clone(), plane.lease())) + { + net.2.ensure_open()?; + return Ok(net); + } + + tokio::select! { + _ = self.runtime_guard.closed() => { + return Err(DataPlaneError::new( + DataPlaneErrorKind::InstanceStopped, + "data-plane runtime stopped while waiting for the stack", + )); + } + result = deadline.run(async { + ready + .wait_for(|ready| *ready) + .await + .map(|_| ()) + .map_err(|_| DataPlaneError::new( + DataPlaneErrorKind::InstanceStopped, + "data-plane readiness channel closed", + )) + }) => { + result?; + } + } + } + } + + pub async fn data_plane_tcp_connect( + &self, + dst_addr: SocketAddr, + timeout: Duration, + ) -> DataPlaneResult { + self.connect_tcp(dst_addr, DataPlaneTcpConnectOptions::public(timeout)) + .await + } + + pub async fn data_plane_tcp_bind( + &self, + local_port: u16, + timeout: Duration, + ) -> DataPlaneResult { + self.data_plane_tcp_bind_with_deadline(local_port, DataPlaneDeadline::from_timeout(timeout)) + .await + } + + async fn data_plane_tcp_bind_with_deadline( + &self, + local_port: u16, + deadline: DataPlaneDeadline, + ) -> DataPlaneResult { + let data_plane_ref = self.acquire_data_plane_ref()?; + let (ipv4_addr, smoltcp_net, generation) = self.wait_data_plane_net(deadline).await?; + let bind_addr = SocketAddr::new(IpAddr::V4(ipv4_addr.address()), local_port); + let listener = generation + .while_open(deadline.run(smoltcp_net.tcp_bind(bind_addr))) + .await?; + generation.ensure_open()?; + let local_addr = listener.local_addr().map_err(DataPlaneError::from)?; + let listen_route = FlowLease::try_register( + self.entries.clone(), + FlowKey { + src: local_addr, + dst: SocketAddr::new(IpAddr::V4(Ipv4Addr::UNSPECIFIED), 0), + kind: TCP_LISTEN_ENTRY, + }, + FlowData::DataPlaneRoute, + ) + .ok_or_else(|| { + DataPlaneError::new( + DataPlaneErrorKind::AddressInUse, + "data-plane TCP listener already exists", + ) + })?; + + Ok(DataPlaneTcpListener { + listener, + local_addr, + flows: self.entries.clone(), + _listen_flow: listen_route, + data_plane_lease: data_plane_ref, + generation, + }) + } + + pub async fn data_plane_udp_bind( + &self, + local_port: u16, + timeout: Duration, + ) -> DataPlaneResult { + self.data_plane_udp_bind_with_deadline(local_port, DataPlaneDeadline::from_timeout(timeout)) + .await + } + + async fn data_plane_udp_bind_with_deadline( + &self, + local_port: u16, + deadline: DataPlaneDeadline, + ) -> DataPlaneResult { + let data_plane_ref = self.acquire_data_plane_ref()?; + let (ipv4_addr, smoltcp_net, generation) = self.wait_data_plane_net(deadline).await?; + let reservation_options = UdpBindOptions::port_lease(SocketAddr::new( + IpAddr::V4(Ipv4Addr::UNSPECIFIED), + local_port, + )) + .with_context(self.socket_context.clone()); + let reservation = generation + .while_open(deadline.run(self.host.bind_udp(reservation_options))) + .await?; + let reserved_port = reservation + .local_addr() + .map_err(DataPlaneError::from)? + .port(); + let bind_addr = SocketAddr::new(IpAddr::V4(ipv4_addr.address()), reserved_port); + let smol = generation + .while_open(deadline.run(smoltcp_net.udp_bind(bind_addr))) + .await?; + generation.ensure_open()?; + let local_addr = smol.local_addr().map_err(DataPlaneError::from)?; + let socket = Arc::new(DataPlaneUdpIo(smol)); + + Ok(DataPlaneUdpSocket { + socket, + flows: self.entries.clone(), + routes: std::sync::Mutex::new(std::collections::HashMap::new()), + local_addr, + _reservation: reservation, + _data_plane_lease: data_plane_ref, + generation, + }) + } +} diff --git a/easytier-core/src/gateway/dataplane/operation.rs b/easytier-core/src/gateway/dataplane/operation.rs new file mode 100644 index 00000000..1ede188e --- /dev/null +++ b/easytier-core/src/gateway/dataplane/operation.rs @@ -0,0 +1,145 @@ +//! Stable operation, completion, and result types for one data-plane session. + +use std::net::SocketAddr; + +use crate::foundation::operation_broker::OperationId; + +use super::DataPlaneErrorKind; + +#[derive(Clone, Copy, Debug, Eq, Hash, Ord, PartialEq, PartialOrd)] +#[repr(transparent)] +pub struct DataPlaneOperationId(OperationId); + +impl DataPlaneOperationId { + pub fn from_raw(value: u64) -> Option { + OperationId::from_raw(value).map(Self) + } + + pub fn get(self) -> u64 { + self.0.get() + } + + pub(super) fn from_broker(operation_id: OperationId) -> Self { + Self(operation_id) + } + + pub(super) fn broker_id(self) -> OperationId { + self.0 + } +} + +#[derive(Clone, Copy, Debug, Eq, Hash, Ord, PartialEq, PartialOrd)] +#[repr(transparent)] +pub struct DataPlaneResourceId(u64); + +impl DataPlaneResourceId { + pub fn from_raw(value: u64) -> Option { + (value != 0).then_some(Self(value)) + } + + pub fn get(self) -> u64 { + self.0 + } +} + +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +#[repr(u16)] +pub enum DataPlaneOperationKind { + TcpConnect = 1, + TcpBind = 2, + TcpAccept = 3, + TcpRead = 4, + TcpWrite = 5, + UdpBind = 6, + UdpReceive = 7, + UdpSend = 8, +} + +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +pub enum DataPlaneCompletionStatus { + Success, + Error(DataPlaneErrorKind), +} + +impl DataPlaneCompletionStatus { + pub fn code(self) -> u16 { + match self { + Self::Success => 0, + Self::Error(kind) => kind as u16, + } + } +} + +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +pub struct DataPlaneCompletionDescriptor { + pub operation_id: DataPlaneOperationId, + pub kind: DataPlaneOperationKind, + pub status: DataPlaneCompletionStatus, +} + +#[derive(Debug)] +pub enum DataPlaneOperationResult { + TcpConnected { + stream: DataPlaneResourceId, + local_addr: SocketAddr, + peer_addr: SocketAddr, + }, + TcpBound { + listener: DataPlaneResourceId, + local_addr: SocketAddr, + }, + TcpAccepted { + stream: DataPlaneResourceId, + local_addr: SocketAddr, + peer_addr: SocketAddr, + }, + TcpRead { + data: Vec, + eof: bool, + }, + TcpWritten { + len: usize, + }, + UdpBound { + socket: DataPlaneResourceId, + local_addr: SocketAddr, + }, + UdpReceived { + data: Vec, + peer_addr: SocketAddr, + truncated: bool, + }, + UdpSent { + len: usize, + }, +} + +impl DataPlaneOperationResult { + pub(super) fn retained_bytes(&self) -> usize { + match self { + Self::TcpRead { data, .. } | Self::UdpReceived { data, .. } => data.capacity(), + _ => 0, + } + } + + pub(super) fn payload_bytes(&self) -> usize { + match self { + Self::TcpRead { data, .. } | Self::UdpReceived { data, .. } => data.len(), + _ => 0, + } + } + + pub(super) fn created_resource(&self) -> Option { + match self { + Self::TcpConnected { stream, .. } | Self::TcpAccepted { stream, .. } => Some(*stream), + Self::TcpBound { listener, .. } => Some(*listener), + Self::UdpBound { socket, .. } => Some(*socket), + Self::TcpRead { .. } + | Self::TcpWritten { .. } + | Self::UdpReceived { .. } + | Self::UdpSent { .. } => None, + } + } +} + +pub type DataPlaneOperationOutcome = Result; diff --git a/easytier-core/src/gateway/dataplane/packet.rs b/easytier-core/src/gateway/dataplane/packet.rs new file mode 100644 index 00000000..1d78f4a0 --- /dev/null +++ b/easytier-core/src/gateway/dataplane/packet.rs @@ -0,0 +1,377 @@ +//! Peer-packet classification for registered data-plane flows. + +use std::net::{IpAddr, Ipv4Addr, SocketAddr}; + +use pnet_packet::{ + Packet, ip::IpNextHeaderProtocols, ipv4::Ipv4Packet, tcp::TcpPacket, udp::UdpPacket, +}; + +use crate::{ + gateway::proxy::ip_reassembler::{IpReassembler, SmolIpv4Packet}, + packet::{PacketType, ZCPacket}, +}; + +use super::flow::{FlowKey, FlowKind, FlowTable}; + +#[derive(Clone, Debug, Eq, PartialEq)] +enum ClassifiedPeerPacket { + Tcp { + entry: FlowKey, + listen_entry: FlowKey, + flags: u8, + }, + Udp { + entry: FlowKey, + }, + FragmentedUdp { + source: Ipv4Addr, + }, + Unsupported, +} +#[derive(Clone, Debug, Eq, PartialEq)] +pub(crate) enum PeerPacketRoute { + Pass, + Unmatched { + entry: FlowKey, + tcp_flags: Option, + }, + Deliver { + entry: FlowKey, + tcp_flags: Option, + }, + FragmentedUdp { + source: Ipv4Addr, + mirror: bool, + }, +} +fn classify_peer_ipv4_payload(payload: &[u8]) -> ClassifiedPeerPacket { + let Some(ipv4) = Ipv4Packet::new(payload) else { + return ClassifiedPeerPacket::Unsupported; + }; + if ipv4.get_version() != 4 { + return ClassifiedPeerPacket::Unsupported; + } + + match ipv4.get_next_level_protocol() { + IpNextHeaderProtocols::Tcp => { + let Some(tcp) = TcpPacket::new(ipv4.payload()) else { + return ClassifiedPeerPacket::Unsupported; + }; + let entry = FlowKey { + dst: SocketAddr::new(ipv4.get_source().into(), tcp.get_source()), + src: SocketAddr::new(ipv4.get_destination().into(), tcp.get_destination()), + kind: FlowKind::Tcp, + }; + let listen_entry = FlowKey { + src: entry.src, + dst: SocketAddr::new(IpAddr::V4(Ipv4Addr::UNSPECIFIED), 0), + kind: FlowKind::TcpListen, + }; + ClassifiedPeerPacket::Tcp { + entry, + listen_entry, + flags: tcp.get_flags(), + } + } + IpNextHeaderProtocols::Udp => { + let smol_ipv4 = SmolIpv4Packet::new_unchecked(ipv4.packet()); + if IpReassembler::is_packet_fragmented(&smol_ipv4) { + return ClassifiedPeerPacket::FragmentedUdp { + source: ipv4.get_source(), + }; + } + let Some(udp) = UdpPacket::new(ipv4.payload()) else { + return ClassifiedPeerPacket::Unsupported; + }; + ClassifiedPeerPacket::Udp { + entry: FlowKey { + dst: SocketAddr::new(ipv4.get_source().into(), udp.get_source()), + src: SocketAddr::new(ipv4.get_destination().into(), udp.get_destination()), + kind: FlowKind::Udp, + }, + } + } + _ => ClassifiedPeerPacket::Unsupported, + } +} +impl FlowTable { + pub fn route_peer_packet( + &self, + packet: &ZCPacket, + allow_tcp_listen_fallback: bool, + ) -> PeerPacketRoute { + let Some(header) = packet.peer_manager_header() else { + return PeerPacketRoute::Pass; + }; + let is_modified_source = matches!( + header.packet_type, + x if x == PacketType::DataWithKcpSrcModified as u8 + || x == PacketType::DataWithQuicSrcModified as u8 + ); + if header.packet_type != PacketType::Data as u8 && !is_modified_source { + return PeerPacketRoute::Pass; + } + if is_modified_source && header.from_peer_id != header.to_peer_id { + return PeerPacketRoute::Pass; + } + + self.route_peer_ipv4_payload(packet.payload(), allow_tcp_listen_fallback) + } + + pub fn route_peer_ipv4_payload( + &self, + payload: &[u8], + allow_tcp_listen_fallback: bool, + ) -> PeerPacketRoute { + let (entry, tcp_flags) = match classify_peer_ipv4_payload(payload) { + ClassifiedPeerPacket::Tcp { + entry, + listen_entry, + flags, + } => { + let entry = if allow_tcp_listen_fallback && !self.contains_key(&entry) { + listen_entry + } else { + entry + }; + (entry, Some(flags)) + } + ClassifiedPeerPacket::Udp { entry } => (entry, None), + ClassifiedPeerPacket::FragmentedUdp { source } => { + return PeerPacketRoute::FragmentedUdp { + source, + mirror: self.contains_destination_ip(source.into()), + }; + } + ClassifiedPeerPacket::Unsupported => return PeerPacketRoute::Pass, + }; + + if self.contains_key(&entry) { + PeerPacketRoute::Deliver { entry, tcp_flags } + } else { + PeerPacketRoute::Unmatched { entry, tcp_flags } + } + } +} + +#[cfg(test)] +mod tests { + use std::net::{IpAddr, Ipv4Addr, SocketAddr}; + + use pnet_packet::{ + MutablePacket, + ip::IpNextHeaderProtocols, + ipv4::MutableIpv4Packet, + tcp::{MutableTcpPacket, TcpFlags}, + udp::MutableUdpPacket, + }; + + use super::*; + use crate::packet::{PacketType, ZCPacket}; + + fn ipv4_packet(protocol: pnet_packet::ip::IpNextHeaderProtocol, payload_len: usize) -> Vec { + let mut packet = vec![0; 20 + payload_len]; + let packet_len = packet.len() as u16; + let mut ipv4 = MutableIpv4Packet::new(&mut packet).unwrap(); + ipv4.set_version(4); + ipv4.set_header_length(5); + ipv4.set_total_length(packet_len); + ipv4.set_source(Ipv4Addr::new(10, 1, 1, 2)); + ipv4.set_destination(Ipv4Addr::new(10, 2, 2, 3)); + ipv4.set_next_level_protocol(protocol); + packet + } + #[test] + fn classifies_tcp_and_listen_keys() { + let mut packet = ipv4_packet(IpNextHeaderProtocols::Tcp, 20); + let mut ipv4 = MutableIpv4Packet::new(&mut packet).unwrap(); + let mut tcp = MutableTcpPacket::new(ipv4.payload_mut()).unwrap(); + tcp.set_source(1234); + tcp.set_destination(4321); + tcp.set_flags(TcpFlags::SYN); + + assert_eq!( + classify_peer_ipv4_payload(&packet), + ClassifiedPeerPacket::Tcp { + entry: FlowKey { + src: "10.2.2.3:4321".parse().unwrap(), + dst: "10.1.1.2:1234".parse().unwrap(), + kind: FlowKind::Tcp, + }, + listen_entry: FlowKey { + src: "10.2.2.3:4321".parse().unwrap(), + dst: SocketAddr::new(IpAddr::V4(Ipv4Addr::UNSPECIFIED), 0), + kind: FlowKind::TcpListen, + }, + flags: TcpFlags::SYN, + } + ); + } + #[test] + fn classifies_udp_and_fragmented_udp() { + let mut packet = ipv4_packet(IpNextHeaderProtocols::Udp, 8); + let mut ipv4 = MutableIpv4Packet::new(&mut packet).unwrap(); + let mut udp = MutableUdpPacket::new(ipv4.payload_mut()).unwrap(); + udp.set_source(1234); + udp.set_destination(4321); + assert_eq!( + classify_peer_ipv4_payload(&packet), + ClassifiedPeerPacket::Udp { + entry: FlowKey { + src: "10.2.2.3:4321".parse().unwrap(), + dst: "10.1.1.2:1234".parse().unwrap(), + kind: FlowKind::Udp, + } + } + ); + + let mut fragmented = ipv4_packet(IpNextHeaderProtocols::Udp, 8); + MutableIpv4Packet::new(&mut fragmented) + .unwrap() + .set_fragment_offset(1); + assert_eq!( + classify_peer_ipv4_payload(&fragmented), + ClassifiedPeerPacket::FragmentedUdp { + source: Ipv4Addr::new(10, 1, 1, 2), + } + ); + } + #[test] + fn rejects_malformed_and_unsupported_packets() { + assert_eq!( + classify_peer_ipv4_payload(&[]), + ClassifiedPeerPacket::Unsupported + ); + assert_eq!( + classify_peer_ipv4_payload(&ipv4_packet(IpNextHeaderProtocols::Icmp, 8)), + ClassifiedPeerPacket::Unsupported + ); + } + #[test] + fn flow_table_routes_tcp_exact_and_listen_fallback() { + let mut packet = ipv4_packet(IpNextHeaderProtocols::Tcp, 20); + let mut ipv4 = MutableIpv4Packet::new(&mut packet).unwrap(); + let mut tcp = MutableTcpPacket::new(ipv4.payload_mut()).unwrap(); + tcp.set_source(1234); + tcp.set_destination(4321); + tcp.set_flags(TcpFlags::SYN); + + let exact = FlowKey { + src: "10.2.2.3:4321".parse().unwrap(), + dst: "10.1.1.2:1234".parse().unwrap(), + kind: FlowKind::Tcp, + }; + let listen = FlowKey { + src: exact.src, + dst: SocketAddr::new(IpAddr::V4(Ipv4Addr::UNSPECIFIED), 0), + kind: FlowKind::TcpListen, + }; + let table = FlowTable::default(); + + assert_eq!( + table.route_peer_ipv4_payload(&packet, false), + PeerPacketRoute::Unmatched { + entry: exact.clone(), + tcp_flags: Some(TcpFlags::SYN), + } + ); + + table.insert(listen.clone(), ()); + assert_eq!( + table.route_peer_ipv4_payload(&packet, true), + PeerPacketRoute::Deliver { + entry: listen, + tcp_flags: Some(TcpFlags::SYN), + } + ); + + table.insert(exact.clone(), ()); + assert_eq!( + table.route_peer_ipv4_payload(&packet, true), + PeerPacketRoute::Deliver { + entry: exact, + tcp_flags: Some(TcpFlags::SYN), + } + ); + } + #[test] + fn flow_table_routes_fragmented_udp_by_source_ip() { + let mut packet = ipv4_packet(IpNextHeaderProtocols::Udp, 8); + MutableIpv4Packet::new(&mut packet) + .unwrap() + .set_fragment_offset(1); + let table = FlowTable::default(); + + assert_eq!( + table.route_peer_ipv4_payload(&packet, false), + PeerPacketRoute::FragmentedUdp { + source: Ipv4Addr::new(10, 1, 1, 2), + mirror: false, + } + ); + + table.insert( + FlowKey { + src: "10.2.2.3:4321".parse().unwrap(), + dst: "10.1.1.2:1234".parse().unwrap(), + kind: FlowKind::Udp, + }, + (), + ); + assert_eq!( + table.route_peer_ipv4_payload(&packet, false), + PeerPacketRoute::FragmentedUdp { + source: Ipv4Addr::new(10, 1, 1, 2), + mirror: true, + } + ); + } + #[test] + fn flow_table_routes_loopback_modified_source_packets() { + let mut payload = ipv4_packet(IpNextHeaderProtocols::Tcp, 20); + let mut ipv4 = MutableIpv4Packet::new(&mut payload).unwrap(); + let mut tcp = MutableTcpPacket::new(ipv4.payload_mut()).unwrap(); + tcp.set_source(1234); + tcp.set_destination(4321); + let entry = FlowKey { + src: "10.2.2.3:4321".parse().unwrap(), + dst: "10.1.1.2:1234".parse().unwrap(), + kind: FlowKind::Tcp, + }; + let table = FlowTable::default(); + table.insert(entry.clone(), ()); + + for packet_type in [ + PacketType::DataWithKcpSrcModified, + PacketType::DataWithQuicSrcModified, + ] { + let mut packet = ZCPacket::new_with_payload(&payload); + packet.fill_peer_manager_hdr(7, 7, packet_type as u8); + assert_eq!( + table.route_peer_packet(&packet, false), + PeerPacketRoute::Deliver { + entry: entry.clone(), + tcp_flags: Some(0), + } + ); + } + } + #[test] + fn flow_table_passes_non_loopback_or_malformed_modified_source_packets() { + let table = FlowTable::<()>::default(); + let mut non_loopback = + ZCPacket::new_with_payload(&ipv4_packet(IpNextHeaderProtocols::Tcp, 20)); + non_loopback.fill_peer_manager_hdr(7, 8, PacketType::DataWithKcpSrcModified as u8); + assert_eq!( + table.route_peer_packet(&non_loopback, false), + PeerPacketRoute::Pass + ); + + let mut malformed = ZCPacket::new_with_payload(&[0u8; 8]); + malformed.fill_peer_manager_hdr(7, 7, PacketType::DataWithQuicSrcModified as u8); + assert_eq!( + table.route_peer_packet(&malformed, false), + PeerPacketRoute::Pass + ); + } +} diff --git a/easytier-core/src/gateway/dataplane/resource.rs b/easytier-core/src/gateway/dataplane/resource.rs new file mode 100644 index 00000000..a06b9e95 --- /dev/null +++ b/easytier-core/src/gateway/dataplane/resource.rs @@ -0,0 +1,165 @@ +//! Shared lifetime and close guards held by data-plane resources. + +use std::{ + future::Future, + pin::Pin, + sync::{Arc, Mutex}, +}; + +use crossbeam::atomic::AtomicCell; +use tokio::sync::Notify; +use tokio_util::sync::CancellationToken; + +use super::{DataPlaneError, DataPlaneErrorKind, DataPlaneResult}; + +pub(super) struct DataPlaneLease { + consumers: Arc, +} + +pub(super) struct DataPlaneConsumers { + refs: Mutex, + notifier: Notify, +} + +pub(crate) struct DataPlaneConsumerLease { + _lease: DataPlaneLease, +} + +impl DataPlaneConsumerLease { + pub(super) fn new(lease: DataPlaneLease) -> Self { + Self { _lease: lease } + } +} + +impl DataPlaneLease { + pub(super) fn acquire(consumers: Arc) -> Self { + consumers.acquire(); + Self { consumers } + } +} + +impl Clone for DataPlaneLease { + fn clone(&self) -> Self { + self.consumers.clone_ref(); + Self { + consumers: self.consumers.clone(), + } + } +} + +impl Drop for DataPlaneLease { + fn drop(&mut self) { + self.consumers.release(); + } +} + +impl DataPlaneConsumers { + pub(super) fn new() -> Self { + Self { + refs: Mutex::new(0), + notifier: Notify::new(), + } + } + + pub(super) fn has_consumers(&self) -> bool { + *self.refs.lock().unwrap() != 0 + } + + pub(super) async fn changed(&self) { + self.notifier.notified().await; + } + + fn acquire(&self) { + let mut refs = self.refs.lock().unwrap(); + *refs = refs + .checked_add(1) + .expect("data-plane consumer reference count overflow"); + drop(refs); + self.notifier.notify_one(); + } + + fn clone_ref(&self) { + let mut refs = self.refs.lock().unwrap(); + *refs = refs + .checked_add(1) + .expect("data-plane consumer reference count overflow"); + } + + fn release(&self) { + let mut refs = self.refs.lock().unwrap(); + debug_assert_ne!(*refs, 0); + *refs -= 1; + let final_ref = *refs == 0; + drop(refs); + if final_ref { + self.notifier.notify_one(); + } + } +} + +#[derive(Clone)] +pub(super) struct DataPlaneIoGuard { + closed: CancellationToken, + close_kind: Arc>>, +} + +impl DataPlaneIoGuard { + pub(super) fn new() -> Self { + Self { + closed: CancellationToken::new(), + close_kind: Arc::new(AtomicCell::new(None)), + } + } + + pub(super) fn ensure_open(&self) -> DataPlaneResult<()> { + if self.closed.is_cancelled() { + return Err(self.closed_error()); + } + Ok(()) + } + + pub(super) fn is_closed(&self) -> bool { + self.closed.is_cancelled() + } + + pub(super) async fn closed(&self) { + self.closed.cancelled().await; + } + + pub(super) async fn while_open( + &self, + future: impl Future>, + ) -> DataPlaneResult { + self.ensure_open()?; + tokio::select! { + biased; + _ = self.closed() => Err(self.closed_error()), + result = future => result, + } + } + + pub(super) fn closed_future(&self) -> Pin + Send>> { + let closed = self.closed.clone(); + Box::pin(async move { + closed.cancelled().await; + }) + } + + pub(super) fn close(&self, kind: DataPlaneErrorKind) { + if self.close_kind.compare_exchange(None, Some(kind)).is_ok() { + self.closed.cancel(); + } + } + + pub(super) fn closed_error(&self) -> DataPlaneError { + let kind = self + .close_kind + .load() + .unwrap_or(DataPlaneErrorKind::HandleClosed); + DataPlaneError::new(kind, "data-plane I/O resource closed") + } + + pub(super) fn closed_io_error(&self) -> std::io::Error { + self.closed_error().into_io_error() + } +} diff --git a/easytier-core/src/gateway/dataplane/route.rs b/easytier-core/src/gateway/dataplane/route.rs new file mode 100644 index 00000000..356da537 --- /dev/null +++ b/easytier-core/src/gateway/dataplane/route.rs @@ -0,0 +1,166 @@ +//! Pure route selection for TCP data-plane connections. + +use super::{DataPlaneError, DataPlaneErrorKind, DataPlaneResult}; + +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +pub(crate) enum DataPlaneRoutePolicy { + OverlayOnly, + OverlayOrDirect, +} + +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +pub(crate) enum DataPlaneTransportPreference { + SmoltcpOnly, + PreferKcp, +} + +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +pub(super) enum DataPlaneTcpRoute { + LocalEndpoint, + LocalHost, + Smoltcp, + Kcp, + Direct, +} + +#[derive(Clone, Copy, Debug)] +pub(super) struct DataPlaneTcpRouteInput { + pub policy: DataPlaneRoutePolicy, + pub transport: DataPlaneTransportPreference, + pub local_endpoint: bool, + pub local_virtual_destination: bool, + pub overlay_destination: bool, + pub smoltcp_ready: bool, + pub kcp_ready: bool, + pub kcp_allowed: bool, +} + +impl DataPlaneTcpRouteInput { + pub(super) fn select(self) -> DataPlaneResult { + if self.local_endpoint { + return self + .smoltcp_ready + .then_some(DataPlaneTcpRoute::LocalEndpoint) + .ok_or_else(path_not_ready); + } + if self.local_virtual_destination { + return Ok(DataPlaneTcpRoute::LocalHost); + } + if self.overlay_destination { + if self.transport == DataPlaneTransportPreference::PreferKcp + && self.kcp_ready + && self.kcp_allowed + { + return Ok(DataPlaneTcpRoute::Kcp); + } + return self + .smoltcp_ready + .then_some(DataPlaneTcpRoute::Smoltcp) + .ok_or_else(path_not_ready); + } + if self.policy == DataPlaneRoutePolicy::OverlayOrDirect { + return Ok(DataPlaneTcpRoute::Direct); + } + Err(DataPlaneError::new( + DataPlaneErrorKind::NoOverlayRoute, + "destination has no EasyTier overlay route", + )) + } +} + +fn path_not_ready() -> DataPlaneError { + DataPlaneError::new( + DataPlaneErrorKind::PathNotReady, + "selected EasyTier data-plane path is not ready", + ) +} + +#[cfg(test)] +mod tests { + use super::*; + + fn overlay() -> DataPlaneTcpRouteInput { + DataPlaneTcpRouteInput { + policy: DataPlaneRoutePolicy::OverlayOnly, + transport: DataPlaneTransportPreference::SmoltcpOnly, + local_endpoint: false, + local_virtual_destination: false, + overlay_destination: true, + smoltcp_ready: true, + kcp_ready: false, + kcp_allowed: false, + } + } + + #[test] + fn overlay_only_rejects_unrelated_destination() { + let input = DataPlaneTcpRouteInput { + overlay_destination: false, + ..overlay() + }; + + assert_eq!( + input.select().unwrap_err().kind(), + DataPlaneErrorKind::NoOverlayRoute + ); + } + + #[test] + fn gateway_policy_can_select_direct_host_route() { + let input = DataPlaneTcpRouteInput { + policy: DataPlaneRoutePolicy::OverlayOrDirect, + overlay_destination: false, + ..overlay() + }; + + assert_eq!(input.select().unwrap(), DataPlaneTcpRoute::Direct); + } + + #[test] + fn smoltcp_only_ignores_available_kcp() { + let input = DataPlaneTcpRouteInput { + kcp_ready: true, + kcp_allowed: true, + ..overlay() + }; + + assert_eq!(input.select().unwrap(), DataPlaneTcpRoute::Smoltcp); + } + + #[test] + fn gateway_preference_selects_allowed_ready_kcp() { + let input = DataPlaneTcpRouteInput { + transport: DataPlaneTransportPreference::PreferKcp, + kcp_ready: true, + kcp_allowed: true, + ..overlay() + }; + + assert_eq!(input.select().unwrap(), DataPlaneTcpRoute::Kcp); + } + + #[test] + fn known_overlay_without_a_ready_path_is_not_direct() { + let input = DataPlaneTcpRouteInput { + policy: DataPlaneRoutePolicy::OverlayOrDirect, + smoltcp_ready: false, + ..overlay() + }; + + assert_eq!( + input.select().unwrap_err().kind(), + DataPlaneErrorKind::PathNotReady + ); + } + + #[test] + fn local_endpoint_precedes_local_host_mapping() { + let input = DataPlaneTcpRouteInput { + local_endpoint: true, + local_virtual_destination: true, + ..overlay() + }; + + assert_eq!(input.select().unwrap(), DataPlaneTcpRoute::LocalEndpoint); + } +} diff --git a/easytier-core/src/gateway/dataplane/session.rs b/easytier-core/src/gateway/dataplane/session.rs new file mode 100644 index 00000000..df9c6047 --- /dev/null +++ b/easytier-core/src/gateway/dataplane/session.rs @@ -0,0 +1,1540 @@ +//! Instance-scoped data-plane resources, operations, and completion delivery. + +use std::{ + collections::{HashMap, HashSet}, + future::Future, + net::SocketAddr, + sync::{Arc, Condvar, Mutex, MutexGuard, Weak}, + time::Duration, +}; + +use tokio::{ + io::{AsyncReadExt, AsyncWriteExt, ReadHalf, WriteHalf}, + sync::Mutex as AsyncMutex, +}; +use tokio_util::sync::CancellationToken; + +use crate::{ + foundation::operation_broker::{ + AccessError as OperationAccessError, AdmissionError, OperationBroker, + }, + socket::{ + tcp::{VirtualTcpListenerFactory, VirtualTcpSocketFactory}, + udp::VirtualUdpSocketFactory, + }, +}; + +use super::{ + DataPlaneConsumerLease, DataPlaneDeadline, DataPlaneError, DataPlaneErrorKind, DataPlaneResult, + DataPlaneRuntime, DataPlaneTcpConnectOptions, DataPlaneTcpListener, DataPlaneTcpStream, + DataPlaneUdpSocket, + operation::{ + DataPlaneCompletionDescriptor, DataPlaneCompletionStatus, DataPlaneOperationId, + DataPlaneOperationKind, DataPlaneOperationOutcome, DataPlaneOperationResult, + DataPlaneResourceId, + }, +}; + +const DEFAULT_MAX_RESOURCES: usize = 4_096; +const DEFAULT_MAX_OPERATIONS: usize = 4_096; +const DEFAULT_MAX_RESULT_BYTES: usize = 64 * 1024 * 1024; +const DEFAULT_MAX_READ_SIZE: usize = 1024 * 1024; + +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +pub struct DataPlaneSessionLimits { + pub max_resources: usize, + pub max_operations: usize, + pub max_result_bytes: usize, + pub max_read_size: usize, +} + +impl Default for DataPlaneSessionLimits { + fn default() -> Self { + Self { + max_resources: DEFAULT_MAX_RESOURCES, + max_operations: DEFAULT_MAX_OPERATIONS, + max_result_bytes: DEFAULT_MAX_RESULT_BYTES, + max_read_size: DEFAULT_MAX_READ_SIZE, + } + } +} + +struct TcpResource { + read: AsyncMutex>, + write: AsyncMutex>, +} + +struct UdpResource { + socket: Arc, + read: AsyncMutex<()>, + write: AsyncMutex<()>, +} + +#[derive(Clone)] +enum ResourceIo { + Tcp(Arc), + TcpListener(Arc>), + Udp(Arc), +} + +struct ResourceEntry { + io: ResourceIo, + pending_operations: HashSet, +} + +#[derive(Default)] +struct ResourceTable { + next_resource_id: u64, + entries: HashMap, + reserved: usize, +} + +struct OperationMetadata { + target: Option, + reserved_result_bytes: usize, + reserves_resource: bool, +} + +type DataPlaneOperationBroker = + OperationBroker; + +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +enum SessionLifecycle { + Created, + Running, + Stopped, +} + +struct SessionState { + lifecycle: SessionLifecycle, + resources: ResourceTable, + broker: DataPlaneOperationBroker, + retained_result_bytes: usize, +} + +impl SessionState { + fn new(max_operations: usize) -> Self { + Self { + lifecycle: SessionLifecycle::Created, + resources: ResourceTable { + next_resource_id: 1, + ..ResourceTable::default() + }, + broker: OperationBroker::new(max_operations), + retained_result_bytes: 0, + } + } +} + +enum PendingOperationResult { + TcpConnected { + stream: DataPlaneTcpStream, + peer_addr: SocketAddr, + }, + TcpBound(DataPlaneTcpListener), + TcpAccepted { + stream: DataPlaneTcpStream, + peer_addr: SocketAddr, + }, + TcpRead { + data: Vec, + eof: bool, + }, + TcpWritten(usize), + UdpBound(DataPlaneUdpSocket), + UdpReceived { + data: Vec, + peer_addr: SocketAddr, + truncated: bool, + }, + UdpSent(usize), +} + +pub struct DataPlaneSession +where + H: VirtualTcpSocketFactory + VirtualTcpListenerFactory + VirtualUdpSocketFactory, +{ + runtime: Weak>, + consumer_lease: Mutex>, + limits: DataPlaneSessionLimits, + state: Mutex, + completion_condvar: Condvar, + completion_notify: tokio::sync::Notify, +} + +impl DataPlaneSession +where + H: VirtualTcpSocketFactory + VirtualTcpListenerFactory + VirtualUdpSocketFactory, +{ + pub(crate) fn new(runtime: &Arc>) -> Arc { + Self::with_runtime(Arc::downgrade(runtime), DataPlaneSessionLimits::default()) + } + + fn with_runtime( + runtime: Weak>, + limits: DataPlaneSessionLimits, + ) -> Arc { + Arc::new(Self { + runtime, + consumer_lease: Mutex::new(None), + limits, + state: Mutex::new(SessionState::new(limits.max_operations)), + completion_condvar: Condvar::new(), + completion_notify: tokio::sync::Notify::new(), + }) + } + + fn lock_state(&self) -> MutexGuard<'_, SessionState> { + self.state.lock().unwrap_or_else(|error| error.into_inner()) + } + + fn error(kind: DataPlaneErrorKind, message: &'static str) -> DataPlaneError { + DataPlaneError::new(kind, message) + } + + fn ensure_executor() -> DataPlaneResult<()> { + tokio::runtime::Handle::try_current() + .map(|_| ()) + .map_err(|_| { + Self::error( + DataPlaneErrorKind::PathNotReady, + "data-plane operations require an active Tokio runtime", + ) + }) + } + + fn runtime(&self) -> DataPlaneResult>> { + let runtime = self.runtime.upgrade().ok_or_else(|| { + Self::error( + DataPlaneErrorKind::InstanceStopped, + "data-plane runtime is no longer available", + ) + })?; + let mut lease = self + .consumer_lease + .lock() + .unwrap_or_else(|error| error.into_inner()); + if lease.is_none() { + lease.replace(runtime.acquire_consumer_lease()?); + } + Ok(runtime) + } + + pub(crate) fn start(&self) -> DataPlaneResult<()> { + let mut state = self.lock_state(); + match state.lifecycle { + SessionLifecycle::Created => { + state.lifecycle = SessionLifecycle::Running; + Ok(()) + } + SessionLifecycle::Running => Ok(()), + SessionLifecycle::Stopped => Err(Self::error( + DataPlaneErrorKind::InstanceStopped, + "data-plane session is stopped", + )), + } + } + + pub fn limits(&self) -> DataPlaneSessionLimits { + self.limits + } + + fn next_resource_id(resources: &mut ResourceTable) -> DataPlaneResult { + let id = DataPlaneResourceId::from_raw(resources.next_resource_id).ok_or_else(|| { + Self::error( + DataPlaneErrorKind::ResourceLimit, + "data-plane resource ID space is exhausted", + ) + })?; + resources.next_resource_id = + resources.next_resource_id.checked_add(1).ok_or_else(|| { + Self::error( + DataPlaneErrorKind::ResourceLimit, + "data-plane resource ID space is exhausted", + ) + })?; + Ok(id) + } + + fn admit_locked( + &self, + state: &mut SessionState, + kind: DataPlaneOperationKind, + target: Option, + reserved_result_bytes: usize, + reserves_resource: bool, + ) -> DataPlaneResult<(DataPlaneOperationId, CancellationToken)> { + if state.lifecycle != SessionLifecycle::Running { + return Err(Self::error( + DataPlaneErrorKind::InstanceStopped, + "data-plane session is not running", + )); + } + let retained_result_bytes = state + .retained_result_bytes + .checked_add(reserved_result_bytes) + .ok_or_else(|| { + Self::error( + DataPlaneErrorKind::ResourceLimit, + "data-plane result-byte accounting overflow", + ) + })?; + if retained_result_bytes > self.limits.max_result_bytes { + return Err(Self::error( + DataPlaneErrorKind::ResourceLimit, + "data-plane retained-result limit reached", + )); + } + if reserves_resource { + let resource_usage = state + .resources + .entries + .len() + .checked_add(state.resources.reserved) + .ok_or_else(|| { + Self::error( + DataPlaneErrorKind::ResourceLimit, + "data-plane resource accounting overflow", + ) + })?; + if resource_usage >= self.limits.max_resources { + return Err(Self::error( + DataPlaneErrorKind::ResourceLimit, + "too many open data-plane resources", + )); + } + } + + let admission = state + .broker + .admit( + kind, + OperationMetadata { + target, + reserved_result_bytes, + reserves_resource, + }, + ) + .map_err(|error| match error { + AdmissionError::AtCapacity => Self::error( + DataPlaneErrorKind::ResourceLimit, + "too many outstanding data-plane operations", + ), + AdmissionError::IdExhausted => Self::error( + DataPlaneErrorKind::ResourceLimit, + "data-plane operation ID space is exhausted", + ), + })?; + let operation_id = DataPlaneOperationId::from_broker(admission.id); + state.retained_result_bytes = retained_result_bytes; + if reserves_resource { + state.resources.reserved += 1; + } + if let Some(resource_id) = target { + state + .resources + .entries + .get_mut(&resource_id) + .expect("target resource was validated before admission") + .pending_operations + .insert(operation_id); + } + Ok((operation_id, admission.cancellation)) + } + + fn require_read_size(&self, max_len: usize) -> DataPlaneResult<()> { + if max_len > self.limits.max_read_size { + return Err(Self::error( + DataPlaneErrorKind::ResourceLimit, + "data-plane read allocation exceeds the per-operation limit", + )); + } + Ok(()) + } + + fn resource_io( + state: &SessionState, + resource_id: DataPlaneResourceId, + ) -> DataPlaneResult { + state + .resources + .entries + .get(&resource_id) + .map(|resource| resource.io.clone()) + .ok_or_else(|| { + Self::error( + DataPlaneErrorKind::HandleClosed, + "data-plane resource is closed", + ) + }) + } + + fn require_tcp( + state: &SessionState, + resource_id: DataPlaneResourceId, + ) -> DataPlaneResult> { + match Self::resource_io(state, resource_id)? { + ResourceIo::Tcp(resource) => Ok(resource), + ResourceIo::TcpListener(_) | ResourceIo::Udp(_) => Err(Self::error( + DataPlaneErrorKind::HandleClosed, + "data-plane resource is not a TCP stream", + )), + } + } + + fn require_tcp_listener( + state: &SessionState, + resource_id: DataPlaneResourceId, + ) -> DataPlaneResult>> { + match Self::resource_io(state, resource_id)? { + ResourceIo::TcpListener(resource) => Ok(resource), + ResourceIo::Tcp(_) | ResourceIo::Udp(_) => Err(Self::error( + DataPlaneErrorKind::HandleClosed, + "data-plane resource is not a TCP listener", + )), + } + } + + fn require_udp( + state: &SessionState, + resource_id: DataPlaneResourceId, + ) -> DataPlaneResult> { + match Self::resource_io(state, resource_id)? { + ResourceIo::Udp(resource) => Ok(resource), + ResourceIo::Tcp(_) | ResourceIo::TcpListener(_) => Err(Self::error( + DataPlaneErrorKind::HandleClosed, + "data-plane resource is not a UDP socket", + )), + } + } + + async fn run_operation( + cancel: CancellationToken, + deadline: DataPlaneDeadline, + future: impl Future>, + ) -> DataPlaneResult + where + E: Into, + { + tokio::select! { + biased; + _ = cancel.cancelled() => Err(Self::error( + DataPlaneErrorKind::Cancelled, + "data-plane operation cancelled", + )), + result = deadline.run(future) => result, + } + } + + fn spawn_operation( + self: &Arc, + operation_id: DataPlaneOperationId, + future: impl Future> + Send + 'static, + ) { + let session = Arc::downgrade(self); + tokio::spawn(async move { + let result = future.await; + if let Some(session) = session.upgrade() { + session.complete_operation(operation_id, result); + } + }); + } + + pub fn submit_tcp_connect( + self: &Arc, + peer_addr: SocketAddr, + timeout: Option, + ) -> DataPlaneResult { + Self::ensure_executor()?; + let deadline = DataPlaneDeadline::from_optional_timeout(timeout); + let runtime = self.runtime()?; + let (operation_id, cancel) = { + let mut state = self.lock_state(); + self.admit_locked( + &mut state, + DataPlaneOperationKind::TcpConnect, + None, + 0, + true, + )? + }; + self.spawn_operation(operation_id, async move { + let options = DataPlaneTcpConnectOptions::public_with_deadline(deadline); + let stream = + Self::run_operation(cancel, deadline, runtime.connect_tcp(peer_addr, options)) + .await?; + Ok(PendingOperationResult::TcpConnected { stream, peer_addr }) + }); + Ok(operation_id) + } + + pub fn submit_tcp_bind( + self: &Arc, + local_port: u16, + timeout: Option, + ) -> DataPlaneResult { + Self::ensure_executor()?; + let deadline = DataPlaneDeadline::from_optional_timeout(timeout); + let runtime = self.runtime()?; + let (operation_id, cancel) = { + let mut state = self.lock_state(); + self.admit_locked(&mut state, DataPlaneOperationKind::TcpBind, None, 0, true)? + }; + self.spawn_operation(operation_id, async move { + let listener = Self::run_operation( + cancel, + deadline, + runtime.data_plane_tcp_bind_with_deadline(local_port, deadline), + ) + .await?; + Ok(PendingOperationResult::TcpBound(listener)) + }); + Ok(operation_id) + } + + pub fn submit_tcp_accept( + self: &Arc, + listener_id: DataPlaneResourceId, + timeout: Option, + ) -> DataPlaneResult { + Self::ensure_executor()?; + let deadline = DataPlaneDeadline::from_optional_timeout(timeout); + let (listener, operation_id, cancel) = { + let mut state = self.lock_state(); + let listener = Self::require_tcp_listener(&state, listener_id)?; + let (operation_id, cancel) = self.admit_locked( + &mut state, + DataPlaneOperationKind::TcpAccept, + Some(listener_id), + 0, + true, + )?; + (listener, operation_id, cancel) + }; + self.spawn_operation(operation_id, async move { + let (stream, peer_addr) = Self::run_operation(cancel, deadline, async move { + listener.lock().await.accept().await + }) + .await?; + Ok(PendingOperationResult::TcpAccepted { stream, peer_addr }) + }); + Ok(operation_id) + } + + pub fn submit_tcp_read( + self: &Arc, + stream_id: DataPlaneResourceId, + max_len: usize, + timeout: Option, + ) -> DataPlaneResult { + Self::ensure_executor()?; + self.require_read_size(max_len)?; + let deadline = DataPlaneDeadline::from_optional_timeout(timeout); + let (stream, operation_id, cancel) = { + let mut state = self.lock_state(); + let stream = Self::require_tcp(&state, stream_id)?; + let (operation_id, cancel) = self.admit_locked( + &mut state, + DataPlaneOperationKind::TcpRead, + Some(stream_id), + max_len, + false, + )?; + (stream, operation_id, cancel) + }; + if max_len == 0 { + self.complete_operation( + operation_id, + Ok(PendingOperationResult::TcpRead { + data: Vec::new(), + eof: false, + }), + ); + return Ok(operation_id); + } + self.spawn_operation(operation_id, async move { + let (data, eof) = Self::run_operation(cancel, deadline, async move { + let mut data = vec![0u8; max_len]; + let len = stream.read.lock().await.read(&mut data).await?; + data.truncate(len); + Ok::<_, std::io::Error>((data, len == 0)) + }) + .await?; + Ok(PendingOperationResult::TcpRead { data, eof }) + }); + Ok(operation_id) + } + + pub fn submit_tcp_write( + self: &Arc, + stream_id: DataPlaneResourceId, + data: Vec, + timeout: Option, + ) -> DataPlaneResult { + Self::ensure_executor()?; + let deadline = DataPlaneDeadline::from_optional_timeout(timeout); + let (stream, operation_id, cancel) = { + let mut state = self.lock_state(); + let stream = Self::require_tcp(&state, stream_id)?; + let (operation_id, cancel) = self.admit_locked( + &mut state, + DataPlaneOperationKind::TcpWrite, + Some(stream_id), + 0, + false, + )?; + (stream, operation_id, cancel) + }; + self.spawn_operation(operation_id, async move { + let len = Self::run_operation(cancel, deadline, async move { + stream.write.lock().await.write(&data).await + }) + .await?; + Ok(PendingOperationResult::TcpWritten(len)) + }); + Ok(operation_id) + } + + pub fn submit_udp_bind( + self: &Arc, + local_port: u16, + timeout: Option, + ) -> DataPlaneResult { + Self::ensure_executor()?; + let deadline = DataPlaneDeadline::from_optional_timeout(timeout); + let runtime = self.runtime()?; + let (operation_id, cancel) = { + let mut state = self.lock_state(); + self.admit_locked(&mut state, DataPlaneOperationKind::UdpBind, None, 0, true)? + }; + self.spawn_operation(operation_id, async move { + let socket = Self::run_operation( + cancel, + deadline, + runtime.data_plane_udp_bind_with_deadline(local_port, deadline), + ) + .await?; + Ok(PendingOperationResult::UdpBound(socket)) + }); + Ok(operation_id) + } + + pub fn submit_udp_receive( + self: &Arc, + socket_id: DataPlaneResourceId, + max_len: usize, + timeout: Option, + ) -> DataPlaneResult { + Self::ensure_executor()?; + self.require_read_size(max_len)?; + let deadline = DataPlaneDeadline::from_optional_timeout(timeout); + let (socket, operation_id, cancel) = { + let mut state = self.lock_state(); + let socket = Self::require_udp(&state, socket_id)?; + let (operation_id, cancel) = self.admit_locked( + &mut state, + DataPlaneOperationKind::UdpReceive, + Some(socket_id), + max_len, + false, + )?; + (socket, operation_id, cancel) + }; + self.spawn_operation(operation_id, async move { + let (data, peer_addr, truncated) = Self::run_operation(cancel, deadline, async move { + let _read = socket.read.lock().await; + socket.socket.recv_from_limited(max_len).await + }) + .await?; + Ok(PendingOperationResult::UdpReceived { + data, + peer_addr, + truncated, + }) + }); + Ok(operation_id) + } + + pub fn submit_udp_send( + self: &Arc, + socket_id: DataPlaneResourceId, + peer_addr: SocketAddr, + data: Vec, + timeout: Option, + ) -> DataPlaneResult { + Self::ensure_executor()?; + let deadline = DataPlaneDeadline::from_optional_timeout(timeout); + let (socket, operation_id, cancel) = { + let mut state = self.lock_state(); + let socket = Self::require_udp(&state, socket_id)?; + let (operation_id, cancel) = self.admit_locked( + &mut state, + DataPlaneOperationKind::UdpSend, + Some(socket_id), + 0, + false, + )?; + (socket, operation_id, cancel) + }; + self.spawn_operation(operation_id, async move { + let len = Self::run_operation(cancel, deadline, async move { + let _write = socket.write.lock().await; + socket.socket.send_to(&data, peer_addr).await + }) + .await?; + Ok(PendingOperationResult::UdpSent(len)) + }); + Ok(operation_id) + } + + fn unlink_target_locked( + resources: &mut ResourceTable, + operation_id: DataPlaneOperationId, + target: Option, + ) { + if let Some(resource_id) = target + && let Some(resource) = resources.entries.get_mut(&resource_id) + { + resource.pending_operations.remove(&operation_id); + } + } + + fn release_result_reservation_locked( + retained_result_bytes: &mut usize, + metadata: &mut OperationMetadata, + ) { + *retained_result_bytes = + retained_result_bytes.saturating_sub(metadata.reserved_result_bytes); + metadata.reserved_result_bytes = 0; + } + + fn release_resource_reservation_locked( + resources: &mut ResourceTable, + metadata: &mut OperationMetadata, + ) { + if metadata.reserves_resource { + resources.reserved = resources.reserved.saturating_sub(1); + metadata.reserves_resource = false; + } + } + + fn insert_tcp_resource_locked( + resources: &mut ResourceTable, + stream: DataPlaneTcpStream, + ) -> DataPlaneResult { + let resource_id = Self::next_resource_id(resources)?; + let (read, write) = tokio::io::split(stream); + resources.entries.insert( + resource_id, + ResourceEntry { + io: ResourceIo::Tcp(Arc::new(TcpResource { + read: AsyncMutex::new(read), + write: AsyncMutex::new(write), + })), + pending_operations: HashSet::new(), + }, + ); + Ok(resource_id) + } + + fn insert_listener_resource_locked( + resources: &mut ResourceTable, + listener: DataPlaneTcpListener, + ) -> DataPlaneResult { + let resource_id = Self::next_resource_id(resources)?; + resources.entries.insert( + resource_id, + ResourceEntry { + io: ResourceIo::TcpListener(Arc::new(AsyncMutex::new(listener))), + pending_operations: HashSet::new(), + }, + ); + Ok(resource_id) + } + + fn insert_udp_resource_locked( + resources: &mut ResourceTable, + socket: DataPlaneUdpSocket, + ) -> DataPlaneResult { + let resource_id = Self::next_resource_id(resources)?; + resources.entries.insert( + resource_id, + ResourceEntry { + io: ResourceIo::Udp(Arc::new(UdpResource { + socket: Arc::new(socket), + read: AsyncMutex::new(()), + write: AsyncMutex::new(()), + })), + pending_operations: HashSet::new(), + }, + ); + Ok(resource_id) + } + + fn finalize_success_locked( + resources: &mut ResourceTable, + retained_result_bytes: &mut usize, + metadata: &mut OperationMetadata, + result: PendingOperationResult, + ) -> DataPlaneResult { + let result = match result { + PendingOperationResult::TcpConnected { stream, peer_addr } => { + let local_addr = stream.local_addr(); + let stream = Self::insert_tcp_resource_locked(resources, stream)?; + DataPlaneOperationResult::TcpConnected { + stream, + local_addr, + peer_addr, + } + } + PendingOperationResult::TcpBound(listener) => { + let local_addr = listener.local_addr(); + let listener = Self::insert_listener_resource_locked(resources, listener)?; + DataPlaneOperationResult::TcpBound { + listener, + local_addr, + } + } + PendingOperationResult::TcpAccepted { stream, peer_addr } => { + let local_addr = stream.local_addr(); + let stream = Self::insert_tcp_resource_locked(resources, stream)?; + DataPlaneOperationResult::TcpAccepted { + stream, + local_addr, + peer_addr, + } + } + PendingOperationResult::TcpRead { data, eof } => { + DataPlaneOperationResult::TcpRead { data, eof } + } + PendingOperationResult::TcpWritten(len) => DataPlaneOperationResult::TcpWritten { len }, + PendingOperationResult::UdpBound(socket) => { + let local_addr = socket.local_addr(); + let socket = Self::insert_udp_resource_locked(resources, socket)?; + DataPlaneOperationResult::UdpBound { socket, local_addr } + } + PendingOperationResult::UdpReceived { + data, + peer_addr, + truncated, + } => DataPlaneOperationResult::UdpReceived { + data, + peer_addr, + truncated, + }, + PendingOperationResult::UdpSent(len) => DataPlaneOperationResult::UdpSent { len }, + }; + + let retained_bytes = result.retained_bytes(); + if retained_bytes > metadata.reserved_result_bytes { + return Err(Self::error( + DataPlaneErrorKind::ResourceLimit, + "data-plane operation exceeded its reserved result bytes", + )); + } + *retained_result_bytes -= metadata.reserved_result_bytes - retained_bytes; + metadata.reserved_result_bytes = retained_bytes; + Self::release_resource_reservation_locked(resources, metadata); + Ok(result) + } + + fn complete_operation( + &self, + operation_id: DataPlaneOperationId, + result: DataPlaneResult, + ) { + let mut state = self.lock_state(); + let SessionState { + resources, + broker, + retained_result_bytes, + .. + } = &mut *state; + let notify = broker.complete_with(operation_id.broker_id(), |kind, metadata| { + Self::unlink_target_locked(resources, operation_id, metadata.target); + match result { + Ok(result) => match Self::finalize_success_locked( + resources, + retained_result_bytes, + metadata, + result, + ) { + Ok(result) => Ok(result), + Err(error) => { + Self::release_result_reservation_locked(retained_result_bytes, metadata); + Self::release_resource_reservation_locked(resources, metadata); + Err(error.kind()) + } + }, + Err(error) => { + tracing::debug!( + ?operation_id, + ?kind, + error = %error, + "data-plane operation failed" + ); + Self::release_result_reservation_locked(retained_result_bytes, metadata); + Self::release_resource_reservation_locked(resources, metadata); + Err(error.kind()) + } + } + }); + drop(state); + if notify { + self.notify_completion(); + } + } + + fn queue_error_locked( + state: &mut SessionState, + operation_id: DataPlaneOperationId, + error_kind: DataPlaneErrorKind, + ) -> bool { + let SessionState { + resources, + broker, + retained_result_bytes, + .. + } = state; + broker.cancel_with(operation_id.broker_id(), |_, metadata| { + Self::unlink_target_locked(resources, operation_id, metadata.target); + Self::release_result_reservation_locked(retained_result_bytes, metadata); + Self::release_resource_reservation_locked(resources, metadata); + Err(error_kind) + }) + } + + fn close_resource_locked( + state: &mut SessionState, + resource_id: DataPlaneResourceId, + pending_kind: DataPlaneErrorKind, + ) -> bool { + let Some(resource) = state.resources.entries.remove(&resource_id) else { + return false; + }; + let mut notify = false; + for operation_id in resource.pending_operations { + notify |= Self::queue_error_locked(state, operation_id, pending_kind); + } + notify + } + + fn discard_outcome_locked( + state: &mut SessionState, + outcome: DataPlaneOperationOutcome, + ) -> bool { + match outcome { + Ok(result) => result.created_resource().is_some_and(|resource_id| { + Self::close_resource_locked(state, resource_id, DataPlaneErrorKind::HandleClosed) + }), + Err(_) => false, + } + } + + fn notify_completion(&self) { + self.completion_condvar.notify_all(); + self.completion_notify.notify_one(); + } + + pub fn cancel_operation(&self, operation_id: DataPlaneOperationId) { + let mut state = self.lock_state(); + let notify = + Self::queue_error_locked(&mut state, operation_id, DataPlaneErrorKind::Cancelled); + drop(state); + if notify { + self.notify_completion(); + } + } + + pub fn free_operation(&self, operation_id: DataPlaneOperationId) { + let mut state = self.lock_state(); + let Some(mut released) = state.broker.free(operation_id.broker_id()) else { + return; + }; + Self::unlink_target_locked(&mut state.resources, operation_id, released.metadata.target); + Self::release_result_reservation_locked( + &mut state.retained_result_bytes, + &mut released.metadata, + ); + Self::release_resource_reservation_locked(&mut state.resources, &mut released.metadata); + let notify = released + .outcome + .is_some_and(|outcome| Self::discard_outcome_locked(&mut state, outcome)); + drop(state); + if notify { + self.notify_completion(); + } + } + + pub fn close_resource(&self, resource_id: DataPlaneResourceId) { + let mut state = self.lock_state(); + let notify = + Self::close_resource_locked(&mut state, resource_id, DataPlaneErrorKind::HandleClosed); + drop(state); + if notify { + self.notify_completion(); + } + } + + pub fn drain_completions(&self, max_count: usize) -> Vec { + let mut state = self.lock_state(); + state + .broker + .drain(max_count, |outcome| match outcome { + Ok(_) => DataPlaneCompletionStatus::Success, + Err(kind) => DataPlaneCompletionStatus::Error(*kind), + }) + .into_iter() + .map(|completion| DataPlaneCompletionDescriptor { + operation_id: DataPlaneOperationId::from_broker(completion.operation_id), + kind: completion.kind, + status: completion.status, + }) + .collect() + } + + pub fn has_completions(&self) -> bool { + self.lock_state().broker.has_completions() + } + + pub fn completion_wait(&self, timeout: Option) -> bool { + self.completion_wait_after_ready(timeout, || {}) + } + + fn completion_wait_after_ready(&self, timeout: Option, ready: impl FnOnce()) -> bool { + let state = self.lock_state(); + if state.broker.has_completions() { + return true; + } + if state.lifecycle == SessionLifecycle::Stopped { + return false; + } + let wake_generation = state.broker.wake_generation(); + ready(); + + let state = match timeout { + Some(timeout) => { + self.completion_condvar + .wait_timeout_while(state, timeout, |state| { + !state.broker.has_completions() + && state.lifecycle != SessionLifecycle::Stopped + && state.broker.wake_generation() == wake_generation + }) + .unwrap_or_else(|error| error.into_inner()) + .0 + } + None => self + .completion_condvar + .wait_while(state, |state| { + !state.broker.has_completions() + && state.lifecycle != SessionLifecycle::Stopped + && state.broker.wake_generation() == wake_generation + }) + .unwrap_or_else(|error| error.into_inner()), + }; + state.broker.has_completions() + } + + pub fn discard_all(&self) { + let mut state = self.lock_state(); + state.broker.discard_all(); + state.resources.entries.clear(); + state.resources.reserved = 0; + state.retained_result_bytes = 0; + drop(state); + self.completion_condvar.notify_all(); + self.completion_notify.notify_one(); + } + + pub async fn completion_notified(&self) { + loop { + if self.has_completions() { + return; + } + { + let state = self.lock_state(); + if state.lifecycle == SessionLifecycle::Stopped { + return; + } + } + self.completion_notify.notified().await; + } + } + + pub fn result_retained_bytes( + &self, + operation_id: DataPlaneOperationId, + ) -> DataPlaneResult { + let state = self.lock_state(); + state + .broker + .with_drained(operation_id.broker_id(), |_, metadata, _| { + metadata.reserved_result_bytes + }) + .map_err(Self::operation_access_error) + } + + pub fn result_payload_bytes( + &self, + operation_id: DataPlaneOperationId, + ) -> DataPlaneResult { + let state = self.lock_state(); + state + .broker + .with_drained(operation_id.broker_id(), |_, _, outcome| match outcome { + Ok(result) => result.payload_bytes(), + Err(_) => 0, + }) + .map_err(Self::operation_access_error) + } + + pub fn operation_kind( + &self, + operation_id: DataPlaneOperationId, + ) -> DataPlaneResult { + let state = self.lock_state(); + state + .broker + .with_drained(operation_id.broker_id(), |kind, _, _| kind) + .map_err(Self::operation_access_error) + } + + pub fn take_result_with( + &self, + operation_id: DataPlaneOperationId, + take: impl FnOnce(&DataPlaneOperationOutcome) -> Option, + ) -> DataPlaneResult> { + let mut state = self.lock_state(); + let taken = state + .broker + .take_with(operation_id.broker_id(), take) + .map_err(Self::operation_access_error)?; + let Some(mut taken) = taken else { + return Ok(None); + }; + Self::release_result_reservation_locked( + &mut state.retained_result_bytes, + &mut taken.metadata, + ); + Ok(Some(taken.value)) + } + + fn operation_access_error(error: OperationAccessError) -> DataPlaneError { + match error { + OperationAccessError::Missing => Self::error( + DataPlaneErrorKind::HandleClosed, + "data-plane operation result is unavailable", + ), + OperationAccessError::NotDrained => Self::error( + DataPlaneErrorKind::PathNotReady, + "data-plane operation result has not been drained", + ), + } + } + + pub(crate) fn stop(&self) { + let mut state = self.lock_state(); + if state.lifecycle == SessionLifecycle::Stopped { + return; + } + state.lifecycle = SessionLifecycle::Stopped; + state.broker.invalidate_waiters(); + let pending = state.broker.pending_ids(); + let mut notify = false; + for operation_id in pending { + notify |= Self::queue_error_locked( + &mut state, + DataPlaneOperationId::from_broker(operation_id), + DataPlaneErrorKind::InstanceStopped, + ); + } + state.resources.entries.clear(); + drop(state); + self.consumer_lease + .lock() + .unwrap_or_else(|error| error.into_inner()) + .take(); + self.completion_condvar.notify_all(); + self.completion_notify.notify_one(); + if notify { + tracing::trace!("data-plane session queued stop completions"); + } + } + + #[cfg(test)] + fn admit_test_operation( + &self, + kind: DataPlaneOperationKind, + reserved_result_bytes: usize, + ) -> DataPlaneResult { + let mut state = self.lock_state(); + self.admit_locked(&mut state, kind, None, reserved_result_bytes, false) + .map(|(operation_id, _)| operation_id) + } + + #[cfg(test)] + fn admit_test_resource_operation( + &self, + kind: DataPlaneOperationKind, + ) -> DataPlaneResult { + let mut state = self.lock_state(); + self.admit_locked(&mut state, kind, None, 0, true) + .map(|(operation_id, _)| operation_id) + } + + #[cfg(test)] + fn complete_test_operation( + &self, + operation_id: DataPlaneOperationId, + result: DataPlaneResult, + ) { + self.complete_operation(operation_id, result); + } +} + +impl Drop for DataPlaneSession +where + H: VirtualTcpSocketFactory + VirtualTcpListenerFactory + VirtualUdpSocketFactory, +{ + fn drop(&mut self) { + let state = self + .state + .get_mut() + .unwrap_or_else(|error| error.into_inner()); + state.broker.discard_all(); + state.resources.entries.clear(); + } +} + +#[cfg(test)] +mod tests { + use std::{ + sync::{Arc, Barrier}, + thread, + }; + + use crate::host::testkit::TestHost; + + use super::*; + + fn session_with_limits(limits: DataPlaneSessionLimits) -> Arc> { + let session = DataPlaneSession::with_runtime(Weak::new(), limits); + session.start().unwrap(); + session + } + + fn session() -> Arc> { + session_with_limits(DataPlaneSessionLimits::default()) + } + + fn successful_write(len: usize) -> DataPlaneResult { + Ok(PendingOperationResult::TcpWritten(len)) + } + + #[test] + fn completion_is_drained_once_and_result_is_taken_once() { + let session = session(); + let operation_id = session + .admit_test_operation(DataPlaneOperationKind::TcpWrite, 0) + .unwrap(); + session.complete_test_operation(operation_id, successful_write(7)); + + let completions = session.drain_completions(8); + assert_eq!(completions.len(), 1); + assert_eq!(completions[0].operation_id, operation_id); + assert_eq!(completions[0].status, DataPlaneCompletionStatus::Success); + assert!(session.drain_completions(8).is_empty()); + + let preserved = session + .take_result_with(operation_id, |_| None::) + .unwrap(); + assert_eq!(preserved, None); + let len = session + .take_result_with(operation_id, |outcome| match outcome { + Ok(DataPlaneOperationResult::TcpWritten { len }) => Some(*len), + _ => None, + }) + .unwrap(); + assert_eq!(len, Some(7)); + assert_eq!( + session + .take_result_with(operation_id, |_| Some(())) + .unwrap_err() + .kind(), + DataPlaneErrorKind::HandleClosed + ); + } + + #[test] + fn free_pending_discards_the_late_completion() { + let session = session(); + let operation_id = session + .admit_test_operation(DataPlaneOperationKind::TcpWrite, 0) + .unwrap(); + + session.free_operation(operation_id); + session.complete_test_operation(operation_id, successful_write(3)); + + assert!(session.drain_completions(8).is_empty()); + assert_eq!(session.lock_state().broker.len(), 0); + } + + #[test] + fn free_queued_removes_its_completion_descriptor() { + let session = session(); + let operation_id = session + .admit_test_operation(DataPlaneOperationKind::TcpWrite, 0) + .unwrap(); + session.complete_test_operation(operation_id, successful_write(3)); + + session.free_operation(operation_id); + + assert!(session.drain_completions(8).is_empty()); + assert_eq!(session.lock_state().broker.len(), 0); + } + + #[test] + fn cancel_and_complete_race_has_one_terminal_outcome() { + for _ in 0..128 { + let session = session(); + let operation_id = session + .admit_test_operation(DataPlaneOperationKind::TcpWrite, 0) + .unwrap(); + let barrier = Arc::new(Barrier::new(3)); + + let cancel_session = session.clone(); + let cancel_barrier = barrier.clone(); + let cancel = thread::spawn(move || { + cancel_barrier.wait(); + cancel_session.cancel_operation(operation_id); + }); + let complete_session = session.clone(); + let complete_barrier = barrier.clone(); + let complete = thread::spawn(move || { + complete_barrier.wait(); + complete_session.complete_test_operation(operation_id, successful_write(9)); + }); + barrier.wait(); + cancel.join().unwrap(); + complete.join().unwrap(); + + let completions = session.drain_completions(8); + assert_eq!(completions.len(), 1); + assert!(matches!( + completions[0].status, + DataPlaneCompletionStatus::Success + | DataPlaneCompletionStatus::Error(DataPlaneErrorKind::Cancelled) + )); + } + } + + #[test] + fn blocking_wait_observes_completion_before_and_after_wait_starts() { + let session = session(); + let first = session + .admit_test_operation(DataPlaneOperationKind::TcpWrite, 0) + .unwrap(); + session.complete_test_operation(first, successful_write(1)); + assert!(session.completion_wait(Some(Duration::ZERO))); + session.drain_completions(1); + + let second = session + .admit_test_operation(DataPlaneOperationKind::TcpWrite, 0) + .unwrap(); + let waiter = { + let session = session.clone(); + thread::spawn(move || session.completion_wait(Some(Duration::from_secs(1)))) + }; + thread::yield_now(); + session.complete_test_operation(second, successful_write(2)); + assert!(waiter.join().unwrap()); + } + + #[test] + fn stop_wakes_waiters_and_preserves_terminal_completion() { + let session = session(); + let operation_id = session + .admit_test_operation(DataPlaneOperationKind::TcpWrite, 0) + .unwrap(); + let waiter = { + let session = session.clone(); + thread::spawn(move || session.completion_wait(None)) + }; + + session.stop(); + + assert!(waiter.join().unwrap()); + let completion = session.drain_completions(1).pop().unwrap(); + assert_eq!(completion.operation_id, operation_id); + assert_eq!( + completion.status, + DataPlaneCompletionStatus::Error(DataPlaneErrorKind::InstanceStopped) + ); + } + + #[test] + fn discard_all_wakes_waiters_without_a_completion() { + let session = session(); + let barrier = Arc::new(Barrier::new(2)); + let waiter = { + let session = session.clone(); + let barrier = barrier.clone(); + thread::spawn(move || { + session.completion_wait_after_ready(None, || { + barrier.wait(); + }) + }) + }; + barrier.wait(); + + session.discard_all(); + + assert!(!waiter.join().unwrap()); + assert!(session.drain_completions(1).is_empty()); + } + + #[test] + fn admission_enforces_operation_and_result_limits() { + let session = session_with_limits(DataPlaneSessionLimits { + max_resources: 1, + max_operations: 1, + max_result_bytes: 4, + max_read_size: 4, + }); + session + .admit_test_operation(DataPlaneOperationKind::TcpRead, 4) + .unwrap(); + + assert_eq!( + session + .admit_test_operation(DataPlaneOperationKind::TcpWrite, 0) + .unwrap_err() + .kind(), + DataPlaneErrorKind::ResourceLimit + ); + assert_eq!( + session.require_read_size(5).unwrap_err().kind(), + DataPlaneErrorKind::ResourceLimit + ); + } + + #[test] + fn terminal_errors_release_result_and_resource_reservations() { + let session = session_with_limits(DataPlaneSessionLimits { + max_resources: 1, + max_operations: 4, + max_result_bytes: 4, + max_read_size: 4, + }); + let read = session + .admit_test_operation(DataPlaneOperationKind::TcpRead, 4) + .unwrap(); + let resource = session + .admit_test_resource_operation(DataPlaneOperationKind::TcpConnect) + .unwrap(); + + assert_eq!( + session + .admit_test_operation(DataPlaneOperationKind::TcpRead, 1) + .unwrap_err() + .kind(), + DataPlaneErrorKind::ResourceLimit + ); + assert_eq!( + session + .admit_test_resource_operation(DataPlaneOperationKind::TcpBind) + .unwrap_err() + .kind(), + DataPlaneErrorKind::ResourceLimit + ); + + session.cancel_operation(read); + session.complete_test_operation( + resource, + Err(DataPlaneError::new( + DataPlaneErrorKind::Io, + "synthetic failure", + )), + ); + + session + .admit_test_operation(DataPlaneOperationKind::TcpRead, 4) + .unwrap(); + session + .admit_test_resource_operation(DataPlaneOperationKind::TcpBind) + .unwrap(); + } + + #[test] + fn short_read_accounts_for_retained_vector_capacity() { + let session = session_with_limits(DataPlaneSessionLimits { + max_resources: 1, + max_operations: 2, + max_result_bytes: 4, + max_read_size: 4, + }); + let operation_id = session + .admit_test_operation(DataPlaneOperationKind::TcpRead, 4) + .unwrap(); + let mut data = Vec::with_capacity(4); + data.push(7); + + session.complete_test_operation( + operation_id, + Ok(PendingOperationResult::TcpRead { data, eof: false }), + ); + + assert_eq!(session.lock_state().retained_result_bytes, 4); + session.drain_completions(1); + assert_eq!(session.result_retained_bytes(operation_id).unwrap(), 4); + assert_eq!(session.result_payload_bytes(operation_id).unwrap(), 1); + session + .take_result_with(operation_id, |_| Some(())) + .unwrap(); + assert_eq!(session.lock_state().retained_result_bytes, 0); + } + + #[tokio::test] + async fn spawned_operation_does_not_keep_session_alive() { + let session = session(); + let (operation_id, cancel) = { + let mut state = session.lock_state(); + session + .admit_locked(&mut state, DataPlaneOperationKind::TcpWrite, None, 0, false) + .unwrap() + }; + let (cancelled_tx, cancelled_rx) = tokio::sync::oneshot::channel(); + session.spawn_operation(operation_id, async move { + cancel.cancelled().await; + let _ = cancelled_tx.send(()); + Err(DataPlaneError::new( + DataPlaneErrorKind::Cancelled, + "synthetic cancellation", + )) + }); + let weak = Arc::downgrade(&session); + + drop(session); + + tokio::time::timeout(Duration::from_secs(1), cancelled_rx) + .await + .unwrap() + .unwrap(); + assert!(weak.upgrade().is_none()); + } + + #[tokio::test] + async fn deadline_elapsed_before_polling_is_expired() { + let deadline = DataPlaneDeadline::from_timeout(Duration::from_millis(1)); + thread::sleep(Duration::from_millis(5)); + + assert_eq!( + deadline + .run(async { Ok::<_, DataPlaneError>(()) }) + .await + .unwrap_err() + .kind(), + DataPlaneErrorKind::DeadlineExceeded + ); + } +} diff --git a/easytier-core/src/gateway/dataplane/stack.rs b/easytier-core/src/gateway/dataplane/stack.rs new file mode 100644 index 00000000..bc78fe83 --- /dev/null +++ b/easytier-core/src/gateway/dataplane/stack.rs @@ -0,0 +1,130 @@ +//! smoltcp stack generation and its peer-packet bridge. + +use std::{ + net::IpAddr, + sync::{Arc, Weak}, +}; + +use pnet_packet::ipv4::Ipv4Packet; +use tokio::{ + sync::{Mutex, mpsc}, + task::JoinSet, +}; + +use crate::{ + foundation::task::reap_joinset_background, + gateway::smoltcp::{BufferSize, Net, NetConfig, channel_device}, + packet::ZCPacket, + peers::peer_manager::PeerManagerCore, +}; + +use super::{DataPlaneErrorKind, DataPlaneIoGuard}; + +pub(super) struct SmoltcpPlane { + pub(super) ipv4_addr: cidr::Ipv4Inet, + pub(super) net: Arc, + generation: DataPlaneIoGuard, + _forward_tasks: Arc>>, +} + +impl SmoltcpPlane { + pub(super) fn new( + ipv4_addr: cidr::Ipv4Inet, + peer_manager: Weak, + packet_recv: Arc>>, + ) -> Self { + let mut forward_tasks = JoinSet::new(); + let mut capabilities = smoltcp::phy::DeviceCapabilities::default(); + // Fragment offsets are expressed in eight-byte units. + capabilities.max_transmission_unit = 1284; + capabilities.medium = smoltcp::phy::Medium::Ip; + let (device, stack_sink, mut stack_stream) = + channel_device::ChannelDevice::new(capabilities); + + forward_tasks.spawn(async move { + let mut packet_recv = packet_recv.lock().await; + while let Some(packet) = packet_recv.recv().await { + tracing::trace!(?packet, "deliver peer packet to smoltcp"); + if let Err(error) = stack_sink.send(Ok(packet.payload().to_vec())).await { + tracing::error!(?error, "deliver peer packet to smoltcp failed"); + } + } + tracing::debug!("peer-to-smoltcp bridge stopped"); + }); + + forward_tasks.spawn(async move { + while let Some(data) = stack_stream.recv().await { + let Some(ipv4) = Ipv4Packet::new(&data) else { + tracing::error!(?data, "smoltcp emitted a non-IPv4 packet"); + continue; + }; + let destination = ipv4.get_destination(); + let Some(peer_manager) = peer_manager.upgrade() else { + tracing::debug!("smoltcp-to-peer bridge lost PeerManager"); + return; + }; + if let Err(error) = peer_manager + .send_msg_by_ip( + ZCPacket::new_with_payload(&data), + IpAddr::V4(destination), + false, + ) + .await + { + tracing::error!(?error, "deliver smoltcp packet to peer failed"); + } + } + tracing::debug!("smoltcp-to-peer bridge stopped"); + }); + + let interface_config = smoltcp::iface::Config::new(smoltcp::wire::HardwareAddress::Ip); + let net = Net::new( + device, + NetConfig::new( + interface_config, + format!("{}/{}", ipv4_addr.address(), ipv4_addr.network_length()) + .parse() + .expect("validated IPv4 prefix"), + vec![ + ipv4_addr + .address() + .to_string() + .parse() + .expect("validated IPv4 address"), + ], + Some(BufferSize { + tcp_rx_size: 1024 * 128, + tcp_tx_size: 1024 * 128, + ..Default::default() + }), + ), + ); + + let forward_tasks = Arc::new(std::sync::Mutex::new(forward_tasks)); + forward_tasks.lock().unwrap().spawn(reap_joinset_background( + forward_tasks.clone(), + "SmoltcpPlane", + )); + + Self { + ipv4_addr, + net: Arc::new(net), + generation: DataPlaneIoGuard::new(), + _forward_tasks: forward_tasks, + } + } + + pub(super) fn lease(&self) -> DataPlaneIoGuard { + self.generation.clone() + } + + pub(super) fn close(&self, kind: DataPlaneErrorKind) { + self.generation.close(kind); + } +} + +impl Drop for SmoltcpPlane { + fn drop(&mut self) { + self.close(DataPlaneErrorKind::HandleClosed); + } +} diff --git a/easytier-core/src/gateway/dataplane/tcp.rs b/easytier-core/src/gateway/dataplane/tcp.rs new file mode 100644 index 00000000..dea9a521 --- /dev/null +++ b/easytier-core/src/gateway/dataplane/tcp.rs @@ -0,0 +1,160 @@ +//! TCP stream and listener resources exposed by the data plane. + +use std::{ + future::Future, + net::SocketAddr, + pin::Pin, + task::{Context, Poll}, +}; + +use tokio::io::{AsyncRead, AsyncWrite}; + +use crate::gateway::smoltcp::TcpListener; + +use super::{ + DataPlaneIoGuard, DataPlaneLease, DataPlaneTcpIo, FlowData, FlowKey, FlowLease, FlowSet, + TCP_ENTRY, +}; + +/// Tracks how an established stream keeps its inbound flow alive. +pub(super) enum DataPlaneTcpStreamRoute { + Outbound { _flow: FlowLease }, + Accepted { _flow: FlowLease }, + External, +} + +/// A TCP stream created by the data-plane API. +pub struct DataPlaneTcpStream { + pub(super) stream: DataPlaneTcpIo, + pub(super) local_addr: SocketAddr, + pub(super) _route: DataPlaneTcpStreamRoute, + pub(super) _data_plane_lease: Option, + generation: DataPlaneIoGuard, + read_closed: Pin + Send>>, + write_closed: Pin + Send>>, +} + +/// A TCP listener created by the data-plane API. +pub struct DataPlaneTcpListener { + pub(super) listener: TcpListener, + pub(super) local_addr: SocketAddr, + pub(super) flows: FlowSet, + pub(super) _listen_flow: FlowLease, + pub(super) data_plane_lease: DataPlaneLease, + pub(super) generation: DataPlaneIoGuard, +} + +impl DataPlaneTcpStream { + pub(super) fn new( + stream: DataPlaneTcpIo, + local_addr: SocketAddr, + data_plane_lease: Option, + route: DataPlaneTcpStreamRoute, + generation: DataPlaneIoGuard, + ) -> Self { + Self { + stream, + local_addr, + _data_plane_lease: data_plane_lease, + _route: route, + read_closed: generation.closed_future(), + write_closed: generation.closed_future(), + generation, + } + } + + pub fn local_addr(&self) -> SocketAddr { + self.local_addr + } +} + +impl DataPlaneTcpListener { + pub fn local_addr(&self) -> SocketAddr { + self.local_addr + } + + pub async fn accept(&mut self) -> Result<(DataPlaneTcpStream, SocketAddr), std::io::Error> { + self.generation + .ensure_open() + .map_err(|error| error.into_io_error())?; + let generation = self.generation.clone(); + let (stream, peer_addr) = tokio::select! { + biased; + _ = generation.closed() => { + return Err(generation.closed_io_error()); + } + result = self.listener.accept() => result?, + }; + generation + .ensure_open() + .map_err(|error| error.into_io_error())?; + let local_addr = stream.local_addr()?; + let (flow, _) = FlowLease::register( + self.flows.clone(), + FlowKey { + src: local_addr, + dst: peer_addr, + kind: TCP_ENTRY, + }, + FlowData::DataPlaneRoute, + ); + let accepted = DataPlaneTcpStream::new( + Box::new(stream), + local_addr, + Some(self.data_plane_lease.clone()), + DataPlaneTcpStreamRoute::Accepted { _flow: flow }, + generation, + ); + Ok((accepted, peer_addr)) + } +} + +impl AsyncRead for DataPlaneTcpStream { + fn poll_read( + self: Pin<&mut Self>, + cx: &mut Context<'_>, + buf: &mut tokio::io::ReadBuf<'_>, + ) -> Poll> { + let this = self.get_mut(); + if this.generation.ensure_open().is_err() || this.read_closed.as_mut().poll(cx).is_ready() { + return Poll::Ready(Err(this.generation.closed_io_error())); + } + Pin::new(&mut this.stream).poll_read(cx, buf) + } +} + +impl AsyncWrite for DataPlaneTcpStream { + fn poll_write( + self: Pin<&mut Self>, + cx: &mut Context<'_>, + buf: &[u8], + ) -> Poll> { + let this = self.get_mut(); + if this.generation.ensure_open().is_err() || this.write_closed.as_mut().poll(cx).is_ready() + { + return Poll::Ready(Err(this.generation.closed_io_error())); + } + Pin::new(&mut this.stream).poll_write(cx, buf) + } + + fn poll_flush(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + let this = self.get_mut(); + if this.generation.ensure_open().is_err() || this.write_closed.as_mut().poll(cx).is_ready() + { + return Poll::Ready(Err(this.generation.closed_io_error())); + } + Pin::new(&mut this.stream).poll_flush(cx) + } + + fn poll_shutdown( + self: Pin<&mut Self>, + cx: &mut Context<'_>, + ) -> Poll> { + let this = self.get_mut(); + if this.generation.ensure_open().is_err() || this.write_closed.as_mut().poll(cx).is_ready() + { + return Poll::Ready(Err(this.generation.closed_io_error())); + } + Pin::new(&mut this.stream).poll_shutdown(cx) + } +} diff --git a/easytier-core/src/gateway/dataplane/tests.rs b/easytier-core/src/gateway/dataplane/tests.rs new file mode 100644 index 00000000..b706d55e --- /dev/null +++ b/easytier-core/src/gateway/dataplane/tests.rs @@ -0,0 +1,839 @@ +use std::net::{IpAddr, Ipv4Addr, SocketAddr}; + +use pnet_packet::{ + MutablePacket, + ip::IpNextHeaderProtocols, + ipv4::{self, MutableIpv4Packet}, + tcp::{self, MutableTcpPacket, TcpFlags}, +}; +use tokio::io::{AsyncReadExt, AsyncWriteExt}; + +use super::*; +use crate::{ + config::peers::PeerRuntimeSnapshot, + config::{IpPrefix, NetworkIdentity}, + host::testkit::TestHost, + peers::{ + PacketRecvChanReceiver, create_packet_recv_chan, peer_manager::PortablePeerManagerConfig, + }, + tunnel::ring::RingTunnelRegistry, +}; + +fn test_gateway() -> Arc> { + let runtime_config = CoreRuntimeConfigStore::new( + crate::config::runtime::CoreRuntimeConfig::default(), + Arc::new(PeerRuntimeSnapshot::default()), + ); + let host = Arc::new(TestHost::default()); + let (packet_sender, packet_recv) = mpsc::channel(16); + Arc::new(DataPlaneRuntime { + operation: Mutex::new(()), + runtime_started: AtomicBool::new(false), + runtime_guard: DataPlaneIoGuard::new(), + runtime_config, + peer_manager: Weak::new(), + transport_proxy: None, + host: host.clone(), + socket_context: SocketContext::default(), + runtime_tasks: Arc::new(std::sync::Mutex::new(JoinSet::new())), + packet_sender, + packet_recv: Arc::new(Mutex::new(packet_recv)), + net: Arc::new(Mutex::new(None)), + entries: Arc::new(FlowTable::default()), + data_plane_consumers: Arc::new(DataPlaneConsumers::new()), + data_plane_net_ready: tokio::sync::watch::channel(false).0, + pipeline_guard: Mutex::new(None), + }) +} + +struct DataPlaneEndpoint { + gateway: Arc>, + peer_manager: Arc, + _packet_receiver: PacketRecvChanReceiver, + ip: cidr::Ipv4Inet, +} + +fn data_plane_endpoint(host: Arc, ip: cidr::Ipv4Inet) -> DataPlaneEndpoint { + const NETWORK_NAME: &str = "gateway-data-plane"; + + let mut runtime = PeerRuntimeSnapshot::default().runtime; + runtime.core.node.peer_id = None; + runtime.core.node.network_name = NETWORK_NAME.to_owned(); + runtime.core.routes.ipv4 = Some( + IpPrefix::new(IpAddr::V4(ip.address()), ip.network_length()) + .expect("test IPv4 prefix should be valid"), + ); + runtime.network_identity = NetworkIdentity { + network_name: NETWORK_NAME.to_owned(), + network_secret: Some("shared-secret".to_owned()), + network_secret_digest: None, + }; + let peer_config = PortablePeerManagerConfig::new(runtime); + let runtime_config = CoreRuntimeConfigStore::new( + crate::config::runtime::CoreRuntimeConfig::default(), + Arc::new(peer_config.snapshot.clone()), + ); + let (packet_sender, packet_receiver) = create_packet_recv_chan(); + let peer_manager = Arc::new( + PeerManagerCore::new_portable_for_test(peer_config, packet_sender) + .expect("build portable peer manager"), + ); + let gateway = DataPlaneRuntime::new( + runtime_config, + peer_manager.clone(), + None, + host, + SocketContext::default(), + ); + + DataPlaneEndpoint { + gateway, + peer_manager, + _packet_receiver: packet_receiver, + ip, + } +} + +async fn setup_data_plane_pair() -> (DataPlaneEndpoint, DataPlaneEndpoint) { + let host = Arc::new(TestHost::default()); + let a = data_plane_endpoint(host.clone(), "10.126.126.1/24".parse().unwrap()); + let b = loop { + let b = data_plane_endpoint(host.clone(), "10.126.126.2/24".parse().unwrap()); + if b.peer_manager.my_peer_id() != a.peer_manager.my_peer_id() { + break b; + } + }; + + let (run_a, run_b) = tokio::join!(a.peer_manager.run(), b.peer_manager.run()); + run_a.unwrap(); + run_b.unwrap(); + let (start_a, start_b) = tokio::join!(a.gateway.start_runtime(), b.gateway.start_runtime()); + start_a.unwrap(); + start_b.unwrap(); + + let registry = Arc::new(RingTunnelRegistry::default()); + let listener_id = uuid::Uuid::new_v4(); + let mut listener = registry.bind(listener_id).unwrap(); + let client_tunnel = registry.connect(listener_id).unwrap().into_tunnel(); + let server_tunnel = listener.accept().await.unwrap().into_tunnel(); + let (client, server) = tokio::join!( + b.peer_manager.add_client_tunnel(client_tunnel, true), + a.peer_manager.add_tunnel_as_server(server_tunnel, true), + ); + client.unwrap(); + server.unwrap(); + + tokio::time::timeout(Duration::from_secs(10), async { + loop { + if a.peer_manager + .list_route_snapshots() + .await + .iter() + .any(|route| { + route.peer_id == b.peer_manager.my_peer_id() + && route.ipv4_addr == Some(b.ip.into()) + }) + { + break; + } + tokio::time::sleep(Duration::from_millis(10)).await; + } + }) + .await + .expect("Ring peers did not exchange routes"); + + (a, b) +} + +async fn stop_data_plane_pair(a: &DataPlaneEndpoint, b: &DataPlaneEndpoint) { + tokio::join!(a.gateway.stop_runtime(), b.gateway.stop_runtime()); + tokio::join!( + a.peer_manager.clear_resources(), + b.peer_manager.clear_resources() + ); +} + +async fn wait_for_session_completion( + session: &DataPlaneSession, +) -> DataPlaneCompletionDescriptor { + tokio::time::timeout(Duration::from_secs(10), session.completion_notified()) + .await + .expect("data-plane session completion timed out"); + let completions = session.drain_completions(1); + assert_eq!(completions.len(), 1); + completions[0] +} + +fn build_tcp_packet(src: SocketAddr, dst: SocketAddr) -> Vec { + let mut buf = vec![0u8; 40]; + let src_ip = match src.ip() { + IpAddr::V4(ip) => ip, + IpAddr::V6(_) => panic!("test only supports ipv4"), + }; + let dst_ip = match dst.ip() { + IpAddr::V4(ip) => ip, + IpAddr::V6(_) => panic!("test only supports ipv4"), + }; + + { + let mut ip_packet = MutableIpv4Packet::new(&mut buf).unwrap(); + ip_packet.set_version(4); + ip_packet.set_header_length(5); + ip_packet.set_total_length(40); + ip_packet.set_ttl(64); + ip_packet.set_next_level_protocol(IpNextHeaderProtocols::Tcp); + ip_packet.set_source(src_ip); + ip_packet.set_destination(dst_ip); + + let mut tcp_packet = MutableTcpPacket::new(ip_packet.payload_mut()).unwrap(); + tcp_packet.set_source(src.port()); + tcp_packet.set_destination(dst.port()); + tcp_packet.set_data_offset(5); + tcp_packet.set_flags(TcpFlags::SYN | TcpFlags::ACK); + tcp_packet.set_window(65535); + tcp_packet.set_checksum(tcp::ipv4_checksum( + &tcp_packet.to_immutable(), + &src_ip, + &dst_ip, + )); + + ip_packet.set_checksum(ipv4::checksum(&ip_packet.to_immutable())); + } + + buf +} + +fn build_udp_followup_fragment(src: Ipv4Addr, dst: Ipv4Addr) -> Vec { + let mut buf = vec![0u8; 28]; + { + let mut ip_packet = MutableIpv4Packet::new(&mut buf).unwrap(); + ip_packet.set_version(4); + ip_packet.set_header_length(5); + ip_packet.set_total_length(28); + ip_packet.set_ttl(64); + ip_packet.set_next_level_protocol(IpNextHeaderProtocols::Udp); + ip_packet.set_fragment_offset(1); + ip_packet.set_source(src); + ip_packet.set_destination(dst); + ip_packet + .payload_mut() + .copy_from_slice(&[0xde, 0xad, 0xbe, 0xef, 0xca, 0xfe, 0xba, 0xbe]); + + ip_packet.set_checksum(ipv4::checksum(&ip_packet.to_immutable())); + } + + buf +} + +#[tokio::test] +async fn data_plane_tcp_pingpong() { + let (a, b) = setup_data_plane_pair().await; + let timeout = Duration::from_secs(10); + let mut listener = b.gateway.data_plane_tcp_bind(0, timeout).await.unwrap(); + let listen_addr = SocketAddr::new(b.ip.address().into(), listener.local_addr().port()); + + let accept = tokio::spawn(async move { + let (mut stream, _peer) = listener.accept().await.unwrap(); + let mut buf = [0u8; 4]; + stream.read_exact(&mut buf).await.unwrap(); + assert_eq!(&buf, b"ping"); + stream.write_all(b"pong").await.unwrap(); + stream.flush().await.unwrap(); + }); + + let mut client = a + .gateway + .data_plane_tcp_connect(listen_addr, timeout) + .await + .unwrap(); + client.write_all(b"ping").await.unwrap(); + client.flush().await.unwrap(); + let mut buf = [0u8; 4]; + client.read_exact(&mut buf).await.unwrap(); + assert_eq!(&buf, b"pong"); + accept.await.unwrap(); + + stop_data_plane_pair(&a, &b).await; +} + +#[tokio::test] +async fn data_plane_sessions_complete_tcp_operations_end_to_end() { + let (a, b) = setup_data_plane_pair().await; + let session_a = DataPlaneSession::new(&a.gateway); + let session_b = DataPlaneSession::new(&b.gateway); + session_a.start().unwrap(); + session_b.start().unwrap(); + + let bind = session_b + .submit_tcp_bind(0, Some(Duration::from_secs(10))) + .unwrap(); + let completion = wait_for_session_completion(&session_b).await; + assert_eq!(completion.operation_id, bind); + let (listener, listen_addr) = session_b + .take_result_with(bind, |outcome| match outcome { + Ok(DataPlaneOperationResult::TcpBound { + listener, + local_addr, + }) => Some((*listener, *local_addr)), + _ => None, + }) + .unwrap() + .unwrap(); + + let accept = session_b + .submit_tcp_accept(listener, Some(Duration::from_secs(10))) + .unwrap(); + let connect = session_a + .submit_tcp_connect(listen_addr, Some(Duration::from_secs(10))) + .unwrap(); + let (connect_completion, accept_completion) = tokio::join!( + wait_for_session_completion(&session_a), + wait_for_session_completion(&session_b), + ); + assert_eq!(connect_completion.operation_id, connect); + assert_eq!(accept_completion.operation_id, accept); + let client = session_a + .take_result_with(connect, |outcome| match outcome { + Ok(DataPlaneOperationResult::TcpConnected { stream, .. }) => Some(*stream), + _ => None, + }) + .unwrap() + .unwrap(); + let server = session_b + .take_result_with(accept, |outcome| match outcome { + Ok(DataPlaneOperationResult::TcpAccepted { stream, .. }) => Some(*stream), + _ => None, + }) + .unwrap() + .unwrap(); + + let read = session_b + .submit_tcp_read(server, 16, Some(Duration::from_secs(10))) + .unwrap(); + let write = session_a + .submit_tcp_write(client, b"ping".to_vec(), Some(Duration::from_secs(10))) + .unwrap(); + let (write_completion, read_completion) = tokio::join!( + wait_for_session_completion(&session_a), + wait_for_session_completion(&session_b), + ); + assert_eq!(write_completion.operation_id, write); + assert_eq!(read_completion.operation_id, read); + let written = session_a + .take_result_with(write, |outcome| match outcome { + Ok(DataPlaneOperationResult::TcpWritten { len }) => Some(*len), + _ => None, + }) + .unwrap() + .unwrap(); + let received = session_b + .take_result_with(read, |outcome| match outcome { + Ok(DataPlaneOperationResult::TcpRead { data, eof }) if !eof => Some(data.clone()), + _ => None, + }) + .unwrap() + .unwrap(); + assert_eq!(written, 4); + assert_eq!(received, b"ping"); + + let blocked_read = session_b.submit_tcp_read(server, 16, None).unwrap(); + session_b.close_resource(server); + let close_completion = wait_for_session_completion(&session_b).await; + assert_eq!(close_completion.operation_id, blocked_read); + assert_eq!( + close_completion.status, + DataPlaneCompletionStatus::Error(DataPlaneErrorKind::HandleClosed) + ); + let close_error = session_b + .take_result_with(blocked_read, |outcome| outcome.as_ref().err().copied()) + .unwrap() + .unwrap(); + assert_eq!(close_error, DataPlaneErrorKind::HandleClosed); + + let stopped_read = session_a.submit_tcp_read(client, 16, None).unwrap(); + session_a.stop(); + let stop_completion = wait_for_session_completion(&session_a).await; + assert_eq!(stop_completion.operation_id, stopped_read); + assert_eq!( + stop_completion.status, + DataPlaneCompletionStatus::Error(DataPlaneErrorKind::InstanceStopped) + ); + + session_a.close_resource(client); + session_b.close_resource(listener); + session_b.stop(); + stop_data_plane_pair(&a, &b).await; +} + +#[tokio::test] +async fn data_plane_sessions_report_udp_truncation() { + let (a, b) = setup_data_plane_pair().await; + let session_a = DataPlaneSession::new(&a.gateway); + let session_b = DataPlaneSession::new(&b.gateway); + session_a.start().unwrap(); + session_b.start().unwrap(); + + let bind_a = session_a.submit_udp_bind(0, None).unwrap(); + let bind_b = session_b.submit_udp_bind(0, None).unwrap(); + let (completion_a, completion_b) = tokio::join!( + wait_for_session_completion(&session_a), + wait_for_session_completion(&session_b), + ); + assert_eq!(completion_a.operation_id, bind_a); + assert_eq!(completion_b.operation_id, bind_b); + let (socket_a, addr_a) = session_a + .take_result_with(bind_a, |outcome| match outcome { + Ok(DataPlaneOperationResult::UdpBound { socket, local_addr }) => { + Some((*socket, *local_addr)) + } + _ => None, + }) + .unwrap() + .unwrap(); + let (socket_b, addr_b) = session_b + .take_result_with(bind_b, |outcome| match outcome { + Ok(DataPlaneOperationResult::UdpBound { socket, local_addr }) => { + Some((*socket, *local_addr)) + } + _ => None, + }) + .unwrap() + .unwrap(); + + let warmup = session_b + .submit_udp_send(socket_b, addr_a, b"warmup".to_vec(), None) + .unwrap(); + wait_for_session_completion(&session_b).await; + session_b + .take_result_with(warmup, |outcome| match outcome { + Ok(DataPlaneOperationResult::UdpSent { len }) => Some(*len), + _ => None, + }) + .unwrap() + .unwrap(); + + let receive = session_b + .submit_udp_receive(socket_b, 2, Some(Duration::from_secs(10))) + .unwrap(); + let send = session_a + .submit_udp_send( + socket_a, + addr_b, + b"ping".to_vec(), + Some(Duration::from_secs(10)), + ) + .unwrap(); + let (send_completion, receive_completion) = tokio::join!( + wait_for_session_completion(&session_a), + wait_for_session_completion(&session_b), + ); + assert_eq!(send_completion.operation_id, send); + assert_eq!(receive_completion.operation_id, receive); + let (data, peer_addr, truncated) = session_b + .take_result_with(receive, |outcome| match outcome { + Ok(DataPlaneOperationResult::UdpReceived { + data, + peer_addr, + truncated, + }) => Some((data.clone(), *peer_addr, *truncated)), + _ => None, + }) + .unwrap() + .unwrap(); + assert_eq!(data, b"pi"); + assert_eq!(peer_addr, addr_a); + assert!(truncated); + + session_a.close_resource(socket_a); + session_b.close_resource(socket_b); + session_a.stop(); + session_b.stop(); + stop_data_plane_pair(&a, &b).await; +} + +#[tokio::test] +async fn public_tcp_connect_never_falls_back_to_an_unrelated_host() { + let host = Arc::new(TestHost::default()); + let endpoint = data_plane_endpoint(host, "10.126.131.1/24".parse().unwrap()); + endpoint.peer_manager.run().await.unwrap(); + endpoint.gateway.start_runtime().await.unwrap(); + + let error = match endpoint + .gateway + .data_plane_tcp_connect("192.0.2.10:443".parse().unwrap(), Duration::from_secs(1)) + .await + { + Ok(_) => panic!("public data plane unexpectedly used a Host TCP route"), + Err(error) => error, + }; + assert_eq!(error.kind(), DataPlaneErrorKind::NoOverlayRoute); + assert_eq!(endpoint.gateway.host.tcp_binds.load(Ordering::Relaxed), 0); + + endpoint.gateway.stop_runtime().await; + endpoint.peer_manager.clear_resources().await; +} + +#[tokio::test] +async fn listener_and_accepted_stream_own_independent_flow_lifetimes() { + let (a, b) = setup_data_plane_pair().await; + let timeout = Duration::from_secs(10); + let mut listener = b.gateway.data_plane_tcp_bind(0, timeout).await.unwrap(); + let listen_addr = SocketAddr::new(b.ip.address().into(), listener.local_addr().port()); + + let (accepted, client) = tokio::join!( + listener.accept(), + a.gateway.data_plane_tcp_connect(listen_addr, timeout), + ); + let (mut server, peer_addr) = accepted.unwrap(); + let mut client = client.unwrap(); + + assert_eq!(client.local_addr(), peer_addr); + assert_eq!(a.gateway.entries.count(), 1); + assert_eq!(b.gateway.entries.count(), 2); + + drop(listener); + assert_eq!(b.gateway.entries.count(), 1); + + client.write_all(b"after-listener-drop").await.unwrap(); + client.flush().await.unwrap(); + let mut buf = [0u8; 19]; + server.read_exact(&mut buf).await.unwrap(); + assert_eq!(&buf, b"after-listener-drop"); + + drop(client); + assert_eq!(a.gateway.entries.count(), 0); + drop(server); + assert_eq!(b.gateway.entries.count(), 0); + + stop_data_plane_pair(&a, &b).await; +} + +#[tokio::test] +async fn data_plane_udp_pingpong() { + let (a, b) = setup_data_plane_pair().await; + let timeout = Duration::from_secs(10); + let socket_a = a.gateway.data_plane_udp_bind(0, timeout).await.unwrap(); + let socket_b = b.gateway.data_plane_udp_bind(0, timeout).await.unwrap(); + let addr_a = SocketAddr::new(a.ip.address().into(), socket_a.local_addr().port()); + let addr_b = SocketAddr::new(b.ip.address().into(), socket_b.local_addr().port()); + + socket_b.send_to(b"warmup", addr_a).await.unwrap(); + socket_a.send_to(b"ping", addr_b).await.unwrap(); + let mut buf = [0u8; 16]; + let (len, from) = tokio::time::timeout(timeout, socket_b.recv_from(&mut buf)) + .await + .expect("receive ping timed out") + .unwrap(); + assert_eq!(&buf[..len], b"ping"); + assert_eq!(from, addr_a); + + socket_b.send_to(b"pong", addr_a).await.unwrap(); + loop { + let (len, from) = tokio::time::timeout(timeout, socket_a.recv_from(&mut buf)) + .await + .expect("receive pong timed out") + .unwrap(); + if &buf[..len] == b"pong" { + assert_eq!(from, addr_b); + break; + } + } + + stop_data_plane_pair(&a, &b).await; +} + +#[tokio::test] +async fn udp_socket_drop_releases_every_destination_flow() { + let (a, b) = setup_data_plane_pair().await; + let timeout = Duration::from_secs(10); + let socket = a.gateway.data_plane_udp_bind(0, timeout).await.unwrap(); + let first = SocketAddr::new(b.ip.address().into(), 31001); + let second = SocketAddr::new(b.ip.address().into(), 31002); + + socket.send_to(b"one", first).await.unwrap(); + socket.send_to(b"two", second).await.unwrap(); + assert_eq!(a.gateway.entries.count(), 2); + + drop(socket); + assert_eq!(a.gateway.entries.count(), 0); + + stop_data_plane_pair(&a, &b).await; +} + +#[tokio::test] +async fn udp_socket_owns_a_host_port_reservation() { + let host = Arc::new(TestHost::default()); + let endpoint = data_plane_endpoint(host.clone(), "10.126.132.1/24".parse().unwrap()); + endpoint.peer_manager.run().await.unwrap(); + endpoint.gateway.start_runtime().await.unwrap(); + + let socket = endpoint + .gateway + .data_plane_udp_bind(0, Duration::from_secs(1)) + .await + .unwrap(); + + assert_eq!(socket.local_addr().port(), 20002); + assert_eq!(host.udp_binds.load(Ordering::Relaxed), 1); + + drop(socket); + endpoint.gateway.stop_runtime().await; + endpoint.peer_manager.clear_resources().await; +} + +#[tokio::test] +async fn final_data_plane_lease_releases_net_and_same_ipv4_reacquires_it() { + let host = Arc::new(TestHost::default()); + let endpoint = data_plane_endpoint(host, "10.126.127.1/24".parse().unwrap()); + endpoint.peer_manager.run().await.unwrap(); + endpoint.gateway.start_runtime().await.unwrap(); + + let socket = endpoint + .gateway + .data_plane_udp_bind(0, Duration::from_secs(1)) + .await + .unwrap(); + assert!(endpoint.gateway.net.lock().await.is_some()); + + drop(socket); + tokio::time::timeout(Duration::from_secs(1), async { + loop { + if endpoint.gateway.net.lock().await.is_none() { + break; + } + tokio::task::yield_now().await; + } + }) + .await + .expect("final data-plane lease did not release smoltcp net"); + + let socket = endpoint + .gateway + .data_plane_udp_bind(0, Duration::from_secs(1)) + .await + .expect("same IPv4 generation should be recreated"); + assert!(endpoint.gateway.net.lock().await.is_some()); + + drop(socket); + endpoint.gateway.stop_runtime().await; + endpoint.peer_manager.clear_resources().await; +} + +#[tokio::test] +async fn immediate_consumer_reacquire_never_leases_closing_generation() { + let host = Arc::new(TestHost::default()); + let endpoint = data_plane_endpoint(host, "10.126.133.1/24".parse().unwrap()); + endpoint.peer_manager.run().await.unwrap(); + endpoint.gateway.start_runtime().await.unwrap(); + let destination: SocketAddr = "10.126.133.2:31001".parse().unwrap(); + + for _ in 0..20 { + let socket = endpoint + .gateway + .data_plane_udp_bind(0, Duration::from_secs(1)) + .await + .expect("consumer should acquire the current smoltcp generation"); + socket + .send_to(b"generation-probe", destination) + .await + .expect("new consumer must not inherit a closing generation"); + drop(socket); + tokio::task::yield_now().await; + } + + endpoint.gateway.stop_runtime().await; + endpoint.peer_manager.clear_resources().await; +} + +#[tokio::test] +async fn ipv4_change_closes_existing_generation_with_typed_error() { + let host = Arc::new(TestHost::default()); + let endpoint = data_plane_endpoint(host, "10.126.128.1/24".parse().unwrap()); + endpoint.peer_manager.run().await.unwrap(); + endpoint.gateway.start_runtime().await.unwrap(); + let socket = endpoint + .gateway + .data_plane_udp_bind(0, Duration::from_secs(1)) + .await + .unwrap(); + + endpoint.gateway.runtime_config.update_peer_with(|peer| { + peer.runtime.core.routes.ipv4 = + Some(IpPrefix::new("10.126.129.1".parse().unwrap(), 24).unwrap()); + }); + + let mut buf = [0u8; 1]; + let error = tokio::time::timeout(Duration::from_secs(1), socket.recv_from(&mut buf)) + .await + .expect("old generation receive did not wake") + .unwrap_err(); + let data_plane_error = error + .get_ref() + .and_then(|error| error.downcast_ref::()) + .expect("generation close must preserve the typed data-plane error"); + assert_eq!(data_plane_error.kind(), DataPlaneErrorKind::NetworkChanged); + + drop(socket); + endpoint.gateway.stop_runtime().await; + endpoint.peer_manager.clear_resources().await; +} + +#[tokio::test] +async fn readiness_timeout_has_stable_error_kind() { + let host = Arc::new(TestHost::default()); + let endpoint = data_plane_endpoint(host, "10.126.130.1/24".parse().unwrap()); + endpoint + .gateway + .runtime_config + .update_peer_with(|peer| peer.runtime.core.routes.ipv4 = None); + endpoint.peer_manager.run().await.unwrap(); + endpoint.gateway.start_runtime().await.unwrap(); + + let error = match endpoint + .gateway + .data_plane_udp_bind(0, Duration::from_millis(1)) + .await + { + Ok(_) => panic!("data-plane bind unexpectedly succeeded without an IPv4 address"), + Err(error) => error, + }; + assert_eq!(error.kind(), DataPlaneErrorKind::DeadlineExceeded); + + endpoint.gateway.stop_runtime().await; + endpoint.peer_manager.clear_resources().await; +} + +#[tokio::test] +async fn data_plane_consumes_modified_data_when_entry_matches() { + let gateway = test_gateway(); + + let local = SocketAddr::new(IpAddr::V4(Ipv4Addr::new(10, 144, 144, 1)), 40000); + let remote = SocketAddr::new(IpAddr::V4(Ipv4Addr::new(10, 144, 144, 3)), 22); + let entry = FlowKey { + src: local, + dst: remote, + kind: TCP_ENTRY, + }; + gateway.entries.insert( + entry, + FlowData::Tcp { + _reservation: Arc::new(()), + }, + ); + + for packet_type in [ + PacketType::DataWithKcpSrcModified, + PacketType::DataWithQuicSrcModified, + ] { + let mut packet = ZCPacket::new_with_payload(&build_tcp_packet(remote, local)); + packet.fill_peer_manager_hdr(1, 1, packet_type as u8); + + let result = gateway.try_process_packet_from_peer(packet).await; + assert!(result.is_none()); + + let mut receiver = gateway.packet_recv.lock().await; + let received = receiver.try_recv().unwrap(); + assert_eq!( + received.peer_manager_header().unwrap().packet_type, + packet_type as u8 + ); + } +} + +#[tokio::test] +async fn data_plane_passes_through_unmatched_or_malformed_modified_data() { + let gateway = test_gateway(); + gateway.entries.insert( + FlowKey { + src: SocketAddr::new(IpAddr::V4(Ipv4Addr::new(10, 144, 144, 1)), 40000), + dst: SocketAddr::new(IpAddr::V4(Ipv4Addr::new(10, 144, 144, 3)), 22), + kind: TCP_ENTRY, + }, + FlowData::Tcp { + _reservation: Arc::new(()), + }, + ); + + let unmatched_local = SocketAddr::new(IpAddr::V4(Ipv4Addr::new(10, 144, 144, 1)), 40001); + let remote = SocketAddr::new(IpAddr::V4(Ipv4Addr::new(10, 144, 144, 3)), 22); + let mut unmatched_packet = + ZCPacket::new_with_payload(&build_tcp_packet(remote, unmatched_local)); + unmatched_packet.fill_peer_manager_hdr(1, 2, PacketType::DataWithKcpSrcModified as u8); + let result = gateway.try_process_packet_from_peer(unmatched_packet).await; + assert!(result.is_some()); + + let mut malformed_packet = ZCPacket::new_with_payload(&[0u8; 8]); + malformed_packet.fill_peer_manager_hdr(1, 2, PacketType::DataWithQuicSrcModified as u8); + let result = gateway.try_process_packet_from_peer(malformed_packet).await; + assert!(result.is_some()); + + let mut receiver = gateway.packet_recv.lock().await; + assert!(receiver.try_recv().is_err()); +} + +#[tokio::test] +async fn data_plane_passes_through_non_loopback_modified_data_when_entry_matches() { + let gateway = test_gateway(); + + let local = SocketAddr::new(IpAddr::V4(Ipv4Addr::new(10, 144, 144, 1)), 40000); + let remote = SocketAddr::new(IpAddr::V4(Ipv4Addr::new(10, 144, 144, 3)), 22); + let entry = FlowKey { + src: local, + dst: remote, + kind: TCP_ENTRY, + }; + gateway.entries.insert( + entry, + FlowData::Tcp { + _reservation: Arc::new(()), + }, + ); + + let mut packet = ZCPacket::new_with_payload(&build_tcp_packet(remote, local)); + packet.fill_peer_manager_hdr(1, 2, PacketType::DataWithKcpSrcModified as u8); + + let result = gateway.try_process_packet_from_peer(packet).await; + assert!(result.is_some()); + + let mut receiver = gateway.packet_recv.lock().await; + assert!(receiver.try_recv().is_err()); +} + +#[tokio::test] +async fn data_plane_mirrors_fragmented_udp_when_entry_matches() { + let gateway = test_gateway(); + + let local = SocketAddr::new(IpAddr::V4(Ipv4Addr::new(10, 144, 144, 1)), 40000); + let remote = SocketAddr::new(IpAddr::V4(Ipv4Addr::new(10, 144, 144, 3)), 53); + gateway.entries.insert( + FlowKey { + src: local, + dst: remote, + kind: UDP_ENTRY, + }, + FlowData::Udp, + ); + assert_eq!(gateway.entries.count(), 1); + + let mut packet = ZCPacket::new_with_payload(&build_udp_followup_fragment( + match remote.ip() { + IpAddr::V4(ip) => ip, + IpAddr::V6(_) => unreachable!(), + }, + match local.ip() { + IpAddr::V4(ip) => ip, + IpAddr::V6(_) => unreachable!(), + }, + )); + packet.fill_peer_manager_hdr(1, 2, PacketType::Data as u8); + + let result = gateway.try_process_packet_from_peer(packet).await; + assert!(result.is_some()); + + let mut receiver = gateway.packet_recv.lock().await; + let received = receiver.try_recv().unwrap(); + assert_eq!( + received.peer_manager_header().unwrap().packet_type, + PacketType::Data as u8 + ); +} diff --git a/easytier-core/src/gateway/dataplane/udp.rs b/easytier-core/src/gateway/dataplane/udp.rs new file mode 100644 index 00000000..03e5e782 --- /dev/null +++ b/easytier-core/src/gateway/dataplane/udp.rs @@ -0,0 +1,80 @@ +//! UDP socket resources exposed by the data plane. + +use std::{ + any::Any, + collections::{HashMap, hash_map::Entry}, + net::SocketAddr, + sync::{Arc, Mutex}, +}; + +use super::{ + DataPlaneIoGuard, DataPlaneLease, DataPlaneUdpIo, FlowData, FlowKey, FlowLease, FlowSet, + UDP_ENTRY, +}; + +pub struct DataPlaneUdpSocket { + pub(super) socket: Arc, + pub(super) flows: FlowSet, + pub(super) routes: Mutex>>, + pub(super) local_addr: SocketAddr, + pub(super) _reservation: Arc, + pub(super) _data_plane_lease: DataPlaneLease, + pub(super) generation: DataPlaneIoGuard, +} + +impl DataPlaneUdpSocket { + pub fn local_addr(&self) -> SocketAddr { + self.local_addr + } + + pub async fn send_to(&self, buf: &[u8], addr: SocketAddr) -> Result { + self.generation + .ensure_open() + .map_err(|error| error.into_io_error())?; + let key = FlowKey { + src: self.local_addr, + dst: addr, + kind: UDP_ENTRY, + }; + if let Entry::Vacant(route) = self.routes.lock().unwrap().entry(addr) { + let lease = FlowLease::try_register(self.flows.clone(), key, FlowData::Udp) + .ok_or_else(|| { + std::io::Error::new( + std::io::ErrorKind::AddrInUse, + "data-plane UDP flow already exists", + ) + })?; + route.insert(lease); + } + tokio::select! { + biased; + _ = self.generation.closed() => Err(self.generation.closed_io_error()), + result = self.socket.send_to(buf, addr) => result, + } + } + + pub async fn recv_from(&self, buf: &mut [u8]) -> Result<(usize, SocketAddr), std::io::Error> { + self.generation + .ensure_open() + .map_err(|error| error.into_io_error())?; + tokio::select! { + biased; + _ = self.generation.closed() => Err(self.generation.closed_io_error()), + result = self.socket.recv_from(buf) => result, + } + } + + pub(super) async fn recv_from_limited( + &self, + max_len: usize, + ) -> Result<(Vec, SocketAddr, bool), std::io::Error> { + self.generation + .ensure_open() + .map_err(|error| error.into_io_error())?; + tokio::select! { + biased; + _ = self.generation.closed() => Err(self.generation.closed_io_error()), + result = self.socket.recv_from_limited(max_len) => result, + } + } +} diff --git a/easytier-core/src/gateway/dhcp.rs b/easytier-core/src/gateway/dhcp.rs new file mode 100644 index 00000000..76320816 --- /dev/null +++ b/easytier-core/src/gateway/dhcp.rs @@ -0,0 +1,564 @@ +use std::{collections::HashSet, net::Ipv4Addr, sync::Arc, time::Duration}; + +use async_trait::async_trait; +use cidr::Ipv4Inet; +use rand::Rng; +use tokio_util::task::AbortOnDropHandle; + +use crate::{ + config::IpPrefix, config::runtime::CoreRuntimeConfigStore, peers::peer_manager::PeerManagerCore, +}; + +#[cfg(feature = "dhcp-ipv4")] +use tokio::sync::Mutex; + +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub enum DhcpIpv4Decision { + WaitForPeers, + Unchanged, + Change { + previous: Option, + next: Option, + }, +} + +#[derive(Debug)] +pub struct DhcpIpv4Allocator { + default_subnet: Ipv4Inet, + current: Option, +} + +impl Default for DhcpIpv4Allocator { + fn default() -> Self { + Self::new(Ipv4Inet::new(Ipv4Addr::new(10, 126, 126, 0), 24).unwrap()) + } +} + +impl DhcpIpv4Allocator { + pub fn new(default_subnet: Ipv4Inet) -> Self { + Self { + default_subnet, + current: None, + } + } + + pub fn current(&self) -> Option { + self.current + } + + pub fn reset(&mut self) { + self.current = None; + } + + pub fn commit(&mut self, next: Option) { + self.current = next; + } + + pub fn evaluate(&self, has_routes: bool, used_ipv4: &HashSet) -> DhcpIpv4Decision { + if !has_routes { + return DhcpIpv4Decision::WaitForPeers; + } + + let subnet = used_ipv4.iter().next().unwrap_or(&self.default_subnet); + if let Some(current) = self.current + && current.network() == subnet.network() + && !used_ipv4.contains(¤t) + { + return DhcpIpv4Decision::Unchanged; + } + + let next = subnet.network().iter().find(|candidate| { + candidate.address() != subnet.first_address() + && candidate.address() != subnet.last_address() + && !used_ipv4.contains(candidate) + }); + if self.current == next { + return DhcpIpv4Decision::Unchanged; + } + + DhcpIpv4Decision::Change { + previous: self.current, + next, + } + } +} + +#[derive(Clone, Debug, Default, PartialEq, Eq)] +pub struct DhcpIpv4RouteSnapshot { + pub has_routes: bool, + pub used_ipv4: HashSet, +} + +#[async_trait] +pub trait DhcpIpv4RouteSource: Send + Sync + 'static { + async fn dhcp_ipv4_route_snapshot(&self) -> DhcpIpv4RouteSnapshot; +} + +#[async_trait] +impl DhcpIpv4RouteSource for PeerManagerCore { + async fn dhcp_ipv4_route_snapshot(&self) -> DhcpIpv4RouteSnapshot { + let routes = self.get_route().list_routes().await; + let has_routes = !routes.is_empty(); + let used_ipv4 = routes + .into_iter() + .filter_map(|route| route.ipv4_addr.map(Into::into)) + .collect(); + DhcpIpv4RouteSnapshot { + has_routes, + used_ipv4, + } + } +} + +#[async_trait] +pub trait DhcpIpv4Host: Send + Sync + 'static { + fn take_interface_closed(&self) -> bool; + + async fn apply_dhcp_ipv4( + &self, + previous: Option, + next: Option, + ) -> DhcpIpv4ApplyOutcome; + + fn publish_dhcp_ipv4( + &self, + _previous: Option, + _requested: Option, + _actual: Option, + ) { + } +} + +#[cfg(feature = "dhcp-ipv4")] +pub(crate) struct DhcpIpv4Runtime { + task: Mutex>>, +} + +#[cfg(feature = "dhcp-ipv4")] +impl DhcpIpv4Runtime { + pub(crate) fn new() -> Self { + Self { + task: Mutex::new(None), + } + } + + pub(crate) async fn start( + &self, + route_source: Arc, + runtime_config: CoreRuntimeConfigStore, + host: Arc, + ) { + let mut task = self.task.lock().await; + if task.is_none() { + task.replace(DhcpIpv4Service::new(route_source, runtime_config, host).start()); + } + } + + pub(crate) async fn stop(&self) { + self.task.lock().await.take(); + } +} + +pub struct DhcpIpv4ApplyPermit { + _guard: Box, +} + +impl DhcpIpv4ApplyPermit { + pub fn new(guard: impl Send + 'static) -> Self { + Self { + _guard: Box::new(guard), + } + } +} + +pub struct DhcpIpv4ApplyOutcome { + pub actual: Option, + pub result: anyhow::Result<()>, + permit: Option, +} + +impl DhcpIpv4ApplyOutcome { + pub fn applied(actual: Option) -> Self { + Self { + actual, + result: Ok(()), + permit: None, + } + } + + pub fn failed(actual: Option, error: impl Into) -> Self { + Self { + actual, + result: Err(error.into()), + permit: None, + } + } + + pub fn with_permit(mut self, permit: DhcpIpv4ApplyPermit) -> Self { + self.permit = Some(permit); + self + } +} + +pub struct DhcpIpv4Service { + operation: tokio::sync::Mutex<()>, + allocator: std::sync::Mutex, + route_source: Arc, + runtime_config: CoreRuntimeConfigStore, + host: Arc, +} + +impl DhcpIpv4Service { + pub fn new( + route_source: Arc, + runtime_config: CoreRuntimeConfigStore, + host: Arc, + ) -> Arc { + Arc::new(Self { + operation: tokio::sync::Mutex::new(()), + allocator: std::sync::Mutex::new(DhcpIpv4Allocator::default()), + route_source, + runtime_config, + host, + }) + } + + pub fn current(&self) -> Option { + self.allocator.lock().unwrap().current() + } + + pub async fn reconcile_once(&self) -> bool { + let _operation = self.operation.lock().await; + if self.host.take_interface_closed() { + self.allocator.lock().unwrap().reset(); + } + let snapshot = self.route_source.dhcp_ipv4_route_snapshot().await; + let decision = self + .allocator + .lock() + .unwrap() + .evaluate(snapshot.has_routes, &snapshot.used_ipv4); + + let DhcpIpv4Decision::Change { previous, next } = decision else { + return snapshot.has_routes; + }; + tracing::debug!(?previous, ?next, "DHCP IPv4 reconciliation applying change"); + let outcome = self.host.apply_dhcp_ipv4(previous, next).await; + let DhcpIpv4ApplyOutcome { + actual, + result, + permit, + } = outcome; + self.runtime_config.update_peer_with(|peer| { + peer.runtime.core.routes.ipv4 = actual.map(|actual| IpPrefix { + address: actual.address().into(), + prefix_len: actual.network_length(), + }); + }); + match result { + Ok(()) => { + self.allocator.lock().unwrap().commit(actual); + self.host.publish_dhcp_ipv4(previous, next, actual); + } + Err(err) => { + tracing::error!(?previous, ?next, ?actual, ?err, "DHCP IPv4 apply failed"); + } + } + drop(permit); + snapshot.has_routes + } + + pub fn start(self: &Arc) -> AbortOnDropHandle<()> { + let service = self.clone(); + AbortOnDropHandle::new(tokio::spawn(async move { + let mut next_sleep = Duration::ZERO; + loop { + crate::foundation::time::sleep(next_sleep).await; + next_sleep = if service.reconcile_once().await { + Duration::from_secs(rand::thread_rng().gen_range(5..10)) + } else { + Duration::from_secs(1) + }; + } + })) + } +} + +#[cfg(test)] +mod tests { + use std::sync::{ + Mutex, + atomic::{AtomicBool, Ordering}, + }; + + use super::*; + + struct StaticRouteSource { + snapshot: Mutex, + } + + #[async_trait] + impl DhcpIpv4RouteSource for StaticRouteSource { + async fn dhcp_ipv4_route_snapshot(&self) -> DhcpIpv4RouteSnapshot { + self.snapshot.lock().unwrap().clone() + } + } + + type PublishedIpv4Route = (Option, Option, Option); + + #[derive(Default)] + struct RecordingHost { + interface_closed: AtomicBool, + fail_apply: AtomicBool, + hold_apply_permit: AtomicBool, + permit_held: Arc, + published_with_permit: AtomicBool, + runtime_config: Mutex>, + published_runtime_ipv4: Mutex>>, + changes: Mutex, Option)>>, + published: Mutex>, + } + + struct RecordingPermit(Arc); + + impl Drop for RecordingPermit { + fn drop(&mut self) { + self.0.store(false, Ordering::Release); + } + } + + #[async_trait] + impl DhcpIpv4Host for RecordingHost { + fn take_interface_closed(&self) -> bool { + self.interface_closed.swap(false, Ordering::AcqRel) + } + + async fn apply_dhcp_ipv4( + &self, + previous: Option, + next: Option, + ) -> DhcpIpv4ApplyOutcome { + self.changes.lock().unwrap().push((previous, next)); + let mut outcome = if self.fail_apply.load(Ordering::Acquire) { + DhcpIpv4ApplyOutcome::failed(None, anyhow::anyhow!("apply failed")) + } else { + DhcpIpv4ApplyOutcome::applied(next) + }; + if self.hold_apply_permit.load(Ordering::Acquire) { + assert!(!self.permit_held.swap(true, Ordering::AcqRel)); + outcome = outcome.with_permit(DhcpIpv4ApplyPermit::new(RecordingPermit( + self.permit_held.clone(), + ))); + } + outcome + } + + fn publish_dhcp_ipv4( + &self, + previous: Option, + requested: Option, + actual: Option, + ) { + self.published_with_permit + .store(self.permit_held.load(Ordering::Acquire), Ordering::Release); + if let Some(runtime_config) = self.runtime_config.lock().unwrap().as_ref() { + self.published_runtime_ipv4.lock().unwrap().push( + runtime_config + .snapshot() + .peer + .runtime + .core + .routes + .ipv4 + .clone(), + ); + } + self.published + .lock() + .unwrap() + .push((previous, requested, actual)); + } + } + + fn service( + snapshot: DhcpIpv4RouteSnapshot, + host: Arc, + ) -> (Arc, CoreRuntimeConfigStore) { + let runtime_config = CoreRuntimeConfigStore::new( + crate::config::runtime::CoreRuntimeConfig::default(), + Arc::new(crate::config::peers::PeerRuntimeSnapshot::default()), + ); + *host.runtime_config.lock().unwrap() = Some(runtime_config.clone()); + let service = DhcpIpv4Service::new( + Arc::new(StaticRouteSource { + snapshot: Mutex::new(snapshot), + }), + runtime_config.clone(), + host, + ); + (service, runtime_config) + } + + #[test] + fn waits_until_at_least_one_route_exists() { + let allocator = DhcpIpv4Allocator::default(); + + assert_eq!( + allocator.evaluate(false, &HashSet::new()), + DhcpIpv4Decision::WaitForPeers + ); + } + + #[test] + fn uses_default_subnet_when_routes_have_no_ipv4() { + let allocator = DhcpIpv4Allocator::default(); + + assert_eq!( + allocator.evaluate(true, &HashSet::new()), + DhcpIpv4Decision::Change { + previous: None, + next: Some("10.126.126.1/24".parse().unwrap()), + } + ); + } + + #[test] + fn keeps_current_address_when_it_is_free_in_the_selected_subnet() { + let mut allocator = DhcpIpv4Allocator::default(); + allocator.commit(Some("10.1.2.8/24".parse().unwrap())); + let used = HashSet::from(["10.1.2.2/24".parse().unwrap()]); + + assert_eq!(allocator.evaluate(true, &used), DhcpIpv4Decision::Unchanged); + } + + #[test] + fn selects_first_available_host_after_a_conflict() { + let mut allocator = DhcpIpv4Allocator::default(); + allocator.commit(Some("10.1.2.1/24".parse().unwrap())); + let used = HashSet::from([ + "10.1.2.1/24".parse().unwrap(), + "10.1.2.2/24".parse().unwrap(), + ]); + + assert_eq!( + allocator.evaluate(true, &used), + DhcpIpv4Decision::Change { + previous: Some("10.1.2.1/24".parse().unwrap()), + next: Some("10.1.2.3/24".parse().unwrap()), + } + ); + } + + #[test] + fn reset_forgets_the_previous_interface_address() { + let mut allocator = DhcpIpv4Allocator::default(); + allocator.commit(Some("10.1.2.8/24".parse().unwrap())); + + allocator.reset(); + + assert_eq!(allocator.current(), None); + } + + #[tokio::test] + async fn service_commits_only_after_host_apply_succeeds() { + let host = Arc::new(RecordingHost::default()); + let (service, runtime_config) = service( + DhcpIpv4RouteSnapshot { + has_routes: true, + used_ipv4: HashSet::new(), + }, + host.clone(), + ); + + assert!(service.reconcile_once().await); + + assert_eq!(service.current(), Some("10.126.126.1/24".parse().unwrap())); + assert_eq!( + *host.changes.lock().unwrap(), + [(None, Some("10.126.126.1/24".parse().unwrap()))] + ); + assert_eq!( + runtime_config.snapshot().peer.runtime.core.routes.ipv4, + Some(IpPrefix::new("10.126.126.1".parse().unwrap(), 24).unwrap()) + ); + assert_eq!( + *host.published.lock().unwrap(), + [( + None, + Some("10.126.126.1/24".parse().unwrap()), + Some("10.126.126.1/24".parse().unwrap()) + )] + ); + } + + #[tokio::test] + async fn service_holds_host_permit_through_store_and_event_commit() { + let host = Arc::new(RecordingHost::default()); + host.hold_apply_permit.store(true, Ordering::Release); + let (service, _runtime_config) = service( + DhcpIpv4RouteSnapshot { + has_routes: true, + used_ipv4: HashSet::new(), + }, + host.clone(), + ); + + service.reconcile_once().await; + + let expected = Some(IpPrefix::new("10.126.126.1".parse().unwrap(), 24).unwrap()); + assert!(host.published_with_permit.load(Ordering::Acquire)); + assert_eq!( + host.published_runtime_ipv4.lock().unwrap().as_slice(), + &[expected] + ); + assert!(!host.permit_held.load(Ordering::Acquire)); + } + + #[tokio::test] + async fn service_retries_change_when_host_apply_fails() { + let host = Arc::new(RecordingHost::default()); + host.fail_apply.store(true, Ordering::Release); + let (service, runtime_config) = service( + DhcpIpv4RouteSnapshot { + has_routes: true, + used_ipv4: HashSet::new(), + }, + host.clone(), + ); + + service.reconcile_once().await; + service.reconcile_once().await; + + assert_eq!(service.current(), None); + assert_eq!(host.changes.lock().unwrap().len(), 2); + assert_eq!( + runtime_config.snapshot().peer.runtime.core.routes.ipv4, + None + ); + assert!(host.published.lock().unwrap().is_empty()); + } + + #[tokio::test] + async fn interface_close_resets_previous_allocation_before_reapply() { + let host = Arc::new(RecordingHost::default()); + let (service, _runtime_config) = service( + DhcpIpv4RouteSnapshot { + has_routes: true, + used_ipv4: HashSet::new(), + }, + host.clone(), + ); + service.reconcile_once().await; + host.interface_closed.store(true, Ordering::Release); + + service.reconcile_once().await; + + assert_eq!( + host.changes.lock().unwrap().as_slice(), + [ + (None, Some("10.126.126.1/24".parse().unwrap())), + (None, Some("10.126.126.1/24".parse().unwrap())), + ] + ); + } +} diff --git a/easytier-core/src/gateway/magic_dns.rs b/easytier-core/src/gateway/magic_dns.rs new file mode 100644 index 00000000..f5f2dba6 --- /dev/null +++ b/easytier-core/src/gateway/magic_dns.rs @@ -0,0 +1,17 @@ +mod records; + +#[cfg(feature = "proxy-packet")] +mod packet; + +pub use records::{ + MagicDnsRecordSnapshot, MagicDnsRecordStore, MagicDnsRoute, MagicDnsRouteAdvertisement, + MagicDnsRoutePublisher, MagicDnsRouteSnapshot, MagicDnsRouteSource, + run_magic_dns_route_publisher, +}; + +#[cfg(feature = "proxy-packet")] +pub(crate) use packet::magic_dns_packet_filter; +#[cfg(feature = "proxy-packet")] +pub use packet::{ + MagicDnsQuery, MagicDnsQueryResolver, MagicDnsResolverRegistration, process_magic_dns_packet, +}; diff --git a/easytier-core/src/gateway/magic_dns/packet.rs b/easytier-core/src/gateway/magic_dns/packet.rs new file mode 100644 index 00000000..16bc05d9 --- /dev/null +++ b/easytier-core/src/gateway/magic_dns/packet.rs @@ -0,0 +1,468 @@ +use std::{future::Future, net::Ipv4Addr}; + +use async_trait::async_trait; +use pnet_packet::{ + MutablePacket, Packet, + icmp::{self, IcmpPacket, IcmpTypes, MutableIcmpPacket}, + ip::IpNextHeaderProtocols, + ipv4::{self, Ipv4Flags, Ipv4Packet, MutableIpv4Packet}, + udp::{self, MutableUdpPacket, UdpPacket}, +}; + +use crate::{ + config::PeerId, + packet::ZCPacket, + peers::{BoxNicPacketFilter, NicPacketFilter}, +}; + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct MagicDnsQuery { + pub source: std::net::SocketAddr, + pub payload: Vec, +} + +#[async_trait] +pub trait MagicDnsQueryResolver: Send + Sync + 'static { + async fn resolve(&self, query: MagicDnsQuery) -> Option>; +} + +/// Owns one Magic DNS resolver installed in the core NIC pipeline. +/// +/// `close` waits until readers that may already be invoking the resolver have +/// finished, then removes the entry so the resolver can be dropped promptly. +pub struct MagicDnsResolverRegistration { + peer_manager: std::sync::Weak, + pipeline: crate::peers::peer_manager::PipelineRegistrationGuard, + runtime: tokio::runtime::Handle, +} + +impl MagicDnsResolverRegistration { + pub(crate) fn new( + peer_manager: std::sync::Weak, + pipeline: crate::peers::peer_manager::PipelineRegistrationGuard, + runtime: tokio::runtime::Handle, + ) -> Self { + Self { + peer_manager, + pipeline, + runtime, + } + } + + pub async fn close(&self) { + self.pipeline.close(); + if let Some(peer_manager) = self.peer_manager.upgrade() { + peer_manager + .remove_managed_nic_packet_process_pipeline(&self.pipeline) + .await; + } + } +} + +impl Drop for MagicDnsResolverRegistration { + fn drop(&mut self) { + self.pipeline.close(); + let Some(peer_manager) = self.peer_manager.upgrade() else { + return; + }; + let pipeline = self.pipeline.clone(); + self.runtime.spawn(async move { + peer_manager + .remove_managed_nic_packet_process_pipeline(&pipeline) + .await; + }); + } +} + +struct MagicDnsPacketFilter { + fake_ip: Ipv4Addr, + my_peer_id: PeerId, + resolver: std::sync::Arc, +} + +pub(crate) fn magic_dns_packet_filter( + fake_ip: Ipv4Addr, + my_peer_id: PeerId, + resolver: std::sync::Arc, +) -> BoxNicPacketFilter { + Box::new(MagicDnsPacketFilter { + fake_ip, + my_peer_id, + resolver, + }) +} + +#[async_trait] +impl NicPacketFilter for MagicDnsPacketFilter { + async fn try_process_packet_from_nic(&self, packet: &mut ZCPacket) -> bool { + process_magic_dns_packet(packet, self.fake_ip, self.my_peer_id, |query| { + self.resolver.resolve(query) + }) + .await + } + + fn id(&self) -> String { + "magic_dns_server".to_owned() + } +} + +pub async fn process_magic_dns_packet( + packet: &mut ZCPacket, + fake_ip: Ipv4Addr, + my_peer_id: PeerId, + resolve: F, +) -> bool +where + F: FnOnce(MagicDnsQuery) -> Fut, + Fut: Future>>, +{ + if packet.peer_manager_header().is_none() { + return false; + } + let Some(ip_packet) = Ipv4Packet::new(packet.payload()) else { + return false; + }; + if ip_packet.get_version() != 4 || ip_packet.get_destination() != fake_ip { + return false; + } + + let ip_header_length = ip_packet.get_header_length() as usize * 4; + let ip_total_length = ip_packet.get_total_length() as usize; + if ip_header_length < MutableIpv4Packet::minimum_packet_size() + || ip_header_length > ip_total_length + || ip_total_length != packet.payload().len() + || ip_packet.get_fragment_offset() != 0 + || ip_packet.get_flags() & Ipv4Flags::MoreFragments != 0 + { + return false; + } + + let protocol = ip_packet.get_next_level_protocol(); + let source_ip = ip_packet.get_source(); + let destination_ip = ip_packet.get_destination(); + + match protocol { + IpNextHeaderProtocols::Udp => { + let ip_payload = &packet.payload()[ip_header_length..ip_total_length]; + let Some(udp_packet) = UdpPacket::new(ip_payload) else { + return false; + }; + let udp_length = udp_packet.get_length() as usize; + if udp_length != ip_payload.len() || udp_length < UdpPacket::minimum_packet_size() { + return false; + } + if udp_packet.get_destination() != 53 { + return false; + } + let source_port = udp_packet.get_source(); + let destination_port = udp_packet.get_destination(); + let query = MagicDnsQuery { + source: std::net::SocketAddr::from((source_ip, source_port)), + payload: udp_packet.payload().to_vec(), + }; + let Some(response) = resolve(query).await else { + return false; + }; + if !apply_udp_response( + packet, + source_ip, + destination_ip, + source_port, + destination_port, + ip_header_length, + &response, + ) { + return false; + } + } + IpNextHeaderProtocols::Icmp => { + let Some(icmp_packet) = IcmpPacket::new(&packet.payload()[ip_header_length..]) else { + return false; + }; + if icmp_packet.get_icmp_type() != IcmpTypes::EchoRequest { + return false; + } + let Some(mut icmp_packet) = + MutableIcmpPacket::new(&mut packet.mut_payload()[ip_header_length..]) + else { + return false; + }; + icmp_packet.set_icmp_type(IcmpTypes::EchoReply); + icmp_packet.set_checksum(icmp::checksum(&icmp_packet.to_immutable())); + } + _ => return false, + } + + let Some(mut ip_packet) = MutableIpv4Packet::new(packet.mut_payload()) else { + return false; + }; + ip_packet.set_source(destination_ip); + ip_packet.set_destination(source_ip); + ip_packet.set_checksum(ipv4::checksum(&ip_packet.to_immutable())); + let payload_length = packet.payload().len() as u32; + let Some(header) = packet.mut_peer_manager_header() else { + return false; + }; + header.to_peer_id = my_peer_id.into(); + header.len.set(payload_length); + true +} + +#[allow(clippy::too_many_arguments)] +fn apply_udp_response( + packet: &mut ZCPacket, + source_ip: Ipv4Addr, + destination_ip: Ipv4Addr, + source_port: u16, + destination_port: u16, + ip_header_length: usize, + response: &[u8], +) -> bool { + let Some(udp_length) = UdpPacket::minimum_packet_size().checked_add(response.len()) else { + return false; + }; + let Some(ip_length) = ip_header_length.checked_add(udp_length) else { + return false; + }; + if ip_length > u16::MAX as usize { + return false; + } + let Some(header_length) = packet.buf_len().checked_sub(packet.payload().len()) else { + return false; + }; + let Some(inner_length) = header_length.checked_add(ip_length) else { + return false; + }; + + if packet.mut_inner().capacity() < inner_length { + packet + .mut_inner() + .truncate(header_length + ip_header_length + UdpPacket::minimum_packet_size()); + } + packet.mut_inner().resize(inner_length, 0); + + let Some(mut ip_packet) = MutableIpv4Packet::new(packet.mut_payload()) else { + return false; + }; + ip_packet.set_total_length(ip_length as u16); + let Some(mut udp_packet) = MutableUdpPacket::new(ip_packet.payload_mut()) else { + return false; + }; + udp_packet.set_length(udp_length as u16); + udp_packet.set_source(destination_port); + udp_packet.set_destination(source_port); + udp_packet.payload_mut().copy_from_slice(response); + udp_packet.set_checksum(udp::ipv4_checksum( + &udp_packet.to_immutable(), + &destination_ip, + &source_ip, + )); + true +} + +#[cfg(test)] +mod tests { + use super::*; + + fn udp_query(payload: &[u8], destination_port: u16) -> ZCPacket { + let mut bytes = vec![0; 20 + 8 + payload.len()]; + { + let mut ip = MutableIpv4Packet::new(&mut bytes).unwrap(); + ip.set_version(4); + ip.set_header_length(5); + ip.set_total_length((20 + 8 + payload.len()) as u16); + ip.set_next_level_protocol(IpNextHeaderProtocols::Udp); + ip.set_source("10.0.0.2".parse().unwrap()); + ip.set_destination("100.100.100.101".parse().unwrap()); + let mut udp = MutableUdpPacket::new(ip.payload_mut()).unwrap(); + udp.set_source(53000); + udp.set_destination(destination_port); + udp.set_length((8 + payload.len()) as u16); + udp.payload_mut().copy_from_slice(payload); + } + ZCPacket::new_with_payload(&bytes) + } + + fn icmp_echo_request() -> ZCPacket { + let mut bytes = vec![0; 20 + 8]; + { + let mut ip = MutableIpv4Packet::new(&mut bytes).unwrap(); + ip.set_version(4); + ip.set_header_length(5); + ip.set_total_length(28); + ip.set_next_level_protocol(IpNextHeaderProtocols::Icmp); + ip.set_source("10.0.0.2".parse().unwrap()); + ip.set_destination("100.100.100.101".parse().unwrap()); + let mut icmp = MutableIcmpPacket::new(ip.payload_mut()).unwrap(); + icmp.set_icmp_type(IcmpTypes::EchoRequest); + } + ZCPacket::new_with_payload(&bytes) + } + #[tokio::test] + async fn packet_engine_rewrites_dns_query_response() { + let mut packet = udp_query(b"query", 53); + let handled = process_magic_dns_packet( + &mut packet, + "100.100.100.101".parse().unwrap(), + 42, + |query| async move { + assert_eq!(query.source, "10.0.0.2:53000".parse().unwrap()); + assert_eq!(query.payload, b"query"); + Some(b"response".to_vec()) + }, + ) + .await; + + assert!(handled); + let ip = Ipv4Packet::new(packet.payload()).unwrap(); + assert_eq!( + ip.get_source(), + "100.100.100.101".parse::().unwrap() + ); + assert_eq!( + ip.get_destination(), + "10.0.0.2".parse::().unwrap() + ); + let udp = UdpPacket::new(ip.payload()).unwrap(); + assert_eq!(udp.get_source(), 53); + assert_eq!(udp.get_destination(), 53000); + assert_eq!(udp.payload(), b"response"); + assert_eq!(packet.get_dst_peer_id(), Some(42)); + assert_eq!( + packet.peer_manager_header().unwrap().len.get() as usize, + packet.payload().len() + ); + } + + #[tokio::test] + async fn packet_engine_rejects_invalid_ipv4_header_without_mutation() { + let mut packet = udp_query(b"query", 53); + MutableIpv4Packet::new(packet.mut_payload()) + .unwrap() + .set_header_length(15); + let original = packet.payload().to_vec(); + + assert!( + !process_magic_dns_packet( + &mut packet, + "100.100.100.101".parse().unwrap(), + 42, + |_| async { panic!("invalid IPv4 header must not invoke DNS") }, + ) + .await + ); + assert_eq!(packet.payload(), original); + } + + #[tokio::test] + async fn packet_engine_rejects_short_zc_packet_without_panicking() { + let mut packet = + ZCPacket::new_from_buf(Default::default(), crate::packet::ZCPacketType::NIC); + + assert!( + !process_magic_dns_packet( + &mut packet, + "100.100.100.101".parse().unwrap(), + 42, + |_| async { panic!("short packet must not invoke DNS") }, + ) + .await + ); + } + + #[tokio::test] + async fn packet_engine_rejects_inconsistent_udp_length_without_mutation() { + let mut packet = udp_query(b"query", 53); + let mut ip = MutableIpv4Packet::new(packet.mut_payload()).unwrap(); + MutableUdpPacket::new(ip.payload_mut()) + .unwrap() + .set_length(8); + let original = packet.payload().to_vec(); + + assert!( + !process_magic_dns_packet( + &mut packet, + "100.100.100.101".parse().unwrap(), + 42, + |_| async { panic!("invalid UDP length must not invoke DNS") }, + ) + .await + ); + assert_eq!(packet.payload(), original); + } + + #[tokio::test] + async fn packet_engine_rejects_fragmented_packets_without_mutation() { + let mut packet = udp_query(b"query", 53); + MutableIpv4Packet::new(packet.mut_payload()) + .unwrap() + .set_flags(Ipv4Flags::MoreFragments); + let original = packet.payload().to_vec(); + + assert!( + !process_magic_dns_packet( + &mut packet, + "100.100.100.101".parse().unwrap(), + 42, + |_| async { panic!("fragmented packet must not invoke DNS") }, + ) + .await + ); + assert_eq!(packet.payload(), original); + } + + #[tokio::test] + async fn packet_engine_rejects_oversized_response_without_mutation() { + let mut packet = udp_query(b"query", 53); + let original = packet.payload().to_vec(); + let original_length = packet.buf_len(); + + assert!( + !process_magic_dns_packet( + &mut packet, + "100.100.100.101".parse().unwrap(), + 42, + |_| async { Some(vec![0; u16::MAX as usize]) }, + ) + .await + ); + assert_eq!(packet.buf_len(), original_length); + assert_eq!(packet.payload(), original); + } + + #[tokio::test] + async fn packet_engine_replies_to_icmp_without_calling_dns() { + let mut packet = icmp_echo_request(); + let handled = process_magic_dns_packet( + &mut packet, + "100.100.100.101".parse().unwrap(), + 7, + |_| async { panic!("ICMP must not invoke DNS") }, + ) + .await; + + assert!(handled); + let ip = Ipv4Packet::new(packet.payload()).unwrap(); + assert_eq!( + ip.get_source(), + "100.100.100.101".parse::().unwrap() + ); + let icmp = pnet_packet::icmp::IcmpPacket::new(ip.payload()).unwrap(); + assert_eq!(icmp.get_icmp_type(), IcmpTypes::EchoReply); + assert_eq!(packet.get_dst_peer_id(), Some(7)); + } + + #[tokio::test] + async fn packet_engine_ignores_non_dns_udp() { + let mut packet = udp_query(b"query", 5353); + assert!( + !process_magic_dns_packet( + &mut packet, + "100.100.100.101".parse().unwrap(), + 42, + |_| async { Some(Vec::new()) }, + ) + .await + ); + } +} diff --git a/easytier-core/src/gateway/magic_dns/records.rs b/easytier-core/src/gateway/magic_dns/records.rs new file mode 100644 index 00000000..7d449868 --- /dev/null +++ b/easytier-core/src/gateway/magic_dns/records.rs @@ -0,0 +1,429 @@ +use std::{collections::BTreeMap, net::Ipv4Addr, sync::Mutex, time::Duration}; + +use async_trait::async_trait; +use quanta::Instant; + +use crate::peers::peer_manager::PeerManagerCore; + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct MagicDnsRoute { + pub hostname: String, + pub ipv4_addr: Option, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct MagicDnsRouteSnapshot { + pub revision: Instant, + pub routes: Vec, + pub zone: String, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct MagicDnsRouteAdvertisement { + pub hostname: String, + pub ipv4_addr: Option, +} + +#[async_trait] +pub trait MagicDnsRouteSource: Send + Sync { + async fn snapshot(&self) -> MagicDnsRouteSnapshot; + async fn revision(&self) -> Instant; +} + +fn magic_dns_route_advertisement( + route: crate::proto::core_peer::peer::Route, +) -> MagicDnsRouteAdvertisement { + MagicDnsRouteAdvertisement { + hostname: route.hostname, + ipv4_addr: route.ipv4_addr, + } +} + +fn magic_dns_route_snapshot( + revision: Instant, + routes: Vec, + local_identity: (String, Option, String), +) -> MagicDnsRouteSnapshot { + let mut routes = routes + .into_iter() + .map(magic_dns_route_advertisement) + .collect::>(); + let (hostname, ipv4_addr, zone) = local_identity; + routes.push(MagicDnsRouteAdvertisement { + hostname, + ipv4_addr, + }); + MagicDnsRouteSnapshot { + revision, + routes, + zone, + } +} + +#[async_trait] +impl MagicDnsRouteSource for PeerManagerCore { + async fn snapshot(&self) -> MagicDnsRouteSnapshot { + let revision = self.get_route().get_peer_info_last_update_time().await; + magic_dns_route_snapshot( + revision, + self.list_route_snapshots().await, + self.dns_route_identity(), + ) + } + + async fn revision(&self) -> Instant { + self.get_route().get_peer_info_last_update_time().await + } +} + +#[async_trait] +pub trait MagicDnsRoutePublisher: Send { + async fn handshake(&mut self) -> anyhow::Result<()>; + async fn heartbeat(&mut self) -> anyhow::Result<()>; + async fn publish(&mut self, snapshot: &MagicDnsRouteSnapshot) -> anyhow::Result<()>; +} + +pub async fn run_magic_dns_route_publisher( + source: &S, + publisher: &mut P, + unchanged_interval: Duration, +) -> anyhow::Result<()> +where + S: MagicDnsRouteSource + ?Sized, + P: MagicDnsRoutePublisher + ?Sized, +{ + let mut published_revision = None; + publisher.handshake().await?; + loop { + publisher.heartbeat().await?; + + let snapshot = source.snapshot().await; + if published_revision == Some(snapshot.revision) { + crate::foundation::time::sleep(unchanged_interval).await; + continue; + } + + publisher.publish(&snapshot).await?; + if source.revision().await == snapshot.revision { + published_revision = Some(snapshot.revision); + } + } +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct MagicDnsRecordSnapshot { + pub zones: BTreeMap>, +} + +#[derive(Debug, Default)] +pub struct MagicDnsRecordStore { + zones: Mutex>>>, +} + +impl MagicDnsRecordStore { + /// Replaces one client's routes within a zone. + /// + /// Returns `true` when the update removed the final client from an + /// existing zone. The host can use that signal to keep an empty zone + /// authoritative. + pub fn replace_client_routes( + &self, + zone: String, + client: String, + routes: Vec, + ) -> bool { + let mut zones = self.zones.lock().unwrap(); + let Some(routes_by_client) = zones.get_mut(&zone) else { + if !routes.is_empty() { + zones.entry(zone).or_default().insert(client, routes); + } + return false; + }; + + routes_by_client.remove(&client); + if !routes.is_empty() { + routes_by_client.insert(client, routes); + } + if !routes_by_client.is_empty() { + return false; + } + zones.remove(&zone); + true + } + + /// Removes a disconnected client from every zone and returns the zones + /// that became empty. + pub fn remove_client(&self, client: &str) -> Vec { + let mut zones = self.zones.lock().unwrap(); + let mut removed_zones = Vec::new(); + zones.retain(|zone, routes_by_client| { + routes_by_client.remove(client); + let retain = !routes_by_client.is_empty(); + if !retain { + removed_zones.push(zone.clone()); + } + retain + }); + removed_zones + } + + pub fn snapshot(&self) -> MagicDnsRecordSnapshot { + let zones = self.zones.lock().unwrap(); + MagicDnsRecordSnapshot { + zones: zones + .iter() + .map(|(zone, routes_by_client)| { + ( + zone.clone(), + routes_by_client + .values() + .flat_map(|routes| routes.iter().cloned()) + .collect(), + ) + }) + .collect(), + } + } +} + +#[cfg(test)] +mod tests { + use std::sync::Arc; + + use super::*; + + #[test] + fn route_advertisement_preserves_untrusted_prefix_without_parsing() { + let ipv4_addr = crate::proto::common::Ipv4Inet { + address: Some("192.0.2.1".parse::().unwrap().into()), + network_length: 33, + }; + + let advertisement = magic_dns_route_advertisement(crate::proto::core_peer::peer::Route { + hostname: "remote".to_owned(), + ipv4_addr: Some(ipv4_addr), + ..Default::default() + }); + + assert_eq!(advertisement.ipv4_addr, Some(ipv4_addr)); + } + + #[test] + fn route_snapshot_appends_local_identity() { + let ipv4_addr: crate::proto::common::Ipv4Inet = + "10.20.0.91/16".parse::().unwrap().into(); + let snapshot = magic_dns_route_snapshot( + Instant::now(), + Vec::new(), + ( + "portable-node".to_owned(), + Some(ipv4_addr), + "et.net.".to_owned(), + ), + ); + + assert_eq!(snapshot.zone, "et.net."); + assert_eq!( + snapshot.routes, + [MagicDnsRouteAdvertisement { + hostname: "portable-node".to_owned(), + ipv4_addr: Some(ipv4_addr), + }] + ); + } + + struct TestRouteSource { + revision: Mutex, + } + + #[async_trait] + impl MagicDnsRouteSource for TestRouteSource { + async fn snapshot(&self) -> MagicDnsRouteSnapshot { + MagicDnsRouteSnapshot { + revision: *self.revision.lock().unwrap(), + routes: vec![MagicDnsRouteAdvertisement { + hostname: "node-a".to_owned(), + ipv4_addr: Some("10.1.0.1/24".parse::().unwrap().into()), + }], + zone: "et.net.".to_owned(), + } + } + + async fn revision(&self) -> Instant { + *self.revision.lock().unwrap() + } + } + + struct TestRoutePublisher { + source: Arc, + heartbeat_calls: usize, + fail_heartbeat_at: usize, + change_revision_on_first_publish: bool, + handshake_calls: usize, + snapshots: Vec, + } + + #[async_trait] + impl MagicDnsRoutePublisher for TestRoutePublisher { + async fn handshake(&mut self) -> anyhow::Result<()> { + self.handshake_calls += 1; + Ok(()) + } + + async fn heartbeat(&mut self) -> anyhow::Result<()> { + self.heartbeat_calls += 1; + if self.heartbeat_calls == self.fail_heartbeat_at { + anyhow::bail!("stop test publisher"); + } + Ok(()) + } + + async fn publish(&mut self, snapshot: &MagicDnsRouteSnapshot) -> anyhow::Result<()> { + self.snapshots.push(snapshot.clone()); + if self.change_revision_on_first_publish && self.snapshots.len() == 1 { + *self.source.revision.lock().unwrap() = snapshot.revision + Duration::from_secs(1); + } + Ok(()) + } + } + + fn test_publisher(source: Arc) -> TestRoutePublisher { + TestRoutePublisher { + source, + heartbeat_calls: 0, + fail_heartbeat_at: 3, + change_revision_on_first_publish: false, + handshake_calls: 0, + snapshots: Vec::new(), + } + } + + #[tokio::test] + async fn route_publisher_skips_unchanged_snapshot() { + let source = Arc::new(TestRouteSource { + revision: Mutex::new(Instant::now()), + }); + let mut publisher = test_publisher(source.clone()); + + let error = run_magic_dns_route_publisher( + source.as_ref(), + &mut publisher, + Duration::from_millis(1), + ) + .await + .unwrap_err(); + + assert!(error.to_string().contains("stop test publisher")); + assert_eq!(publisher.handshake_calls, 1); + assert_eq!(publisher.snapshots.len(), 1); + } + + #[tokio::test] + async fn route_publisher_retries_change_during_publish() { + let source = Arc::new(TestRouteSource { + revision: Mutex::new(Instant::now()), + }); + let mut publisher = test_publisher(source.clone()); + publisher.change_revision_on_first_publish = true; + + let error = run_magic_dns_route_publisher( + source.as_ref(), + &mut publisher, + Duration::from_millis(1), + ) + .await + .unwrap_err(); + + assert!(error.to_string().contains("stop test publisher")); + assert_eq!(publisher.snapshots.len(), 2); + assert_ne!( + publisher.snapshots[0].revision, + publisher.snapshots[1].revision + ); + } + + fn route(hostname: &str, addr: [u8; 4]) -> MagicDnsRoute { + MagicDnsRoute { + hostname: hostname.to_owned(), + ipv4_addr: Some(addr.into()), + } + } + + #[test] + fn replaces_routes_for_the_same_client_without_touching_other_clients() { + let store = MagicDnsRecordStore::default(); + assert!(!store.replace_client_routes( + "et.net.".to_owned(), + "tcp://client-a".to_owned(), + vec![route("old-a", [10, 0, 0, 1])], + )); + assert!(!store.replace_client_routes( + "et.net.".to_owned(), + "tcp://client-b".to_owned(), + vec![route("peer-b", [10, 0, 0, 2])], + )); + assert!(!store.replace_client_routes( + "et.net.".to_owned(), + "tcp://client-a".to_owned(), + vec![route("new-a", [10, 0, 0, 3])], + )); + + let routes = &store.snapshot().zones["et.net."]; + assert_eq!(routes.len(), 2); + assert!(routes.iter().any(|route| route.hostname == "new-a")); + assert!(routes.iter().any(|route| route.hostname == "peer-b")); + assert!(!routes.iter().any(|route| route.hostname == "old-a")); + } + + #[test] + fn empty_update_removes_only_the_target_client_and_reports_empty_zone() { + let store = MagicDnsRecordStore::default(); + store.replace_client_routes( + "et.net.".to_owned(), + "tcp://client-a".to_owned(), + vec![route("peer-a", [10, 0, 0, 1])], + ); + store.replace_client_routes( + "et.net.".to_owned(), + "tcp://client-b".to_owned(), + vec![route("peer-b", [10, 0, 0, 2])], + ); + + assert!(!store.replace_client_routes( + "et.net.".to_owned(), + "tcp://client-a".to_owned(), + Vec::new(), + )); + assert!(store.replace_client_routes( + "et.net.".to_owned(), + "tcp://client-b".to_owned(), + Vec::new(), + )); + assert!(store.snapshot().zones.is_empty()); + } + + #[test] + fn disconnect_removes_client_from_all_zones() { + let store = MagicDnsRecordStore::default(); + for zone in ["a.et.net.", "b.et.net."] { + store.replace_client_routes( + zone.to_owned(), + "tcp://client-a".to_owned(), + vec![route("peer-a", [10, 0, 0, 1])], + ); + } + store.replace_client_routes( + "b.et.net.".to_owned(), + "tcp://client-b".to_owned(), + vec![route("peer-b", [10, 0, 0, 2])], + ); + + assert_eq!( + store.remove_client("tcp://client-a"), + vec!["a.et.net.".to_owned()] + ); + let snapshot = store.snapshot(); + assert_eq!(snapshot.zones.len(), 1); + assert_eq!(snapshot.zones["b.et.net."][0].hostname, "peer-b"); + } +} diff --git a/easytier-core/src/gateway/mod.rs b/easytier-core/src/gateway/mod.rs new file mode 100644 index 00000000..d85ba0eb --- /dev/null +++ b/easytier-core/src/gateway/mod.rs @@ -0,0 +1,32 @@ +//! Packet-plane features: the gateway dataplane, proxy services, and the +//! instance-level network services. See `CONTEXT.md` "Gateway dataplane" +//! and "Module layers". + +#[cfg(feature = "proxy-smoltcp-stack")] +mod dataplane; +pub mod dhcp; +pub mod magic_dns; +#[cfg(feature = "proxy-smoltcp-stack")] +mod port_forward; +pub mod proxy; +#[cfg(feature = "proxy-smoltcp-stack")] +mod smoltcp; +#[cfg(feature = "proxy-smoltcp-stack")] +mod socks5; +#[cfg(feature = "proxy-packet")] +pub mod udp_broadcast; +pub mod vpn_portal; + +#[cfg(feature = "proxy-smoltcp-stack")] +pub(crate) use dataplane::DataPlaneRuntime; +#[cfg(feature = "proxy-smoltcp-stack")] +pub use dataplane::{ + DataPlaneCompletionDescriptor, DataPlaneCompletionStatus, DataPlaneError, DataPlaneErrorKind, + DataPlaneOperationId, DataPlaneOperationKind, DataPlaneOperationOutcome, + DataPlaneOperationResult, DataPlaneResourceId, DataPlaneSession, DataPlaneSessionLimits, + DataPlaneTcpListener, DataPlaneTcpStream, DataPlaneUdpSocket, +}; +#[cfg(feature = "proxy-smoltcp-stack")] +pub(crate) use port_forward::PortForwardAdapter; +#[cfg(feature = "proxy-smoltcp-stack")] +pub(crate) use socks5::Socks5GatewayAdapter; diff --git a/easytier-core/src/gateway/port_forward.rs b/easytier-core/src/gateway/port_forward.rs new file mode 100644 index 00000000..93f4cc54 --- /dev/null +++ b/easytier-core/src/gateway/port_forward.rs @@ -0,0 +1,453 @@ +//! Host port-forward adapter backed by the data-plane runtime. + +use std::{ + net::SocketAddr, + sync::{ + Arc, + atomic::{AtomicBool, Ordering}, + }, + time::Duration, +}; + +use crossbeam::atomic::AtomicCell; +use dashmap::DashMap; +use quanta::Instant; +use tokio::{select, sync::Mutex, task::JoinSet}; +use tokio_util::{ + sync::{CancellationToken, DropGuard}, + task::AbortOnDropHandle, +}; + +use crate::{ + config::{gateway::PortForwardConfig, runtime::CoreRuntimeConfigStore}, + events::{CoreEvent, CoreEventSink}, + foundation::task::reap_joinset_background, + gateway::dataplane::{ + DataPlaneConsumerLease, DataPlaneRuntime, DataPlaneTcpConnectOptions, DataPlaneTcpStream, + DataPlaneUdpSocket, + }, + socket::{ + SocketContext, + tcp::{ + TcpListenOptions, TcpSocketPurpose, VirtualTcpListener, VirtualTcpListenerFactory, + VirtualTcpSocket, VirtualTcpSocketFactory, + }, + udp::{UdpBindOptions, VirtualUdpSocket, VirtualUdpSocketFactory}, + }, +}; + +#[derive(Debug, Eq, PartialEq, Hash, Clone)] +struct UdpClientKey { + client_addr: SocketAddr, + forward: PortForwardConfig, +} + +enum PortForwardUdpFlow +where + H: VirtualUdpSocketFactory, +{ + Host(Arc), + DataPlane(Arc), +} + +impl PortForwardUdpFlow +where + H: VirtualUdpSocketFactory, +{ + async fn send_to(&self, buf: &[u8], addr: SocketAddr) -> std::io::Result { + match self { + Self::Host(socket) => socket.send_to(buf, addr).await, + Self::DataPlane(socket) => socket.send_to(buf, addr).await, + } + } + + async fn recv_from(&self, buf: &mut [u8]) -> std::io::Result<(usize, SocketAddr)> { + match self { + Self::Host(socket) => socket.recv_from(buf).await, + Self::DataPlane(socket) => socket.recv_from(buf).await, + } + } +} + +struct UdpClientInfo +where + H: VirtualUdpSocketFactory, +{ + flow: Arc>, + last_active: AtomicCell, +} + +pub(crate) struct PortForwardAdapter +where + H: VirtualTcpSocketFactory + VirtualTcpListenerFactory + VirtualUdpSocketFactory, +{ + operation: Mutex<()>, + started: AtomicBool, + runtime_config: CoreRuntimeConfigStore, + data_plane: Arc>, + host: Arc, + socket_context: SocketContext, + events: Arc, + tasks: Arc>>, + cancel_tokens: Arc>, + udp_clients: Arc>>>, + udp_response_tasks: Arc>>, + consumer_lease: Mutex>, +} + +impl PortForwardAdapter +where + H: VirtualTcpSocketFactory + VirtualTcpListenerFactory + VirtualUdpSocketFactory, +{ + pub(crate) fn new( + runtime_config: CoreRuntimeConfigStore, + data_plane: Arc>, + host: Arc, + socket_context: SocketContext, + events: Arc, + ) -> Arc { + Arc::new(Self { + operation: Mutex::new(()), + started: AtomicBool::new(false), + runtime_config, + data_plane, + host, + socket_context, + events, + tasks: Arc::new(std::sync::Mutex::new(JoinSet::new())), + cancel_tokens: Arc::new(DashMap::new()), + udp_clients: Arc::new(DashMap::new()), + udp_response_tasks: Arc::new(DashMap::new()), + consumer_lease: Mutex::new(None), + }) + } + + pub(crate) async fn start(&self) -> anyhow::Result<()> { + let _operation = self.operation.lock().await; + if self.started.load(Ordering::Acquire) { + return Ok(()); + } + self.tasks.lock().unwrap().spawn(reap_joinset_background( + self.tasks.clone(), + "port-forward adapter", + )); + self.start_udp_reaper(); + let cfgs = self + .runtime_config + .snapshot() + .services + .gateway + .port_forwards + .clone(); + if let Err(error) = self.apply_port_forwards(&cfgs).await { + self.stop_inner().await; + return Err(error); + } + self.started.store(true, Ordering::Release); + Ok(()) + } + + pub(crate) async fn reload(&self, cfgs: &[PortForwardConfig]) -> anyhow::Result<()> { + let _operation = self.operation.lock().await; + if !self.started.load(Ordering::Acquire) { + return Ok(()); + } + self.apply_port_forwards(cfgs).await + } + + async fn apply_port_forwards(&self, cfgs: &[PortForwardConfig]) -> anyhow::Result<()> { + for cfg in cfgs { + if !matches!(cfg.proto.to_lowercase().as_str(), "tcp" | "udp") { + anyhow::bail!( + "unsupported protocol: {}, only support udp / tcp", + cfg.proto + ); + } + } + + if !cfgs.is_empty() { + let mut consumer_lease = self.consumer_lease.lock().await; + if consumer_lease.is_none() { + consumer_lease.replace(self.data_plane.acquire_consumer_lease()?); + } + } + + self.cancel_tokens.retain(|current, _| { + cfgs.iter().any(|next| { + if next.dst_addr.ip().is_unspecified() { + current.bind_addr == next.bind_addr && current.proto == next.proto + } else { + current == next + } + }) + }); + self.udp_clients + .retain(|key, _| self.cancel_tokens.contains_key(&key.forward)); + self.udp_response_tasks + .retain(|key, _| self.udp_clients.contains_key(key)); + for cfg in cfgs { + if !self.cancel_tokens.contains_key(cfg) { + self.add_port_forward(cfg.clone()).await?; + } + } + if cfgs.is_empty() { + self.consumer_lease.lock().await.take(); + } + Ok(()) + } + + async fn add_port_forward(&self, cfg: PortForwardConfig) -> anyhow::Result<()> { + match cfg.proto.to_lowercase().as_str() { + "tcp" => self.add_tcp_port_forward(&cfg).await?, + "udp" => self.add_udp_port_forward(&cfg).await?, + _ => { + anyhow::bail!( + "unsupported protocol: {}, only support udp / tcp", + cfg.proto + ) + } + } + self.events.emit(CoreEvent::GatewayPortForwardAdded(cfg)); + Ok(()) + } + + async fn add_tcp_port_forward(&self, cfg: &PortForwardConfig) -> anyhow::Result<()> { + let (bind_addr, dst_addr) = (cfg.bind_addr, cfg.dst_addr); + let options = TcpListenOptions::port_forward(bind_addr); + let bind = options + .bind + .clone() + .with_context(self.socket_context.clone()); + let listener = self.host.bind_tcp(options.with_bind(bind)).await?; + let cancel = CancellationToken::new(); + self.cancel_tokens + .insert(cfg.clone(), cancel.clone().drop_guard()); + + let data_plane = self.data_plane.clone(); + let connections = Arc::new(std::sync::Mutex::new(JoinSet::new())); + connections.lock().unwrap().spawn(reap_joinset_background( + connections.clone(), + "TCP port-forward connections", + )); + self.tasks.lock().unwrap().spawn(async move { + loop { + let (incoming, source_addr) = select! { + biased; + _ = cancel.cancelled() => break, + result = listener.accept() => match result { + Ok(accepted) => accepted, + Err(error) => { + tracing::error!(?error, ?bind_addr, "port-forward accept failed"); + continue; + } + }, + }; + let data_plane = data_plane.clone(); + connections.lock().unwrap().spawn(async move { + let options = DataPlaneTcpConnectOptions::gateway( + Duration::from_secs(10), + TcpSocketPurpose::PortForward, + source_addr, + ); + let outgoing = match data_plane.connect_tcp(dst_addr, options).await { + Ok(stream) => stream, + Err(error) => { + tracing::error!(?error, ?dst_addr, "port-forward connect failed"); + return; + } + }; + copy_tcp(incoming, outgoing, dst_addr).await; + }); + } + }); + Ok(()) + } + + async fn add_udp_port_forward(&self, cfg: &PortForwardConfig) -> anyhow::Result<()> { + let (bind_addr, dst_addr) = (cfg.bind_addr, cfg.dst_addr); + let forward = cfg.clone(); + let socket = self + .host + .bind_udp( + UdpBindOptions::port_forward(bind_addr).with_context(self.socket_context.clone()), + ) + .await?; + let cancel = CancellationToken::new(); + self.cancel_tokens + .insert(cfg.clone(), cancel.clone().drop_guard()); + + let data_plane = self.data_plane.clone(); + let host = self.host.clone(); + let socket_context = self.socket_context.clone(); + let udp_clients = self.udp_clients.clone(); + let response_tasks = self.udp_response_tasks.clone(); + self.tasks.lock().unwrap().spawn(async move { + let adapter = UdpFlowFactory { + data_plane, + host, + socket_context, + }; + loop { + let mut buf = vec![0u8; 8192]; + let (len, client_addr) = select! { + biased; + _ = cancel.cancelled() => break, + result = socket.recv_from(&mut buf) => match result { + Ok(packet) => packet, + Err(error) => { + tracing::error!(?error, ?bind_addr, "UDP port-forward receive failed"); + continue; + } + }, + }; + let key = UdpClientKey { + client_addr, + forward: forward.clone(), + }; + let flow = match udp_clients.get(&key) { + Some(client) => client.clone(), + None => { + let flow = match adapter.open(dst_addr).await { + Ok(flow) => flow, + Err(error) => { + tracing::error!( + ?error, + ?dst_addr, + "open UDP data-plane flow failed" + ); + continue; + } + }; + let client = Arc::new(UdpClientInfo { + flow: flow.clone(), + last_active: AtomicCell::new(Instant::now()), + }); + udp_clients.insert(key.clone(), client.clone()); + + let inbound = socket.clone(); + let response_flow = flow.clone(); + let response_client = client_addr; + response_tasks.insert( + key.clone(), + AbortOnDropHandle::new(tokio::spawn(async move { + loop { + let mut buf = vec![0u8; 8192]; + match response_flow.recv_from(&mut buf).await { + Ok((len, remote)) => { + tracing::trace!( + ?remote, + ?response_client, + len, + "forwarding UDP data-plane response" + ); + if let Err(error) = + inbound.send_to(&buf[..len], response_client).await + { + tracing::error!(?error, "send UDP response failed"); + return; + } + } + Err(error) => { + tracing::error!(?error, "receive UDP response failed"); + return; + } + } + } + })), + ); + client + } + }; + flow.last_active.store(Instant::now()); + if let Err(error) = flow.flow.send_to(&buf[..len], dst_addr).await { + tracing::error!(?error, ?dst_addr, "send UDP data-plane packet failed"); + } + } + }); + + Ok(()) + } + + fn start_udp_reaper(&self) { + let udp_clients = self.udp_clients.clone(); + let response_tasks = self.udp_response_tasks.clone(); + self.tasks.lock().unwrap().spawn(async move { + loop { + tokio::time::sleep(Duration::from_secs(30)).await; + let now = Instant::now(); + udp_clients.retain(|_, client| { + now.duration_since(client.last_active.load()).as_secs() < 600 + }); + response_tasks.retain(|key, _| udp_clients.contains_key(key)); + udp_clients.shrink_to_fit(); + response_tasks.shrink_to_fit(); + } + }); + } + + async fn stop_inner(&self) { + self.started.store(false, Ordering::Release); + self.cancel_tokens.clear(); + self.udp_response_tasks.clear(); + self.udp_clients.clear(); + self.consumer_lease.lock().await.take(); + let mut tasks = { + let mut tasks = self.tasks.lock().unwrap(); + std::mem::replace(&mut *tasks, JoinSet::new()) + }; + tasks.shutdown().await; + } + + pub(crate) async fn stop(&self) { + let _operation = self.operation.lock().await; + self.stop_inner().await; + } +} + +struct UdpFlowFactory +where + H: VirtualTcpSocketFactory + VirtualTcpListenerFactory + VirtualUdpSocketFactory, +{ + data_plane: Arc>, + host: Arc, + socket_context: SocketContext, +} + +impl UdpFlowFactory +where + H: VirtualTcpSocketFactory + VirtualTcpListenerFactory + VirtualUdpSocketFactory, +{ + async fn open(&self, dst_addr: SocketAddr) -> anyhow::Result>> { + if self.data_plane.is_local_virtual_ip(dst_addr.ip()) { + let socket = self + .host + .bind_udp( + UdpBindOptions::port_lease("0.0.0.0:0".parse().unwrap()) + .with_context(self.socket_context.clone()), + ) + .await?; + Ok(Arc::new(PortForwardUdpFlow::Host(socket))) + } else { + let socket = self + .data_plane + .data_plane_udp_bind(0, Duration::from_secs(10)) + .await?; + Ok(Arc::new(PortForwardUdpFlow::DataPlane(Arc::new(socket)))) + } + } +} + +async fn copy_tcp(mut incoming: S, mut outgoing: DataPlaneTcpStream, dst_addr: SocketAddr) +where + S: VirtualTcpSocket, +{ + match tokio::io::copy_bidirectional(&mut incoming, &mut outgoing).await { + Ok((from_client, from_server)) => tracing::info!( + ?dst_addr, + from_client, + from_server, + "port-forward connection finished" + ), + Err(error) => tracing::error!(?error, ?dst_addr, "port-forward connection failed"), + } +} diff --git a/easytier-core/src/gateway/proxy/cidr_monitor.rs b/easytier-core/src/gateway/proxy/cidr_monitor.rs new file mode 100644 index 00000000..3e0b6747 --- /dev/null +++ b/easytier-core/src/gateway/proxy/cidr_monitor.rs @@ -0,0 +1,276 @@ +use std::{collections::BTreeSet, sync::Arc, time::Duration}; + +use cidr::Ipv4Cidr; +#[cfg(feature = "proxy-cidr-monitor")] +use tokio::sync::Mutex; +use tokio_util::task::AbortOnDropHandle; + +use crate::{ + config::runtime::{CoreInstanceRuntimeConfig, CoreRuntimeConfigStore}, + events::{CoreEvent, CoreEventSink}, + peers::peer_manager::PeerManagerCore, +}; + +#[derive(Clone, Debug, Default, PartialEq, Eq)] +pub(crate) struct ProxyCidrConfigSnapshot { + pub manual_routes: Option>, + pub vpn_portal_cidr: Option, +} + +impl From<&CoreInstanceRuntimeConfig> for ProxyCidrConfigSnapshot { + fn from(config: &CoreInstanceRuntimeConfig) -> Self { + Self { + manual_routes: config.services.manual_routes.clone(), + vpn_portal_cidr: config.peer.vpn_portal_cidr, + } + } +} + +#[cfg(feature = "proxy-cidr-monitor")] +pub(crate) struct ProxyCidrMonitorRuntime { + enabled: bool, + events: Arc, + task: Mutex>>, +} + +#[cfg(feature = "proxy-cidr-monitor")] +impl ProxyCidrMonitorRuntime { + pub(crate) fn new(enabled: bool, events: Arc) -> Self { + Self { + enabled, + events, + task: Mutex::new(None), + } + } + + pub(crate) fn is_enabled(&self) -> bool { + self.enabled + } + + pub(crate) async fn start( + &self, + peer_manager: &Arc, + runtime_config: CoreRuntimeConfigStore, + ) { + if !self.enabled { + return; + } + let mut task = self.task.lock().await; + if task.is_none() { + task.replace( + ProxyCidrMonitor::new(peer_manager, runtime_config, self.events.clone()).start(), + ); + } + } + + pub(crate) async fn stop(&self) { + self.task.lock().await.take(); + } +} + +#[derive(Clone, Debug, Default, PartialEq, Eq)] +pub struct ProxyCidrDiff { + pub current: BTreeSet, + pub added: Vec, + pub removed: Vec, +} + +pub(crate) fn resolve_proxy_cidrs( + mut peer_routes: BTreeSet, + config: ProxyCidrConfigSnapshot, +) -> BTreeSet { + if let Some(manual_routes) = config.manual_routes { + return manual_routes; + } + if let Some(vpn_portal_cidr) = config.vpn_portal_cidr { + peer_routes.insert(vpn_portal_cidr); + } + peer_routes +} + +pub(crate) fn diff_proxy_cidrs( + previous: &BTreeSet, + current: BTreeSet, +) -> ProxyCidrDiff { + let added = current.difference(previous).copied().collect(); + let removed = previous.difference(¤t).copied().collect(); + ProxyCidrDiff { + current, + added, + removed, + } +} + +pub(crate) async fn collect_proxy_cidrs( + peer_manager: &PeerManagerCore, + config: &CoreInstanceRuntimeConfig, +) -> BTreeSet { + let peer_routes = peer_manager.get_route().list_proxy_cidrs().await; + resolve_proxy_cidrs_from_runtime(peer_routes, config) +} + +fn resolve_proxy_cidrs_from_runtime( + peer_routes: BTreeSet, + config: &CoreInstanceRuntimeConfig, +) -> BTreeSet { + resolve_proxy_cidrs(peer_routes, config.into()) +} + +pub(crate) async fn collect_proxy_cidr_diff( + peer_manager: &PeerManagerCore, + runtime_config: &CoreRuntimeConfigStore, + previous: &BTreeSet, +) -> ProxyCidrDiff { + let config = runtime_config.snapshot(); + collect_proxy_cidr_diff_from_snapshot(peer_manager, config.as_ref(), previous).await +} + +async fn collect_proxy_cidr_diff_from_snapshot( + peer_manager: &PeerManagerCore, + config: &CoreInstanceRuntimeConfig, + previous: &BTreeSet, +) -> ProxyCidrDiff { + let current = collect_proxy_cidrs(peer_manager, config).await; + diff_proxy_cidrs(previous, current) +} + +#[cfg_attr(not(feature = "proxy-cidr-monitor"), allow(dead_code))] +pub(crate) struct ProxyCidrMonitor { + peer_manager: std::sync::Weak, + runtime_config: CoreRuntimeConfigStore, + events: Arc, +} + +#[cfg_attr(not(feature = "proxy-cidr-monitor"), allow(dead_code))] +impl ProxyCidrMonitor { + pub(crate) fn new( + peer_manager: &Arc, + runtime_config: CoreRuntimeConfigStore, + events: Arc, + ) -> Self { + Self { + peer_manager: Arc::downgrade(peer_manager), + runtime_config, + events, + } + } + + pub(crate) fn start(self) -> AbortOnDropHandle<()> { + AbortOnDropHandle::new(tokio::spawn(async move { + let mut current = BTreeSet::new(); + let mut last_update = None; + let mut last_runtime_config: Option> = None; + + loop { + crate::foundation::time::sleep(Duration::from_secs(1)).await; + let Some(peer_manager) = self.peer_manager.upgrade() else { + break; + }; + let update = peer_manager + .get_route() + .get_peer_info_last_update_time() + .await; + let runtime_config = self.runtime_config.snapshot(); + let runtime_config_changed = last_runtime_config + .as_ref() + .map(|previous| !Arc::ptr_eq(previous, &runtime_config)) + .unwrap_or(true); + if last_update == Some(update) && !runtime_config_changed { + continue; + } + last_update = Some(update); + last_runtime_config = Some(runtime_config.clone()); + + let diff = collect_proxy_cidr_diff_from_snapshot( + peer_manager.as_ref(), + runtime_config.as_ref(), + ¤t, + ) + .await; + current = diff.current; + if !diff.added.is_empty() || !diff.removed.is_empty() { + self.events.emit(CoreEvent::ProxyCidrsUpdated { + added: diff.added, + removed: diff.removed, + }); + } + } + })) + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::{ + config::peers::PeerRuntimeSnapshot, + config::runtime::{CoreInstanceRuntimeConfig, CoreRuntimeConfig}, + }; + + fn cidrs(values: &[&str]) -> BTreeSet { + values.iter().map(|value| value.parse().unwrap()).collect() + } + + #[test] + fn manual_routes_override_peer_and_vpn_routes() { + let resolved = resolve_proxy_cidrs( + cidrs(&["10.0.0.0/8"]), + ProxyCidrConfigSnapshot { + manual_routes: Some(cidrs(&["192.0.2.0/24"])), + vpn_portal_cidr: Some("198.51.100.0/24".parse().unwrap()), + }, + ); + assert_eq!(resolved, cidrs(&["192.0.2.0/24"])); + } + + #[test] + fn dynamic_routes_merge_vpn_and_report_ordered_diff() { + let current = resolve_proxy_cidrs( + cidrs(&["10.0.0.0/8"]), + ProxyCidrConfigSnapshot { + manual_routes: None, + vpn_portal_cidr: Some("192.0.2.0/24".parse().unwrap()), + }, + ); + let diff = diff_proxy_cidrs(&cidrs(&["10.0.0.0/8", "172.16.0.0/12"]), current); + + assert_eq!(diff.current, cidrs(&["10.0.0.0/8", "192.0.2.0/24"])); + assert_eq!(diff.added, vec!["192.0.2.0/24".parse().unwrap()]); + assert_eq!(diff.removed, vec!["172.16.0.0/12".parse().unwrap()]); + } + + #[test] + fn runtime_store_update_changes_the_monitor_config_snapshot() { + let initial_peer = PeerRuntimeSnapshot { + vpn_portal_cidr: Some("198.51.100.0/24".parse().unwrap()), + ..Default::default() + }; + let store = CoreRuntimeConfigStore::new( + CoreRuntimeConfig { + manual_routes: Some(cidrs(&["192.0.2.0/24"])), + ..Default::default() + }, + Arc::new(initial_peer), + ); + let initial = store.snapshot(); + + let updated_peer = PeerRuntimeSnapshot { + vpn_portal_cidr: Some("203.0.113.0/24".parse().unwrap()), + ..Default::default() + }; + store.replace(CoreInstanceRuntimeConfig { + services: CoreRuntimeConfig::default(), + peer: Arc::new(updated_peer), + }); + let updated = store.snapshot(); + + assert_eq!( + resolve_proxy_cidrs_from_runtime(cidrs(&["10.0.0.0/8"]), initial.as_ref()), + cidrs(&["192.0.2.0/24"]) + ); + assert_eq!( + resolve_proxy_cidrs_from_runtime(cidrs(&["10.0.0.0/8"]), updated.as_ref()), + cidrs(&["10.0.0.0/8", "203.0.113.0/24"]) + ); + } +} diff --git a/easytier-core/src/gateway/proxy/cidr_table.rs b/easytier-core/src/gateway/proxy/cidr_table.rs new file mode 100644 index 00000000..36fa06da --- /dev/null +++ b/easytier-core/src/gateway/proxy/cidr_table.rs @@ -0,0 +1,164 @@ +use std::net::IpAddr; +#[cfg(any(test, feature = "proxy-packet"))] +use std::net::Ipv4Addr; + +use parking_lot::RwLock; + +use crate::config::{IpPrefix, ProxyNetworkConfig}; + +#[derive(Clone, Debug, Eq, PartialEq)] +pub struct ProxyCidrRule { + pub cidr: cidr::Ipv4Cidr, + pub mapped_cidr: Option, +} + +#[derive(Clone, Debug, Default, Eq, PartialEq)] +pub struct ProxyCidrSnapshot { + pub rules: Vec, +} + +impl ProxyCidrSnapshot { + pub fn from_proxy_networks(networks: &[ProxyNetworkConfig]) -> Self { + Self { + rules: networks.iter().filter_map(proxy_cidr_rule).collect(), + } + } +} + +fn ipv4_cidr(prefix: &IpPrefix) -> Option { + let IpAddr::V4(address) = prefix.address else { + return None; + }; + cidr::Ipv4Cidr::new(address, prefix.prefix_len).ok() +} + +fn proxy_cidr_rule(config: &ProxyNetworkConfig) -> Option { + Some(ProxyCidrRule { + cidr: ipv4_cidr(&config.real)?, + mapped_cidr: config.mapped.as_ref().and_then(ipv4_cidr), + }) +} + +#[derive(Clone, Debug, Eq, PartialEq)] +struct ProxyCidrEntry { + real_cidr: cidr::Ipv4Cidr, + mapped_cidr: cidr::Ipv4Cidr, +} + +#[derive(Debug, Default)] +pub struct ProxyCidrTable { + entries: RwLock>, +} + +impl ProxyCidrTable { + pub fn new() -> Self { + Self::default() + } + + pub fn from_snapshot(snapshot: ProxyCidrSnapshot) -> Self { + let table = Self::new(); + table.update_snapshot(snapshot); + table + } + + pub fn update_snapshot(&self, snapshot: ProxyCidrSnapshot) { + let entries = snapshot + .rules + .into_iter() + .map(|rule| ProxyCidrEntry { + real_cidr: rule.cidr, + mapped_cidr: rule.mapped_cidr.unwrap_or(rule.cidr), + }) + .collect(); + *self.entries.write() = entries; + } + + #[cfg(feature = "proxy-packet")] + pub fn is_empty(&self) -> bool { + self.entries.read().is_empty() + } + + #[cfg(any(test, feature = "proxy-packet"))] + pub fn lookup_v4(&self, ipv4: Ipv4Addr) -> Option { + self.entries + .read() + .iter() + .find_map(|entry| entry.lookup_v4(ipv4)) + } +} + +impl ProxyCidrEntry { + #[cfg(any(test, feature = "proxy-packet"))] + fn lookup_v4(&self, ipv4: Ipv4Addr) -> Option { + if !self.mapped_cidr.contains(&ipv4) { + return None; + } + + if self.mapped_cidr == self.real_cidr { + return Some(ipv4); + } + + let origin_network_bits = self.real_cidr.first().address().to_bits(); + let network_mask = self.mapped_cidr.mask().to_bits(); + let converted_ip = (ipv4.to_bits() & !network_mask) | origin_network_bits; + Some(Ipv4Addr::from(converted_ip)) + } +} + +#[cfg(test)] +mod tests { + use crate::config::{IpPrefix, ProxyNetworkConfig}; + + use super::*; + + #[test] + fn lookup_returns_original_ip_for_unmapped_cidr() { + let table = ProxyCidrTable::from_snapshot(ProxyCidrSnapshot { + rules: vec![ProxyCidrRule { + cidr: "127.0.0.0/24".parse().unwrap(), + mapped_cidr: None, + }], + }); + + assert_eq!( + table.lookup_v4("127.0.0.42".parse().unwrap()), + Some("127.0.0.42".parse().unwrap()) + ); + assert_eq!(table.lookup_v4("127.0.1.42".parse().unwrap()), None); + } + + #[test] + fn lookup_converts_mapped_cidr_to_real_cidr() { + let table = ProxyCidrTable::from_snapshot(ProxyCidrSnapshot { + rules: vec![ProxyCidrRule { + cidr: "127.0.0.0/24".parse().unwrap(), + mapped_cidr: Some("10.10.10.0/24".parse().unwrap()), + }], + }); + + assert_eq!( + table.lookup_v4("10.10.10.42".parse().unwrap()), + Some("127.0.0.42".parse().unwrap()) + ); + } + + #[test] + fn snapshot_normalizes_proxy_network_config() { + let snapshot = ProxyCidrSnapshot::from_proxy_networks(&[ProxyNetworkConfig { + real: IpPrefix { + address: "192.0.2.0".parse().unwrap(), + prefix_len: 24, + }, + mapped: Some(IpPrefix { + address: "198.51.100.0".parse().unwrap(), + prefix_len: 24, + }), + }]); + let table = ProxyCidrTable::from_snapshot(snapshot); + + assert_eq!( + table.lookup_v4("198.51.100.42".parse().unwrap()), + Some("192.0.2.42".parse().unwrap()) + ); + } +} diff --git a/easytier-core/src/gateway/proxy/icmp_host.rs b/easytier-core/src/gateway/proxy/icmp_host.rs new file mode 100644 index 00000000..2a916ea0 --- /dev/null +++ b/easytier-core/src/gateway/proxy/icmp_host.rs @@ -0,0 +1,33 @@ +use std::{ + net::{IpAddr, Ipv4Addr}, + sync::Arc, +}; + +#[derive(Debug, thiserror::Error)] +pub enum ProxyRuntimeError { + #[error(transparent)] + Other(#[from] anyhow::Error), +} + +impl From for ProxyRuntimeError { + fn from(value: std::io::Error) -> Self { + Self::Other(value.into()) + } +} + +#[async_trait::async_trait] +pub trait IcmpProxySocket: Send + Sync + 'static { + async fn send(&self, destination: Ipv4Addr, packet: &[u8]) -> Result<(), ProxyRuntimeError>; + + async fn recv(&self) -> Result<(IpAddr, Vec), ProxyRuntimeError>; + + fn close(&self) {} +} + +#[async_trait::async_trait] +pub trait IcmpProxyHost: Send + Sync + 'static { + async fn open_icmp_v4( + &self, + context: crate::socket::SocketContext, + ) -> Result, ProxyRuntimeError>; +} diff --git a/easytier-core/src/gateway/proxy/icmp_proxy_engine.rs b/easytier-core/src/gateway/proxy/icmp_proxy_engine.rs new file mode 100644 index 00000000..a0e6f51b --- /dev/null +++ b/easytier-core/src/gateway/proxy/icmp_proxy_engine.rs @@ -0,0 +1,516 @@ +use std::{net::Ipv4Addr, sync::Arc, time::Duration}; + +use dashmap::DashMap; +use pnet_packet::{ + Packet, + icmp::{self, IcmpCode, IcmpTypes, MutableIcmpPacket, echo_reply::MutableEchoReplyPacket}, + ip::IpNextHeaderProtocols, + ipv4::Ipv4Packet, +}; +use quanta::Instant; + +use crate::packet::{PacketType, ZCPacket}; + +use super::{ + cidr_table::ProxyCidrTable, + ip_reassembler::{ + ComposeIpv4PacketArgs, IpProtocol, IpReassembler, SmolIpv4Packet, compose_ipv4_packet, + }, +}; + +#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)] +pub struct IcmpProxyContext { + pub virtual_ipv4: Option, + pub enable_exit_node: bool, + pub no_tun: bool, +} + +#[derive(Debug)] +pub enum IcmpProxyAction { + Pass, + SendToSocket { + destination: Ipv4Addr, + packet: Vec, + }, + SendToPeer(Vec), +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] +struct IcmpNatKey { + real_destination: Ipv4Addr, + identifier: u16, + sequence: u16, +} + +#[derive(Debug)] +struct IcmpNatEntry { + source_peer_id: u32, + local_peer_id: u32, + source_ip: Ipv4Addr, + mapped_destination: Ipv4Addr, + started_at: Instant, +} + +#[derive(Debug)] +pub struct IcmpProxyEngine { + cidr_table: Arc, + nat_table: DashMap, + reassembler: IpReassembler, +} + +impl IcmpProxyEngine { + pub fn new(cidr_table: Arc, fragment_timeout: Duration) -> Self { + Self { + cidr_table, + nat_table: DashMap::new(), + reassembler: IpReassembler::new(fragment_timeout), + } + } + + pub fn handle_peer_packet( + &self, + packet: &ZCPacket, + context: IcmpProxyContext, + ) -> IcmpProxyAction { + if self.cidr_table.is_empty() && !context.enable_exit_node && !context.no_tun { + return IcmpProxyAction::Pass; + } + let Some(virtual_ipv4) = context.virtual_ipv4 else { + return IcmpProxyAction::Pass; + }; + let Some(header) = packet.peer_manager_header() else { + return IcmpProxyAction::Pass; + }; + if header.packet_type != PacketType::Data as u8 || header.is_no_proxy() { + return IcmpProxyAction::Pass; + } + let Some(ipv4) = Ipv4Packet::new(packet.payload()) else { + return IcmpProxyAction::Pass; + }; + if ipv4.get_version() != 4 || ipv4.get_next_level_protocol() != IpNextHeaderProtocols::Icmp + { + return IcmpProxyAction::Pass; + } + + let mapped_destination = ipv4.get_destination(); + let real_destination = self.cidr_table.lookup_v4(mapped_destination); + let is_local_no_tun = context.no_tun && mapped_destination == virtual_ipv4; + if real_destination.is_none() && !header.is_exit_node() && !is_local_no_tun { + return IcmpProxyAction::Pass; + } + + let reassembled; + let smol_ipv4 = SmolIpv4Packet::new_unchecked(ipv4.packet()); + let request = if IpReassembler::is_packet_fragmented(&smol_ipv4) { + let Ok(smol_ipv4) = SmolIpv4Packet::new_checked(ipv4.packet()) else { + return IcmpProxyAction::Pass; + }; + reassembled = self.reassembler.add_fragment(&smol_ipv4); + let Some(reassembled) = reassembled.as_ref() else { + return IcmpProxyAction::Pass; + }; + let Some(request) = icmp::echo_request::EchoRequestPacket::new(reassembled) else { + return IcmpProxyAction::Pass; + }; + request + } else { + let Some(request) = icmp::echo_request::EchoRequestPacket::new(ipv4.payload()) else { + return IcmpProxyAction::Pass; + }; + request + }; + if request.get_icmp_type() != IcmpTypes::EchoRequest { + return IcmpProxyAction::Pass; + } + + if is_local_no_tun { + return self.local_reply( + mapped_destination, + ipv4.get_source(), + header.to_peer_id.get(), + header.from_peer_id.get(), + &request, + ); + } + + let real_destination = real_destination.unwrap_or(mapped_destination); + let key = IcmpNatKey { + real_destination, + identifier: request.get_identifier(), + sequence: request.get_sequence_number(), + }; + self.nat_table.insert( + key, + IcmpNatEntry { + source_peer_id: header.from_peer_id.get(), + local_peer_id: header.to_peer_id.get(), + source_ip: ipv4.get_source(), + mapped_destination, + started_at: Instant::now(), + }, + ); + + IcmpProxyAction::SendToSocket { + destination: real_destination, + packet: request.packet().to_vec(), + } + } + + pub fn handle_socket_response(&self, peer_ip: Ipv4Addr, packet: &mut [u8]) -> Vec { + let Some(ipv4) = Ipv4Packet::new(packet) else { + return Vec::new(); + }; + let Some(reply) = icmp::echo_reply::EchoReplyPacket::new(ipv4.payload()) else { + return Vec::new(); + }; + if reply.get_icmp_type() != IcmpTypes::EchoReply { + return Vec::new(); + } + let key = IcmpNatKey { + real_destination: peer_ip, + identifier: reply.get_identifier(), + sequence: reply.get_sequence_number(), + }; + let Some((_, entry)) = self.nat_table.remove(&key) else { + return Vec::new(); + }; + let Some(payload_len) = packet + .len() + .checked_sub(ipv4.get_header_length() as usize * 4) + else { + return Vec::new(); + }; + let ip_id = ipv4.get_identification(); + let mut responses = Vec::new(); + let _ = compose_ipv4_packet( + ComposeIpv4PacketArgs { + buf: packet, + src_v4: &entry.mapped_destination, + dst_v4: &entry.source_ip, + next_protocol: IpProtocol::Icmp, + payload_len, + payload_mtu: 1200, + ip_id, + }, + |buf| { + let mut packet = ZCPacket::new_with_payload(buf); + packet.fill_peer_manager_hdr( + entry.local_peer_id, + entry.source_peer_id, + PacketType::Data as u8, + ); + packet + .mut_peer_manager_header() + .expect("peer manager header") + .set_no_proxy(true); + responses.push(packet); + Ok(()) + }, + ); + responses + } + + pub fn remove_expired_entries(&self, max_age: Duration) { + self.nat_table + .retain(|_, entry| entry.started_at.elapsed() < max_age); + self.nat_table.shrink_to_fit(); + } + + pub fn remove_expired_fragments(&self) { + self.reassembler.remove_expired_packets(); + } + + fn local_reply( + &self, + source: Ipv4Addr, + destination: Ipv4Addr, + source_peer_id: u32, + destination_peer_id: u32, + request: &icmp::echo_request::EchoRequestPacket<'_>, + ) -> IcmpProxyAction { + let mut buffer = vec![0_u8; request.packet().len() + 20]; + let mut reply = MutableEchoReplyPacket::new(&mut buffer[20..]).unwrap(); + reply.set_icmp_type(IcmpTypes::EchoReply); + reply.set_icmp_code(IcmpCode::new(0)); + reply.set_identifier(request.get_identifier()); + reply.set_sequence_number(request.get_sequence_number()); + reply.set_payload(request.payload()); + let mut reply = MutableIcmpPacket::new(&mut buffer[20..]).unwrap(); + reply.set_checksum(icmp::checksum(&reply.to_immutable())); + + let payload_len = buffer.len() - 20; + let mut responses = Vec::new(); + let _ = compose_ipv4_packet( + ComposeIpv4PacketArgs { + buf: &mut buffer, + src_v4: &source, + dst_v4: &destination, + next_protocol: IpProtocol::Icmp, + payload_len, + payload_mtu: 1200, + ip_id: rand::random(), + }, + |buf| { + let mut packet = ZCPacket::new_with_payload(buf); + packet.fill_peer_manager_hdr( + source_peer_id, + destination_peer_id, + PacketType::Data as u8, + ); + responses.push(packet); + Ok(()) + }, + ); + IcmpProxyAction::SendToPeer(responses) + } +} + +#[cfg(test)] +mod tests { + use pnet_packet::{ + MutablePacket as _, + icmp::{MutableIcmpPacket, echo_request::MutableEchoRequestPacket}, + ipv4::{self, MutableIpv4Packet}, + }; + + use super::*; + use crate::gateway::proxy::cidr_table::{ProxyCidrRule, ProxyCidrSnapshot}; + + fn echo_request_with_payload( + source: Ipv4Addr, + destination: Ipv4Addr, + payload: &[u8], + ) -> ZCPacket { + let mut bytes = vec![0_u8; 20 + 8 + payload.len()]; + { + let mut request = MutableEchoRequestPacket::new(&mut bytes[20..]).unwrap(); + request.set_icmp_type(IcmpTypes::EchoRequest); + request.set_identifier(7); + request.set_sequence_number(11); + request.set_payload(payload); + let mut icmp = MutableIcmpPacket::new(&mut bytes[20..]).unwrap(); + icmp.set_checksum(icmp::checksum(&icmp.to_immutable())); + } + { + let packet_len = bytes.len() as u16; + let mut ipv4 = MutableIpv4Packet::new(&mut bytes).unwrap(); + ipv4.set_version(4); + ipv4.set_header_length(5); + ipv4.set_total_length(packet_len); + ipv4.set_ttl(64); + ipv4.set_next_level_protocol(IpNextHeaderProtocols::Icmp); + ipv4.set_source(source); + ipv4.set_destination(destination); + ipv4.set_checksum(ipv4::checksum(&ipv4.to_immutable())); + } + let mut packet = ZCPacket::new_with_payload(&bytes); + packet.fill_peer_manager_hdr(101, 202, PacketType::Data as u8); + packet + } + + fn echo_request(source: Ipv4Addr, destination: Ipv4Addr) -> ZCPacket { + echo_request_with_payload(source, destination, b"ping") + } + + fn engine(rule: Option) -> IcmpProxyEngine { + let table = ProxyCidrTable::from_snapshot(ProxyCidrSnapshot { + rules: rule.into_iter().collect(), + }); + IcmpProxyEngine::new(Arc::new(table), Duration::from_secs(10)) + } + + #[test] + fn inactive_proxy_passes_echo_request() { + let engine = engine(None); + let packet = echo_request("10.0.0.2".parse().unwrap(), "192.0.2.2".parse().unwrap()); + + assert!(matches!( + engine.handle_peer_packet( + &packet, + IcmpProxyContext { + virtual_ipv4: Some("10.0.0.1".parse().unwrap()), + ..Default::default() + } + ), + IcmpProxyAction::Pass + )); + } + + #[test] + fn no_tun_local_request_returns_echo_reply_to_origin_peer() { + let engine = engine(None); + let packet = echo_request("10.0.0.2".parse().unwrap(), "10.0.0.1".parse().unwrap()); + + let IcmpProxyAction::SendToPeer(replies) = engine.handle_peer_packet( + &packet, + IcmpProxyContext { + virtual_ipv4: Some("10.0.0.1".parse().unwrap()), + no_tun: true, + ..Default::default() + }, + ) else { + panic!("expected local reply"); + }; + let [reply] = replies.as_slice() else { + panic!("expected one local reply"); + }; + let header = reply.peer_manager_header().unwrap(); + assert_eq!(header.from_peer_id.get(), 202); + assert_eq!(header.to_peer_id.get(), 101); + let ipv4 = Ipv4Packet::new(reply.payload()).unwrap(); + assert_eq!(ipv4.get_source(), "10.0.0.1".parse::().unwrap()); + assert_eq!( + ipv4.get_destination(), + "10.0.0.2".parse::().unwrap() + ); + let reply = icmp::echo_reply::EchoReplyPacket::new(ipv4.payload()).unwrap(); + assert_eq!(reply.get_identifier(), 7); + assert_eq!(reply.get_sequence_number(), 11); + assert_eq!(reply.payload(), b"ping"); + } + + #[test] + fn mapped_request_and_socket_reply_round_trip() { + let engine = engine(Some(ProxyCidrRule { + cidr: "127.0.0.0/24".parse().unwrap(), + mapped_cidr: Some("10.10.10.0/24".parse().unwrap()), + })); + let packet = echo_request("10.0.0.2".parse().unwrap(), "10.10.10.42".parse().unwrap()); + + let IcmpProxyAction::SendToSocket { + destination, + packet: request, + } = engine.handle_peer_packet( + &packet, + IcmpProxyContext { + virtual_ipv4: Some("10.0.0.1".parse().unwrap()), + ..Default::default() + }, + ) + else { + panic!("expected socket request"); + }; + assert_eq!(destination, "127.0.0.42".parse::().unwrap()); + let request = icmp::echo_request::EchoRequestPacket::new(&request).unwrap(); + assert_eq!(request.payload(), b"ping"); + + let mut response = echo_request(destination, "10.0.0.1".parse().unwrap()) + .payload() + .to_vec(); + { + let mut ipv4 = MutableIpv4Packet::new(&mut response).unwrap(); + let mut reply = MutableEchoReplyPacket::new(ipv4.payload_mut()).unwrap(); + reply.set_icmp_type(IcmpTypes::EchoReply); + let mut icmp = MutableIcmpPacket::new(ipv4.payload_mut()).unwrap(); + icmp.set_checksum(icmp::checksum(&icmp.to_immutable())); + ipv4.set_source(destination); + ipv4.set_checksum(ipv4::checksum(&ipv4.to_immutable())); + } + let replies = engine.handle_socket_response(destination, &mut response); + let [reply] = replies.as_slice() else { + panic!("expected one socket reply"); + }; + let header = reply.peer_manager_header().unwrap(); + assert_eq!(header.from_peer_id.get(), 202); + assert_eq!(header.to_peer_id.get(), 101); + assert!(header.is_no_proxy()); + let ipv4 = Ipv4Packet::new(reply.payload()).unwrap(); + assert_eq!( + ipv4.get_source(), + "10.10.10.42".parse::().unwrap() + ); + assert_eq!( + ipv4.get_destination(), + "10.0.0.2".parse::().unwrap() + ); + } + + #[test] + fn large_local_reply_preserves_all_ipv4_fragments() { + let engine = engine(None); + let packet = echo_request_with_payload( + "10.0.0.2".parse().unwrap(), + "10.0.0.1".parse().unwrap(), + &[0; 2400], + ); + + let IcmpProxyAction::SendToPeer(replies) = engine.handle_peer_packet( + &packet, + IcmpProxyContext { + virtual_ipv4: Some("10.0.0.1".parse().unwrap()), + no_tun: true, + ..Default::default() + }, + ) else { + panic!("expected local replies"); + }; + + assert_eq!(replies.len(), 3); + assert!(replies.iter().all(|packet| { + let header = packet.peer_manager_header().unwrap(); + header.from_peer_id.get() == 202 && header.to_peer_id.get() == 101 + })); + } + + #[test] + fn large_socket_reply_preserves_all_ipv4_fragments() { + let engine = engine(Some(ProxyCidrRule { + cidr: "127.0.0.0/24".parse().unwrap(), + mapped_cidr: Some("10.10.10.0/24".parse().unwrap()), + })); + let destination = "127.0.0.42".parse().unwrap(); + let packet = echo_request_with_payload( + "10.0.0.2".parse().unwrap(), + "10.10.10.42".parse().unwrap(), + &[0; 2400], + ); + assert!(matches!( + engine.handle_peer_packet( + &packet, + IcmpProxyContext { + virtual_ipv4: Some("10.0.0.1".parse().unwrap()), + ..Default::default() + }, + ), + IcmpProxyAction::SendToSocket { .. } + )); + assert!(engine.nat_table.contains_key(&IcmpNatKey { + real_destination: destination, + identifier: 7, + sequence: 11, + })); + + let mut response = + echo_request_with_payload(destination, "10.0.0.1".parse().unwrap(), &[0; 2400]) + .payload() + .to_vec(); + { + let mut ipv4 = MutableIpv4Packet::new(&mut response).unwrap(); + let mut reply = MutableEchoReplyPacket::new(ipv4.payload_mut()).unwrap(); + reply.set_icmp_type(IcmpTypes::EchoReply); + let mut icmp = MutableIcmpPacket::new(ipv4.payload_mut()).unwrap(); + icmp.set_checksum(icmp::checksum(&icmp.to_immutable())); + ipv4.set_source(destination); + // Raw sockets may return a buffer with bytes beyond the IPv4 total + // length. The native implementation composes from the received + // buffer length, so keep that case covered without changing the + // existing in-place composer in this refactor. + ipv4.set_total_length(1220); + ipv4.set_checksum(ipv4::checksum(&ipv4.to_immutable())); + } + let ipv4 = Ipv4Packet::new(&response).unwrap(); + let echo_reply = icmp::echo_reply::EchoReplyPacket::new(ipv4.payload()).unwrap(); + assert_eq!(echo_reply.get_icmp_type(), IcmpTypes::EchoReply); + assert_eq!(echo_reply.get_identifier(), 7); + assert_eq!(echo_reply.get_sequence_number(), 11); + + let replies = engine.handle_socket_response(destination, &mut response); + assert_eq!(replies.len(), 3); + assert!(replies.iter().all(|packet| { + let header = packet.peer_manager_header().unwrap(); + header.from_peer_id.get() == 202 + && header.to_peer_id.get() == 101 + && header.is_no_proxy() + })); + } +} diff --git a/easytier-core/src/gateway/proxy/icmp_proxy_service.rs b/easytier-core/src/gateway/proxy/icmp_proxy_service.rs new file mode 100644 index 00000000..1c167f4c --- /dev/null +++ b/easytier-core/src/gateway/proxy/icmp_proxy_service.rs @@ -0,0 +1,337 @@ +use std::{ + net::{IpAddr, Ipv4Addr}, + sync::{ + Arc, Weak, + atomic::{AtomicBool, Ordering}, + }, + time::Duration, +}; + +use tokio::{ + sync::mpsc::{UnboundedReceiver, UnboundedSender, unbounded_channel}, + task::JoinSet, +}; + +use crate::{ + packet::ZCPacket, + peers::{ + PeerPacketFilter, + peer_manager::{PeerManagerCore, PipelineRegistrationGuard}, + }, +}; + +use super::{ + cidr_table::ProxyCidrTable, + icmp_proxy_engine::{IcmpProxyAction, IcmpProxyContext, IcmpProxyEngine}, + traits::{IcmpProxyRuntime, IcmpProxySocket, ProxyRuntimeError}, +}; + +async fn start_icmp_runtime( + runtime: &R, + no_tun: bool, +) -> Result>, ProxyRuntimeError> { + match runtime.start_icmp().await { + Ok(socket) => Ok(Some(socket)), + Err(err) => { + runtime.stop_icmp(); + if !no_tun { + return Err(err); + } + tracing::warn!(?err, "start ICMP runtime failed without TUN"); + Ok(None) + } + } +} + +pub struct IcmpProxyService { + peer_manager: Arc, + runtime: Arc, + engine: Arc, + response_tx: UnboundedSender, + response_rx: std::sync::Mutex>>, + pipeline_guard: std::sync::Mutex>, + tasks: std::sync::Mutex>, + socket: std::sync::Mutex>>, + runtime_started: AtomicBool, + started: AtomicBool, +} + +impl IcmpProxyService { + pub fn new( + peer_manager: Arc, + runtime: Arc, + cidr_table: Arc, + fragment_timeout: Duration, + ) -> Arc { + let (response_tx, response_rx) = unbounded_channel(); + Arc::new(Self { + peer_manager, + runtime, + engine: Arc::new(IcmpProxyEngine::new(cidr_table, fragment_timeout)), + response_tx, + response_rx: std::sync::Mutex::new(Some(response_rx)), + pipeline_guard: std::sync::Mutex::new(None), + tasks: std::sync::Mutex::new(JoinSet::new()), + socket: std::sync::Mutex::new(None), + runtime_started: AtomicBool::new(false), + started: AtomicBool::new(false), + }) + } + + pub async fn start(self: &Arc) -> Result<(), ProxyRuntimeError> { + if self.started.swap(true, Ordering::AcqRel) { + return Ok(()); + } + + let snapshot = self.runtime.proxy_runtime_snapshot(); + let socket = match start_icmp_runtime(self.runtime.as_ref(), snapshot.no_tun).await { + Ok(socket) => socket, + Err(err) => { + self.started.store(false, Ordering::Release); + return Err(err); + } + }; + let runtime_started = socket.is_some(); + self.runtime_started + .store(runtime_started, Ordering::Release); + self.socket.lock().unwrap().clone_from(&socket); + + if let Some(mut response_rx) = self.response_rx.lock().unwrap().take() { + let peer_manager = self.peer_manager.clone(); + let latency_first = snapshot.latency_first; + self.tasks.lock().unwrap().spawn(async move { + while let Some(mut packet) = response_rx.recv().await { + let Some(header) = packet.mut_peer_manager_header() else { + continue; + }; + header.set_latency_first(latency_first); + let to_peer_id = header.to_peer_id.into(); + if let Err(err) = peer_manager.send_msg_for_proxy(packet, to_peer_id).await { + tracing::error!(?err, "send ICMP proxy response to peer failed"); + } + } + }); + } + + if let Some(socket) = socket { + let service = Arc::downgrade(self); + self.tasks.lock().unwrap().spawn(async move { + loop { + let recv_result = socket.recv().await; + let (peer_ip, mut packet) = match recv_result { + Ok(packet) => packet, + Err(err) => { + tracing::error!(?err, "receive ICMP packet failed"); + continue; + } + }; + if packet.is_empty() { + tracing::error!("received empty ICMP packet"); + break; + } + let IpAddr::V4(peer_ip) = peer_ip else { + continue; + }; + let Some(service) = service.upgrade() else { + break; + }; + service.handle_socket_response(peer_ip, &mut packet); + } + }); + } + + let service = Arc::downgrade(self); + self.tasks.lock().unwrap().spawn(async move { + loop { + crate::foundation::time::sleep(Duration::from_secs(1)).await; + let Some(service) = service.upgrade() else { + break; + }; + service.engine.remove_expired_fragments(); + } + }); + + let guard = self + .peer_manager + .add_managed_packet_process_pipeline(Box::new(IcmpProxyServiceFilter { + service: Arc::downgrade(self), + })) + .await; + self.pipeline_guard.lock().unwrap().replace(guard); + + let service = Arc::downgrade(self); + self.tasks.lock().unwrap().spawn(async move { + loop { + crate::foundation::time::sleep(Duration::from_secs(1)).await; + let Some(service) = service.upgrade() else { + break; + }; + service + .engine + .remove_expired_entries(Duration::from_secs(20)); + } + }); + + Ok(()) + } + + pub fn stop(&self) { + if !self.started.swap(false, Ordering::AcqRel) { + return; + } + if let Some(guard) = self.pipeline_guard.lock().unwrap().take() { + guard.close(); + } + self.tasks.lock().unwrap().abort_all(); + if self.runtime_started.swap(false, Ordering::AcqRel) { + self.runtime.stop_icmp(); + } + self.socket.lock().unwrap().take(); + } + + async fn handle_peer_packet(self: Arc, packet: ZCPacket) -> Option { + let snapshot = self.runtime.proxy_runtime_snapshot(); + match self.engine.handle_peer_packet( + &packet, + IcmpProxyContext { + virtual_ipv4: snapshot.virtual_ipv4, + enable_exit_node: snapshot.enable_exit_node, + no_tun: snapshot.no_tun, + }, + ) { + IcmpProxyAction::Pass => Some(packet), + IcmpProxyAction::SendToSocket { + destination, + packet: request, + } => { + let socket = self.socket.lock().unwrap().clone(); + match socket { + Some(socket) => { + if let Err(err) = socket.send(destination, &request).await { + tracing::error!(?err, "send ICMP packet through runtime failed"); + } + } + None => tracing::error!("send ICMP packet without a runtime socket"), + } + None + } + IcmpProxyAction::SendToPeer(packets) => { + for packet in packets { + if let Err(err) = self.response_tx.send(packet) { + tracing::error!(?err, "queue local ICMP response failed"); + } + } + None + } + } + } +} + +impl IcmpProxyService { + fn handle_socket_response(&self, peer_ip: Ipv4Addr, packet: &mut [u8]) { + for response in self.engine.handle_socket_response(peer_ip, packet) { + if let Err(err) = self.response_tx.send(response) { + tracing::error!(?err, "queue ICMP socket response failed"); + } + } + } +} + +impl Drop for IcmpProxyService { + fn drop(&mut self) { + self.stop(); + } +} + +struct IcmpProxyServiceFilter { + service: Weak>, +} + +#[async_trait::async_trait] +impl PeerPacketFilter for IcmpProxyServiceFilter { + async fn try_process_packet_from_peer(&self, packet: ZCPacket) -> Option { + let Some(service) = self.service.upgrade() else { + return Some(packet); + }; + service.handle_peer_packet(packet).await + } +} + +#[cfg(test)] +mod tests { + use std::{ + net::{IpAddr, Ipv4Addr}, + sync::{ + Arc, + atomic::{AtomicUsize, Ordering}, + }, + }; + + use super::*; + use crate::gateway::proxy::traits::{IcmpProxySocket, ProxyRuntimeInfo, ProxyRuntimeSnapshot}; + + #[derive(Default)] + struct PartialStartRuntime { + stop_count: AtomicUsize, + } + + impl ProxyRuntimeInfo for PartialStartRuntime { + fn proxy_runtime_snapshot(&self) -> ProxyRuntimeSnapshot { + ProxyRuntimeSnapshot::default() + } + + fn is_ip_local_virtual_ip(&self, _ip: &IpAddr) -> bool { + false + } + } + + struct NoopIcmpSocket; + + #[async_trait::async_trait] + impl IcmpProxySocket for NoopIcmpSocket { + async fn send( + &self, + _destination: Ipv4Addr, + _packet: &[u8], + ) -> Result<(), ProxyRuntimeError> { + Ok(()) + } + + async fn recv(&self) -> Result<(IpAddr, Vec), ProxyRuntimeError> { + Err(std::io::Error::other("unused test socket").into()) + } + } + + #[async_trait::async_trait] + impl IcmpProxyRuntime for PartialStartRuntime { + type Socket = NoopIcmpSocket; + + async fn start_icmp(&self) -> Result, ProxyRuntimeError> { + Err(std::io::Error::other("partial start").into()) + } + + fn stop_icmp(&self) { + self.stop_count.fetch_add(1, Ordering::AcqRel); + } + } + + #[tokio::test] + async fn failed_runtime_start_rolls_back_partial_resources() { + let runtime = PartialStartRuntime::default(); + + assert!(start_icmp_runtime(&runtime, false).await.is_err()); + assert_eq!(runtime.stop_count.load(Ordering::Acquire), 1); + } + + #[tokio::test] + async fn no_tun_suppresses_start_error_after_rollback() { + let runtime = PartialStartRuntime::default(); + + let socket = start_icmp_runtime(&runtime, true).await.unwrap(); + assert!(socket.is_none()); + if socket.is_some() { + runtime.stop_icmp(); + } + assert_eq!(runtime.stop_count.load(Ordering::Acquire), 1); + } +} diff --git a/easytier-core/src/gateway/proxy/ip_reassembler.rs b/easytier-core/src/gateway/proxy/ip_reassembler.rs new file mode 100644 index 00000000..875b60d5 --- /dev/null +++ b/easytier-core/src/gateway/proxy/ip_reassembler.rs @@ -0,0 +1,321 @@ +use std::{ + net::Ipv4Addr, + time::{Duration, Instant}, +}; + +use dashmap::DashMap; +use smoltcp::wire::Ipv4Packet; +pub use smoltcp::wire::{IpProtocol, Ipv4Packet as SmolIpv4Packet}; + +#[derive(Debug, Hash, PartialEq, Eq, Clone)] +struct IpReassemblerKey { + source: Ipv4Addr, + destination: Ipv4Addr, + id: u16, +} + +#[derive(Debug, Clone)] +struct IpFragment { + offset: u16, + data: Vec, +} + +#[derive(Debug)] +struct IpPacket { + total_length: Option, + fragments: Vec, +} + +impl IpPacket { + fn new() -> Self { + Self { + total_length: None, + fragments: Vec::new(), + } + } + + fn add_fragment(&mut self, fragment: IpFragment) { + for existing in &self.fragments { + let existing_end = existing.offset + existing.data.len() as u16; + let fragment_end = fragment.offset + fragment.data.len() as u16; + if existing.offset <= fragment.offset && fragment.offset < existing_end { + tracing::trace!( + existing_offset = existing.offset, + fragment_offset = fragment.offset, + existing_len = existing.data.len(), + fragment_len = fragment.data.len(), + "fragment overlap" + ); + return; + } + if fragment.offset <= existing.offset && existing.offset < fragment_end { + tracing::trace!( + existing_offset = existing.offset, + fragment_offset = fragment.offset, + existing_len = existing.data.len(), + fragment_len = fragment.data.len(), + "fragment overlap" + ); + return; + } + } + self.fragments.push(fragment); + } + + fn set_total_length(&mut self, total_length: u16) { + self.total_length = Some(total_length); + } + + fn assemble(&mut self) -> Option> { + let total_length = self.total_length?; + self.fragments.sort_by_key(|fragment| fragment.offset); + + let mut offset = 0; + let mut ret = Vec::with_capacity(total_length as usize); + for fragment in &self.fragments { + if fragment.offset != offset { + return None; + } + ret.extend_from_slice(&fragment.data); + offset += fragment.data.len() as u16; + } + + (offset == total_length).then_some(ret) + } +} + +impl + ?Sized> From<&Ipv4Packet<&T>> for IpFragment { + fn from(packet: &Ipv4Packet<&T>) -> Self { + Self { + offset: packet.frag_offset(), + data: packet.payload().to_vec(), + } + } +} + +#[derive(Debug)] +struct IpReassemblerValue { + packet: IpPacket, + timestamp: Instant, +} + +#[derive(Debug)] +pub struct IpReassembler { + packets: DashMap, + timeout: Duration, +} + +impl IpReassembler { + pub fn new(timeout: Duration) -> Self { + Self { + packets: DashMap::new(), + timeout, + } + } + + pub fn is_packet_fragmented>(packet: &Ipv4Packet) -> bool { + packet.frag_offset() != 0 || packet.more_frags() + } + + pub fn is_last_fragment>(packet: &Ipv4Packet) -> bool { + !packet.more_frags() + } + + pub fn add_fragment + ?Sized>( + &self, + packet: &Ipv4Packet<&T>, + ) -> Option> { + let total_length = packet.total_len() - packet.header_len() as u16; + if total_length != packet.payload().len() as u16 { + tracing::trace!( + ?total_length, + payload_len = ?packet.payload().len(), + "unexpected total length", + ); + return None; + } + + let key = IpReassemblerKey { + source: packet.src_addr(), + destination: packet.dst_addr(), + id: packet.ident(), + }; + let fragment: IpFragment = packet.into(); + + tracing::trace!(?key, offset = fragment.offset, total_length, "add fragment"); + + let mut entry = self.packets.entry(key.clone()).or_insert_with(|| { + let packet = IpPacket::new(); + let timestamp = Instant::now(); + IpReassemblerValue { packet, timestamp } + }); + let value_mut = entry.value_mut(); + + if Self::is_last_fragment(packet) { + value_mut + .packet + .set_total_length(total_length + fragment.offset); + } + + value_mut.packet.add_fragment(fragment); + if let Some(data) = value_mut.packet.assemble() { + drop(entry); + self.packets.remove(&key); + Some(data) + } else { + value_mut.timestamp = Instant::now(); + None + } + } + + pub fn remove_expired_packets(&self) { + let timeout = self.timeout; + self.packets + .retain(|_, value| value.timestamp.elapsed() <= timeout); + self.packets.shrink_to_fit(); + } +} + +pub struct ComposeIpv4PacketArgs<'a> { + pub buf: &'a mut [u8], + pub src_v4: &'a Ipv4Addr, + pub dst_v4: &'a Ipv4Addr, + pub next_protocol: IpProtocol, + pub payload_len: usize, + pub payload_mtu: usize, + pub ip_id: u16, +} + +pub fn compose_ipv4_packet(args: ComposeIpv4PacketArgs, mut cb: F) -> anyhow::Result<()> +where + F: FnMut(&[u8]) -> anyhow::Result<()>, +{ + let total_pieces = args.payload_len.div_ceil(args.payload_mtu); + let mut buf_offset = 0; + let mut fragment_offset = 0; + let mut cur_piece = 0; + while fragment_offset < args.payload_len { + let next_fragment_offset = + std::cmp::min(fragment_offset + args.payload_mtu, args.payload_len); + let fragment_len = next_fragment_offset - fragment_offset; + let packet_len = fragment_len + smoltcp::wire::IPV4_HEADER_LEN; + let mut ipv4_packet = + Ipv4Packet::new_unchecked(&mut args.buf[buf_offset..buf_offset + packet_len]); + ipv4_packet.set_version(4); + ipv4_packet.set_header_len(smoltcp::wire::IPV4_HEADER_LEN as u8); + ipv4_packet.set_total_len(packet_len as u16); + ipv4_packet.set_ident(args.ip_id); + ipv4_packet.clear_flags(); + if total_pieces > 1 { + ipv4_packet.set_more_frags(cur_piece != total_pieces - 1); + ipv4_packet.set_frag_offset(fragment_offset as u16); + } else { + ipv4_packet.set_dont_frag(true); + ipv4_packet.set_frag_offset(0); + } + ipv4_packet.set_dscp(0); + ipv4_packet.set_ecn(0); + ipv4_packet.set_hop_limit(32); + ipv4_packet.set_src_addr(*args.src_v4); + ipv4_packet.set_dst_addr(*args.dst_v4); + ipv4_packet.set_next_header(args.next_protocol); + ipv4_packet.fill_checksum(); + + tracing::trace!(?ipv4_packet, "proxy ipv4 packet composed"); + + cb(ipv4_packet.as_ref())?; + + buf_offset += next_fragment_offset - fragment_offset; + fragment_offset = next_fragment_offset; + cur_piece += 1; + } + Ok(()) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn compose_ipv4_packet_fragments_arbitrary_payload() { + let payload_len = 1932; + let payload_mtu = 1256; + let mut buffer = vec![0xff; smoltcp::wire::IPV4_HEADER_LEN + payload_len]; + let expected_payload = buffer[smoltcp::wire::IPV4_HEADER_LEN..].to_vec(); + let mut fragments = Vec::new(); + + compose_ipv4_packet( + ComposeIpv4PacketArgs { + buf: &mut buffer, + src_v4: &Ipv4Addr::new(10, 1, 2, 3), + dst_v4: &Ipv4Addr::new(10, 1, 2, 4), + next_protocol: IpProtocol::Udp, + payload_len, + payload_mtu, + ip_id: 42, + }, + |fragment| { + fragments.push(fragment.to_vec()); + Ok(()) + }, + ) + .unwrap(); + + assert_eq!(fragments.len(), 2); + let first = Ipv4Packet::new_checked(fragments[0].as_slice()).unwrap(); + let second = Ipv4Packet::new_checked(fragments[1].as_slice()).unwrap(); + assert!(first.more_frags()); + assert_eq!(first.frag_offset(), 0); + assert!(!second.more_frags()); + assert_eq!(second.frag_offset(), payload_mtu as u16); + + let reassembled = first + .payload() + .iter() + .chain(second.payload()) + .copied() + .collect::>(); + assert_eq!(reassembled, expected_payload); + } + + #[test] + fn reassembler() { + let raw_packets = [ + vec![ + 0x45, 0x00, 0x00, 0x1c, 0x1c, 0x46, 0x20, 0x01, 0x40, 0x06, 0xb1, 0xe6, 0xc0, 0xa8, + 0x00, 0x01, 0xc0, 0xa8, 0x00, 0x02, 0x04, 0x05, 0x06, 0x07, 0x04, 0x05, 0x06, 0x07, + ], + vec![ + 0x45, 0x00, 0x00, 0x1c, 0x1c, 0x46, 0x00, 0x02, 0x40, 0x06, 0xb1, 0xe6, 0xc0, 0xa8, + 0x00, 0x01, 0xc0, 0xa8, 0x00, 0x02, 0x08, 0x09, 0x0a, 0x0b, 0x04, 0x05, 0x06, 0x07, + ], + vec![ + 0x45, 0x00, 0x00, 0x1c, 0x1c, 0x46, 0x20, 0x00, 0x40, 0x06, 0xb1, 0xe6, 0xc0, 0xa8, + 0x00, 0x01, 0xc0, 0xa8, 0x00, 0x02, 0x00, 0x01, 0x02, 0x03, 0x04, 0x05, 0x06, 0x07, + ], + vec![ + 0x45, 0x00, 0x00, 0x1c, 0x1c, 0x47, 0x20, 0x00, 0x40, 0x06, 0xb1, 0xe6, 0xc0, 0xa8, + 0x00, 0x01, 0xc0, 0xa8, 0x00, 0x02, 0x00, 0x01, 0x02, 0x03, 0x04, 0x05, 0x06, 0x07, + ], + ]; + + let reassembler = IpReassembler::new(Duration::from_secs(1)); + + for (idx, raw_packet) in raw_packets.iter().enumerate() { + let packet = Ipv4Packet::new_checked(raw_packet.as_slice()).unwrap(); + let ret = reassembler.add_fragment(&packet); + if idx != 2 { + assert!(ret.is_none()); + } else { + assert!(ret.is_some()); + } + } + + reassembler.remove_expired_packets(); + assert_eq!(1, reassembler.packets.len()); + + std::thread::sleep(Duration::from_secs(2)); + reassembler.remove_expired_packets(); + assert_eq!(0, reassembler.packets.len()); + } +} diff --git a/easytier-core/src/gateway/proxy/mod.rs b/easytier-core/src/gateway/proxy/mod.rs new file mode 100644 index 00000000..a3135abb --- /dev/null +++ b/easytier-core/src/gateway/proxy/mod.rs @@ -0,0 +1,34 @@ +pub mod cidr_monitor; +pub(crate) mod cidr_table; +pub mod icmp_host; +#[cfg(feature = "proxy-packet")] +pub(crate) mod proxy_acl; +#[cfg(feature = "proxy-packet")] +pub(crate) mod service; +#[cfg(feature = "proxy-packet")] +pub mod traits; +#[cfg_attr(not(feature = "wrapped-transport"), allow(dead_code))] +pub mod wrapped_transport; + +#[cfg(feature = "proxy-packet")] +pub(crate) mod icmp_proxy_engine; +#[cfg(feature = "proxy-packet")] +pub(crate) mod icmp_proxy_service; +#[cfg(feature = "proxy-packet")] +pub(crate) mod ip_reassembler; +#[cfg(feature = "proxy-packet")] +pub mod tcp_proxy_engine; +#[cfg(feature = "proxy-packet")] +pub(crate) mod tcp_proxy_service; +#[cfg(feature = "proxy-packet")] +pub(crate) mod tcp_socket_connector; +#[cfg(feature = "proxy-packet")] +pub(crate) mod udp_proxy_engine; +#[cfg(feature = "proxy-packet")] +pub(crate) mod udp_proxy_service; +#[cfg(feature = "proxy-packet")] +pub(crate) mod udp_socket_runtime; +#[cfg(feature = "proxy-packet")] +pub(crate) mod wrapped_tcp_proxy; +#[cfg(feature = "proxy-packet")] +pub(crate) mod wrapped_transport_destination; diff --git a/easytier-core/src/gateway/proxy/proxy_acl.rs b/easytier-core/src/gateway/proxy/proxy_acl.rs new file mode 100644 index 00000000..32725d1f --- /dev/null +++ b/easytier-core/src/gateway/proxy/proxy_acl.rs @@ -0,0 +1,49 @@ +use std::sync::Arc; + +use easytier_proto::acl::{Action, ChainType}; +use tokio::io::{AsyncRead, AsyncWrite, copy_bidirectional}; +use tokio_util::io::InspectReader; + +use crate::peers::acl::{filter::AclFilter, processor::PacketInfo}; + +#[derive(Clone)] +pub struct ProxyAclHandler { + pub acl_filter: Arc, + pub packet_info: PacketInfo, + pub chain_type: ChainType, +} + +impl ProxyAclHandler { + pub fn handle_packet(&self, buf: &[u8]) -> anyhow::Result<()> { + self.handle_packet_size(buf.len()) + } + + pub fn handle_packet_size(&self, packet_size: usize) -> anyhow::Result<()> { + let mut packet_info = self.packet_info.clone(); + packet_info.packet_size = packet_size; + let processor = self.acl_filter.get_processor(); + let ret = processor.process_packet(&packet_info, self.chain_type); + self.acl_filter + .handle_acl_result(&ret, &packet_info, self.chain_type, &processor); + if !matches!(ret.action, Action::Allow) { + anyhow::bail!("acl denied"); + } + + Ok(()) + } + + pub async fn copy_bidirection_with_acl( + &self, + src: impl AsyncRead + AsyncWrite + Unpin, + mut dst: impl AsyncRead + AsyncWrite + Unpin, + ) -> anyhow::Result<()> { + let (src_reader, src_writer) = tokio::io::split(src); + let src_reader = InspectReader::new(src_reader, |buf| { + let _ = self.handle_packet(buf); + }); + let mut src = tokio::io::join(src_reader, src_writer); + + copy_bidirectional(&mut src, &mut dst).await?; + Ok(()) + } +} diff --git a/easytier-core/src/gateway/proxy/service.rs b/easytier-core/src/gateway/proxy/service.rs new file mode 100644 index 00000000..126cf7d8 --- /dev/null +++ b/easytier-core/src/gateway/proxy/service.rs @@ -0,0 +1,526 @@ +use std::{ + net::{IpAddr, Ipv4Addr, SocketAddr}, + sync::{ + Arc, + atomic::{AtomicBool, Ordering}, + }, + time::Duration, +}; + +use cidr::Ipv4Inet; +use tokio::sync::Mutex; + +use crate::{ + config::IpPrefix, + config::runtime::{CoreInstanceRuntimeConfig, CoreRuntimeConfigStore}, + connectivity::{ + direct::DirectConnectorHost, hole_punch::tcp::TcpHolePunchHost, protocol::protocol_uses_udp, + }, + foundation::stats::{LabelSet, LabelType, MetricName, StatsManager}, + listener::RunningListenerRegistry, + peers::peer_manager::PeerManagerCore, + process_runtime::ProtectedTcpPortRegistry, + socket::{IpVersion, SocketContext, udp::UdpBindOptions}, +}; + +use super::{ + cidr_table::ProxyCidrTable, + icmp_proxy_service::IcmpProxyService, + tcp_proxy_engine::TcpNatEntrySnapshot, + tcp_proxy_service::TcpProxyService, + tcp_socket_connector::TcpSocketProxyConnector, + traits::{ + IcmpProxyHost, IcmpProxyRuntime, IcmpProxySocket, ProxyRuntimeError, ProxyRuntimeInfo, + ProxyRuntimeSnapshot, TcpProxyConnectContext, TcpProxyRuntime, UdpProxyPolicy, + WrappedTcpDestinationRuntime, + }, + udp_proxy_service::UdpProxyService, + udp_socket_runtime::UdpSocketProxyRuntime, +}; + +const UDP_PROXY_SOCKET_IDLE_TIMEOUT: Duration = Duration::from_secs(120); +const PROXY_FRAGMENT_TIMEOUT: Duration = Duration::from_secs(10); + +fn udp_proxy_bind_options(context: SocketContext) -> UdpBindOptions { + UdpBindOptions::proxy_nat().with_context(context.with_ip_version(IpVersion::V4)) +} + +fn ipv4_inet(prefix: &IpPrefix) -> Option { + let IpAddr::V4(address) = prefix.address else { + return None; + }; + Ipv4Inet::new(address, prefix.prefix_len).ok() +} + +fn smoltcp_proxy_inet() -> Ipv4Inet { + Ipv4Inet::new(Ipv4Addr::new(192, 88, 99, 254), 24) + .expect("smoltcp proxy address must be a valid IPv4 interface") +} + +fn runtime_snapshot( + config: &CoreInstanceRuntimeConfig, + smoltcp_enabled: bool, +) -> ProxyRuntimeSnapshot { + let virtual_inet = config + .peer + .runtime + .core + .routes + .ipv4 + .as_ref() + .and_then(ipv4_inet); + ProxyRuntimeSnapshot { + local_inet: smoltcp_enabled.then(smoltcp_proxy_inet).or(virtual_inet), + virtual_ipv4: virtual_inet.map(|inet| inet.address()), + no_tun: config.services.proxy.no_tun, + enable_exit_node: config.services.proxy.enable_exit_node, + smoltcp_enabled, + latency_first: config.peer.flags.latency_first && !config.peer.flags.p2p_only, + } +} + +pub(crate) struct CoreProxyRuntime +where + H: DirectConnectorHost, +{ + peer_manager: Arc, + host: Arc, + protected_tcp_ports: Arc, + running_listeners: Arc, + config: CoreRuntimeConfigStore, + stats: Arc, + protocol_label: &'static str, + smoltcp_enabled: AtomicBool, +} + +impl CoreProxyRuntime +where + H: DirectConnectorHost, +{ + #[allow(clippy::too_many_arguments)] + pub(crate) fn new( + peer_manager: Arc, + host: Arc, + protected_tcp_ports: Arc, + running_listeners: Arc, + config: CoreRuntimeConfigStore, + protocol_label: &'static str, + ) -> Arc { + Arc::new(Self { + stats: peer_manager.stats_manager(), + peer_manager, + host, + protected_tcp_ports, + running_listeners, + config, + protocol_label, + smoltcp_enabled: AtomicBool::new(false), + }) + } + + pub(crate) fn latch_smoltcp(&self) { + self.smoltcp_enabled.store( + self.config.snapshot().services.proxy.force_smoltcp, + Ordering::Release, + ); + } + + fn should_deny_proxy(&self, destination: SocketAddr, is_udp: bool) -> bool { + let destination_is_local = self.host.is_local_ip(&destination.ip()) + || self.peer_manager.is_local_virtual_ip(&destination.ip()); + if !destination_is_local { + return false; + } + + self.running_listeners + .running_listeners() + .iter() + .any(|listener| { + listener.port() == Some(destination.port()) + && protocol_uses_udp(listener.scheme()) == is_udp + }) + || (!is_udp && self.protected_tcp_ports.contains(destination.port())) + } +} + +impl ProxyRuntimeInfo for CoreProxyRuntime +where + H: DirectConnectorHost, +{ + fn proxy_runtime_snapshot(&self) -> ProxyRuntimeSnapshot { + runtime_snapshot( + self.config.snapshot().as_ref(), + self.smoltcp_enabled.load(Ordering::Acquire), + ) + } + + fn is_ip_local_virtual_ip(&self, ip: &IpAddr) -> bool { + self.peer_manager.is_local_virtual_ip(ip) + } +} + +impl TcpProxyRuntime for CoreProxyRuntime +where + H: DirectConnectorHost, +{ + fn should_deny_tcp_proxy(&self, destination: SocketAddr) -> bool { + self.should_deny_proxy(destination, false) + } + + fn record_tcp_proxy_connect(&self, context: TcpProxyConnectContext, socket_dst: SocketAddr) { + self.stats + .get_counter( + MetricName::TcpProxyConnect, + LabelSet::new() + .with_label_type(LabelType::Protocol(self.protocol_label.to_owned())) + .with_label_type(LabelType::DstIp(socket_dst.ip().to_string())) + .with_label_type(LabelType::MappedDstIp(context.mapped_dst.ip().to_string())), + ) + .inc(); + } +} + +impl WrappedTcpDestinationRuntime for CoreProxyRuntime +where + H: DirectConnectorHost, +{ + fn is_ip_local_virtual_ip(&self, ip: &IpAddr) -> bool { + ProxyRuntimeInfo::is_ip_local_virtual_ip(self, ip) + } + + fn no_tun(&self) -> bool { + self.proxy_runtime_snapshot().no_tun + } + + fn should_deny_tcp_proxy(&self, dst: SocketAddr) -> bool { + TcpProxyRuntime::should_deny_tcp_proxy(self, dst) + } +} + +#[async_trait::async_trait] +impl UdpProxyPolicy for CoreProxyRuntime +where + H: DirectConnectorHost, +{ + fn should_deny_udp_proxy(&self, destination: SocketAddr) -> bool { + self.should_deny_proxy(destination, true) + } + + fn udp_response_ipv4_mtu(&self) -> usize { + self.config.snapshot().services.proxy.udp_response_ipv4_mtu + } +} + +struct CoreIcmpProxyRuntime +where + H: DirectConnectorHost, +{ + policy: Arc>, + host: Arc, + socket: std::sync::Mutex>>, + context: SocketContext, +} + +impl ProxyRuntimeInfo for CoreIcmpProxyRuntime +where + H: DirectConnectorHost, +{ + fn proxy_runtime_snapshot(&self) -> ProxyRuntimeSnapshot { + self.policy.proxy_runtime_snapshot() + } + + fn is_ip_local_virtual_ip(&self, ip: &IpAddr) -> bool { + ProxyRuntimeInfo::is_ip_local_virtual_ip(self.policy.as_ref(), ip) + } +} + +#[async_trait::async_trait] +impl IcmpProxyRuntime for CoreIcmpProxyRuntime +where + H: DirectConnectorHost, +{ + type Socket = dyn IcmpProxySocket; + + async fn start_icmp(&self) -> Result, ProxyRuntimeError> { + let socket = self.host.open_icmp_v4(self.context.clone()).await?; + self.socket.lock().unwrap().replace(socket.clone()); + Ok(socket) + } + + fn stop_icmp(&self) { + if let Some(socket) = self.socket.lock().unwrap().take() { + socket.close(); + } + } +} + +type CoreTcpProxy = TcpProxyService, H, TcpSocketProxyConnector>; +type CoreUdpProxyRuntime = UdpSocketProxyRuntime>; +type CoreUdpProxy = UdpProxyService>; +type CoreIcmpProxy = IcmpProxyService>; + +/// Deep portable proxy Module owned by one `CoreInstance`. +/// +/// The Host supplies socket creation and an optional raw-ICMP capability. Core +/// owns policy, CIDR authority, packet pipelines, NAT entries, and lifecycle. +pub(crate) struct CoreProxyModule +where + H: DirectConnectorHost + TcpHolePunchHost, +{ + operation: Mutex<()>, + runtime: Arc>, + tcp: Arc>, + icmp: Option>>, + udp_runtime: Arc>, + udp: Arc>, + tcp_started: AtomicBool, + icmp_started: AtomicBool, + udp_started: AtomicBool, +} + +impl CoreProxyModule +where + H: DirectConnectorHost + TcpHolePunchHost, +{ + #[allow(clippy::too_many_arguments)] + pub(crate) fn new( + peer_manager: Arc, + host: Arc, + protected_tcp_ports: Arc, + running_listeners: Arc, + config: CoreRuntimeConfigStore, + cidr_table: Arc, + tcp_socket_context: SocketContext, + udp_socket_context: SocketContext, + icmp_socket_context: SocketContext, + icmp_host: Option>, + ) -> Arc { + let runtime = CoreProxyRuntime::new( + peer_manager.clone(), + host.clone(), + protected_tcp_ports, + running_listeners, + config.clone(), + "TCP", + ); + let tcp_connector = Arc::new( + TcpSocketProxyConnector::new(host.clone()) + .with_socket_context(tcp_socket_context.clone()), + ); + let tcp = TcpProxyService::new_with_socket_context( + peer_manager.clone(), + runtime.clone(), + host.clone(), + tcp_connector, + cidr_table.clone(), + tcp_socket_context, + ); + let icmp = icmp_host.map(|host| { + IcmpProxyService::new( + peer_manager.clone(), + Arc::new(CoreIcmpProxyRuntime { + policy: runtime.clone(), + host, + socket: std::sync::Mutex::new(None), + context: icmp_socket_context.with_ip_version(IpVersion::V4), + }), + cidr_table.clone(), + PROXY_FRAGMENT_TIMEOUT, + ) + }); + let udp_runtime = Arc::new(UdpSocketProxyRuntime::new( + host, + runtime.clone(), + udp_proxy_bind_options(udp_socket_context), + UDP_PROXY_SOCKET_IDLE_TIMEOUT, + )); + let udp = UdpProxyService::new( + peer_manager, + udp_runtime.clone(), + cidr_table.clone(), + PROXY_FRAGMENT_TIMEOUT, + ); + + Arc::new(Self { + operation: Mutex::new(()), + runtime, + tcp, + icmp, + udp_runtime, + udp, + tcp_started: AtomicBool::new(false), + icmp_started: AtomicBool::new(false), + udp_started: AtomicBool::new(false), + }) + } + + pub(crate) fn tcp_entry_snapshots(&self) -> Vec { + self.tcp.engine().list_entries() + } + + fn stop_started(&self) { + if self.udp_started.swap(false, Ordering::AcqRel) { + self.udp.stop(); + self.udp_runtime.close_all(); + } + if self.icmp_started.swap(false, Ordering::AcqRel) + && let Some(icmp) = &self.icmp + { + icmp.stop(); + } + if self.tcp_started.swap(false, Ordering::AcqRel) { + self.tcp.stop(); + } + } + + pub(crate) async fn start(&self) -> Result<(), ProxyRuntimeError> { + let _operation = self.operation.lock().await; + if self.tcp_started.load(Ordering::Acquire) { + return Ok(()); + } + + self.runtime.latch_smoltcp(); + self.tcp_started.store(true, Ordering::Release); + if let Err(error) = self.tcp.start(true).await { + self.stop_started(); + return Err(error); + } + + if let Some(icmp) = &self.icmp { + self.icmp_started.store(true, Ordering::Release); + if let Err(error) = icmp.start().await { + self.icmp_started.store(false, Ordering::Release); + if self + .runtime + .config + .snapshot() + .services + .proxy + .icmp_failure_is_fatal + { + self.stop_started(); + return Err(error); + } + tracing::warn!(?error, "optional ICMP proxy runtime failed to start"); + } + } + + self.udp_started.store(true, Ordering::Release); + self.udp.start().await; + Ok(()) + } + + pub(crate) async fn stop(&self) { + let _operation = self.operation.lock().await; + self.stop_started(); + } +} + +#[cfg(test)] +mod tests { + use std::net::{IpAddr, Ipv4Addr}; + + use crate::{ + config::gateway::ProxyRuntimeConfig, + config::peers::{PeerRuntimeConfig, PeerRuntimeSnapshot}, + config::runtime::{CoreInstanceRuntimeConfig, CoreRuntimeConfig}, + config::{CoreConfig, IpPrefix, PeerPolicyConfig, ProxyNetworkConfig, RouteConfig}, + }; + + use super::*; + + fn test_config() -> CoreInstanceRuntimeConfig { + let mut peer = PeerRuntimeSnapshot::new( + PeerRuntimeConfig { + core: CoreConfig { + routes: RouteConfig { + ipv4: Some(IpPrefix { + address: IpAddr::V4(Ipv4Addr::new(10, 1, 2, 3)), + prefix_len: 24, + }), + proxy_networks: vec![ProxyNetworkConfig { + real: IpPrefix { + address: IpAddr::V4(Ipv4Addr::new(192, 0, 2, 0)), + prefix_len: 24, + }, + mapped: Some(IpPrefix { + address: IpAddr::V4(Ipv4Addr::new(198, 51, 100, 0)), + prefix_len: 24, + }), + }], + ..Default::default() + }, + peer_policy: PeerPolicyConfig { + latency_first: true, + ..Default::default() + }, + ..Default::default() + }, + network_identity: Default::default(), + stun_info: Default::default(), + feature_flags: Default::default(), + secure_mode: None, + host_routing: Default::default(), + }, + Default::default(), + ); + peer.flags.latency_first = true; + CoreInstanceRuntimeConfig { + services: CoreRuntimeConfig { + proxy: ProxyRuntimeConfig { + enable_exit_node: true, + no_tun: true, + ..Default::default() + }, + ..Default::default() + }, + peer: Arc::new(peer), + } + } + + #[test] + fn runtime_snapshot_uses_submitted_policy_and_latched_smoltcp() { + let config = test_config(); + + let kernel = runtime_snapshot(&config, false); + assert_eq!(kernel.local_inet.unwrap().to_string(), "10.1.2.3/24"); + assert_eq!(kernel.virtual_ipv4, Some(Ipv4Addr::new(10, 1, 2, 3))); + assert!(kernel.enable_exit_node); + assert!(kernel.no_tun); + assert!(kernel.latency_first); + + let smoltcp = runtime_snapshot(&config, true); + assert_eq!(smoltcp.local_inet, Some(smoltcp_proxy_inet())); + assert_eq!(smoltcp.virtual_ipv4, kernel.virtual_ipv4); + } + + #[test] + fn listener_protocol_classification_matches_native_proxy_guard() { + for scheme in ["udp", "wg", "quic"] { + assert!(protocol_uses_udp(scheme)); + } + for scheme in ["tcp", "ws", "wss", "faketcp"] { + assert!(!protocol_uses_udp(scheme)); + } + } + + #[test] + fn udp_proxy_bind_options_preserve_the_datagram_context() { + let context = SocketContext::default() + .with_socket_mark(Some(73)) + .with_netns(Some(crate::socket::NetNamespace::new("udp-proxy"))); + + let options = udp_proxy_bind_options(context); + + assert_eq!(options.context.ip_version, IpVersion::V4); + assert_eq!(options.context.socket_mark, Some(73)); + assert_eq!( + options.context.netns.as_ref().map(|netns| netns.token()), + Some("udp-proxy") + ); + assert_eq!( + options.purpose, + crate::socket::udp::UdpSocketPurpose::ProxyNat + ); + } +} diff --git a/easytier-core/src/gateway/proxy/tcp_proxy_engine.rs b/easytier-core/src/gateway/proxy/tcp_proxy_engine.rs new file mode 100644 index 00000000..5ea13d40 --- /dev/null +++ b/easytier-core/src/gateway/proxy/tcp_proxy_engine.rs @@ -0,0 +1,603 @@ +use std::{ + net::{IpAddr, Ipv4Addr, SocketAddr, SocketAddrV4}, + sync::{ + Arc, + atomic::{AtomicU16, Ordering}, + }, + time::{Duration, Instant, SystemTime, UNIX_EPOCH}, +}; + +use cidr::Ipv4Inet; +use crossbeam::atomic::AtomicCell; +use dashmap::DashMap; +use smoltcp::wire::{IpAddress, IpProtocol, Ipv4Packet, TcpPacket}; + +use crate::packet::{PacketType, ZCPacket}; + +use super::cidr_table::ProxyCidrTable; + +pub(crate) type TcpNatEntryId = uuid::Uuid; + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub(crate) enum TcpProxyMode { + Tcp, + KcpSrc, + QuicSrc, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum TcpNatEntryState { + SynReceived, + ConnectingDst, + Connected, + ClosingSrc, + ClosingDst, + Closed, +} + +#[derive(Debug, Clone)] +pub struct TcpNatEntrySnapshot { + pub src: SocketAddr, + pub dst: SocketAddr, + pub mapped_dst: SocketAddr, + pub start_time: u64, + pub state: TcpNatEntryState, +} + +#[derive(Debug)] +pub(crate) struct TcpNatEntry { + id: TcpNatEntryId, + src: SocketAddr, + real_dst: SocketAddr, + mapped_dst: SocketAddr, + start_time: Instant, + start_time_unix_secs: u64, + state: AtomicCell, +} + +impl TcpNatEntry { + fn new(src: SocketAddr, real_dst: SocketAddr, mapped_dst: SocketAddr) -> Self { + Self { + id: uuid::Uuid::new_v4(), + src, + real_dst, + mapped_dst, + start_time: Instant::now(), + start_time_unix_secs: SystemTime::now() + .duration_since(UNIX_EPOCH) + .map(|duration| duration.as_secs()) + .unwrap_or_default(), + state: AtomicCell::new(TcpNatEntryState::SynReceived), + } + } + + pub fn id(&self) -> TcpNatEntryId { + self.id + } + + pub fn src(&self) -> SocketAddr { + self.src + } + + pub fn real_dst(&self) -> SocketAddr { + self.real_dst + } + + pub fn mapped_dst(&self) -> SocketAddr { + self.mapped_dst + } + + pub fn state(&self) -> TcpNatEntryState { + self.state.load() + } + + pub fn set_state(&self, state: TcpNatEntryState) { + self.state.store(state); + } + + fn snapshot(&self) -> TcpNatEntrySnapshot { + TcpNatEntrySnapshot { + src: self.src, + dst: self.real_dst, + mapped_dst: self.mapped_dst, + start_time: self.start_time_unix_secs, + state: self.state(), + } + } +} + +#[derive(Clone, Copy, Debug)] +pub(crate) struct TcpProxyPeerContext { + pub local_inet: Option, + pub virtual_ipv4: Option, + pub local_port: u16, + pub enable_exit_node: bool, + pub no_tun: bool, + pub smoltcp_enabled: bool, +} + +#[derive(Clone, Copy, Debug)] +pub(crate) struct TcpProxyNicContext { + pub local_inet: Option, + pub local_port: u16, + pub my_peer_id: u32, + pub smoltcp_enabled: bool, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub(crate) enum TcpProxyPacketAction { + Handled { new_syn: bool }, + Pass, +} + +#[derive(Debug)] +pub(crate) struct TcpProxyEngine { + cidr_table: Arc, + local_port: AtomicU16, + syn_map: DashMap>, + conn_map: DashMap>, + addr_conn_map: DashMap>, +} + +impl TcpProxyEngine { + pub fn new(cidr_table: Arc) -> Self { + Self { + cidr_table, + local_port: AtomicU16::new(0), + syn_map: DashMap::new(), + conn_map: DashMap::new(), + addr_conn_map: DashMap::new(), + } + } + + pub fn set_local_port(&self, port: u16) { + self.local_port.store(port, Ordering::Relaxed); + } + + pub fn local_port(&self) -> u16 { + self.local_port.load(Ordering::Relaxed) + } + + pub fn check_packet_from_peer_fast( + &self, + mode: TcpProxyMode, + ctx: &TcpProxyPeerContext, + ) -> bool { + match mode { + TcpProxyMode::Tcp => !self.cidr_table.is_empty() || ctx.enable_exit_node || ctx.no_tun, + TcpProxyMode::KcpSrc | TcpProxyMode::QuicSrc => true, + } + } + + pub fn try_handle_peer_packet( + &self, + mode: TcpProxyMode, + packet: &mut ZCPacket, + ctx: TcpProxyPeerContext, + ) -> TcpProxyPacketAction { + if !self.check_packet_from_peer_fast(mode, &ctx) { + return TcpProxyPacketAction::Pass; + } + + let Some(local_inet) = ctx.local_inet else { + return TcpProxyPacketAction::Pass; + }; + let local_ip = local_inet.address(); + let Some(hdr) = packet.peer_manager_header() else { + return TcpProxyPacketAction::Pass; + }; + if hdr.is_no_proxy() { + return TcpProxyPacketAction::Pass; + } + + let allowed_packet_type = match mode { + TcpProxyMode::Tcp => hdr.packet_type == PacketType::Data as u8, + TcpProxyMode::KcpSrc => { + hdr.packet_type == PacketType::DataWithKcpSrcModified as u8 + && hdr.from_peer_id == hdr.to_peer_id + } + TcpProxyMode::QuicSrc => { + hdr.packet_type == PacketType::DataWithQuicSrcModified as u8 + && hdr.from_peer_id == hdr.to_peer_id + } + }; + if !allowed_packet_type { + return TcpProxyPacketAction::Pass; + } + + let Ok(ip_packet) = Ipv4Packet::new_checked(packet.payload()) else { + return TcpProxyPacketAction::Pass; + }; + if ip_packet.version() != 4 || ip_packet.next_header() != IpProtocol::Tcp { + return TcpProxyPacketAction::Pass; + } + let origin_ip = ip_packet.dst_addr(); + + let Some(real_dst_ip) = + self.real_dst_ip_for_mode(mode, origin_ip, hdr.is_exit_node(), &ctx) + else { + return TcpProxyPacketAction::Pass; + }; + + let hdr = packet + .mut_peer_manager_header() + .expect("peer manager header"); + hdr.packet_type = PacketType::Data as u8; + + let payload_bytes = packet.mut_payload(); + let ip_packet = Ipv4Packet::new_checked(&payload_bytes[..]).expect("checked ipv4 packet"); + let tcp_packet = TcpPacket::new_checked(ip_packet.payload()).expect("checked tcp packet"); + + let source_ip = ip_packet.src_addr(); + let source_port = tcp_packet.src_port(); + let src = SocketAddr::V4(SocketAddrV4::new(source_ip, source_port)); + + let mut new_syn = false; + if tcp_packet.syn() && !tcp_packet.ack() { + let dest_ip = ip_packet.dst_addr(); + let dest_port = tcp_packet.dst_port(); + let mapped_dst = SocketAddr::V4(SocketAddrV4::new(dest_ip, dest_port)); + let real_dst = SocketAddr::V4(SocketAddrV4::new(real_dst_ip, dest_port)); + + let old_val = self + .syn_map + .insert(src, Arc::new(TcpNatEntry::new(src, real_dst, mapped_dst))); + tracing::info!( + ?src, + ?real_dst, + ?mapped_dst, + old_entry = ?old_val, + "tcp syn received" + ); + new_syn = true; + } else if !self.addr_conn_map.contains_key(&src) && !self.syn_map.contains_key(&src) { + return TcpProxyPacketAction::Pass; + } + + let mut ip_packet = Ipv4Packet::new_checked(payload_bytes).expect("checked ipv4 packet"); + if !ctx.smoltcp_enabled && source_ip == local_ip { + ip_packet.set_src_addr(Self::fake_local_ipv4(&local_inet)); + } + ip_packet.set_dst_addr(local_ip); + let source = ip_packet.src_addr(); + { + let mut tcp_packet = + TcpPacket::new_checked(ip_packet.payload_mut()).expect("checked tcp packet"); + tcp_packet.set_dst_port(ctx.local_port); + tcp_packet.fill_checksum(&IpAddress::Ipv4(source), &IpAddress::Ipv4(local_ip)); + } + ip_packet.fill_checksum(); + + tracing::trace!(?source, ?local_ip, ?packet, "tcp packet after modified"); + TcpProxyPacketAction::Handled { new_syn } + } + + pub fn try_process_packet_from_nic( + &self, + zc_packet: &mut ZCPacket, + ctx: TcpProxyNicContext, + ) -> bool { + let Some(local_inet) = ctx.local_inet else { + return false; + }; + let local_ip = local_inet.address(); + + let data = zc_packet.payload(); + let Ok(ip_packet) = Ipv4Packet::new_checked(data) else { + return false; + }; + if ip_packet.version() != 4 + || ip_packet.src_addr() != local_ip + || ip_packet.next_header() != IpProtocol::Tcp + { + return false; + } + + let Ok(tcp_packet) = TcpPacket::new_checked(ip_packet.payload()) else { + return false; + }; + if tcp_packet.src_port() != ctx.local_port { + return false; + } + + let mut dst_addr = SocketAddr::V4(SocketAddrV4::new( + ip_packet.dst_addr(), + tcp_packet.dst_port(), + )); + let mut need_transform_dst = false; + + if !ctx.smoltcp_enabled && dst_addr.ip() == Self::fake_local_ipv4(&local_inet) { + dst_addr.set_ip(IpAddr::V4(local_ip)); + need_transform_dst = true; + } + + tracing::trace!(?dst_addr, "tcp packet try find entry"); + let entry = if let Some(entry) = self.addr_conn_map.get(&dst_addr) { + entry.clone() + } else { + let Some(syn_entry) = self.syn_map.get(&dst_addr) else { + return false; + }; + syn_entry.clone() + }; + assert_eq!(entry.src, dst_addr); + + let IpAddr::V4(mapped_dst_ip) = entry.mapped_dst.ip() else { + panic!("v4 nat entry src ip is not v4"); + }; + + let hdr = zc_packet + .mut_peer_manager_header() + .expect("peer manager header"); + hdr.set_no_proxy(true); + if need_transform_dst { + hdr.to_peer_id = ctx.my_peer_id.into(); + } + + let mut ip_packet = + Ipv4Packet::new_checked(zc_packet.mut_payload()).expect("checked ipv4 packet"); + ip_packet.set_src_addr(mapped_dst_ip); + if need_transform_dst { + ip_packet.set_dst_addr(local_ip); + } + let dst = ip_packet.dst_addr(); + + { + let mut tcp_packet = + TcpPacket::new_checked(ip_packet.payload_mut()).expect("checked tcp packet"); + tcp_packet.set_src_port(entry.real_dst.port()); + tcp_packet.fill_checksum(&IpAddress::Ipv4(mapped_dst_ip), &IpAddress::Ipv4(dst)); + } + ip_packet.fill_checksum(); + + tracing::trace!(?dst_addr, nat_entry = ?entry, packet = ?ip_packet, "tcp packet after modified"); + true + } + + pub fn accept_connection( + &self, + mut socket_addr: SocketAddr, + virtual_inet: Option, + ) -> Option> { + if let Some(my_ip_inet) = virtual_inet { + let my_ip = my_ip_inet.address(); + if socket_addr.ip() == Self::fake_local_ipv4(&my_ip_inet) { + socket_addr.set_ip(IpAddr::V4(my_ip)); + } + } + + let (_, entry) = self.syn_map.remove(&socket_addr)?; + if entry.state() != TcpNatEntryState::SynReceived { + return None; + } + + entry.set_state(TcpNatEntryState::ConnectingDst); + self.addr_conn_map.insert(entry.src, entry.clone()); + let old_nat_val = self.conn_map.insert(entry.id, entry.clone()); + assert!(old_nat_val.is_none()); + Some(entry) + } + + pub fn remove_entry(&self, entry_id: TcpNatEntryId) { + let Some((_, entry)) = self.conn_map.remove(&entry_id) else { + return; + }; + self.addr_conn_map + .remove_if(&entry.src, |_, current| current.id == entry.id); + if self.conn_map.capacity() - self.conn_map.len() > 16 { + self.conn_map.shrink_to_fit(); + } + if self.addr_conn_map.capacity() - self.addr_conn_map.len() > 16 { + self.addr_conn_map.shrink_to_fit(); + } + } + + pub fn cleanup_expired_syn(&self, timeout: Duration) { + self.syn_map.retain(|_, entry| { + if entry.start_time.elapsed() > timeout { + tracing::warn!(?entry, "syn nat entry expired"); + entry.set_state(TcpNatEntryState::Closed); + false + } else { + true + } + }); + self.syn_map.shrink_to_fit(); + } + + pub fn is_tcp_proxy_connection(&self, src: SocketAddr) -> bool { + self.syn_map.contains_key(&src) || self.addr_conn_map.contains_key(&src) + } + + pub fn list_entries(&self) -> Vec { + let mut entries = Vec::new(); + for entry in self.syn_map.iter() { + entries.push(entry.value().snapshot()); + } + for entry in self.conn_map.iter() { + entries.push(entry.value().snapshot()); + } + entries + } + + pub fn fake_local_ipv4(local_ip: &Ipv4Inet) -> Ipv4Addr { + local_ip.first_address() + } + + fn real_dst_ip_for_mode( + &self, + mode: TcpProxyMode, + origin_ip: Ipv4Addr, + is_exit_node: bool, + ctx: &TcpProxyPeerContext, + ) -> Option { + match mode { + TcpProxyMode::Tcp => { + if let Some(real_ip) = self.cidr_table.lookup_v4(origin_ip) { + return Some(real_ip); + } + let no_tun_local_virtual_ip = ctx.no_tun && Some(origin_ip) == ctx.virtual_ipv4; + (is_exit_node || no_tun_local_virtual_ip).then_some(origin_ip) + } + TcpProxyMode::KcpSrc | TcpProxyMode::QuicSrc => Some(origin_ip), + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::{ + gateway::proxy::cidr_table::{ProxyCidrRule, ProxyCidrSnapshot}, + packet::PeerManagerHeader, + }; + use smoltcp::wire::{IpAddress, TcpPacket}; + + fn build_tcp_packet(src: SocketAddrV4, dst: SocketAddrV4, syn: bool, ack: bool) -> ZCPacket { + let mut raw = vec![0; smoltcp::wire::IPV4_HEADER_LEN + smoltcp::wire::TCP_HEADER_LEN]; + { + let mut ipv4 = Ipv4Packet::new_unchecked(&mut raw); + ipv4.set_version(4); + ipv4.set_header_len(smoltcp::wire::IPV4_HEADER_LEN as u8); + ipv4.set_total_len( + (smoltcp::wire::IPV4_HEADER_LEN + smoltcp::wire::TCP_HEADER_LEN) as u16, + ); + ipv4.set_hop_limit(64); + ipv4.set_next_header(IpProtocol::Tcp); + ipv4.set_src_addr(*src.ip()); + ipv4.set_dst_addr(*dst.ip()); + ipv4.fill_checksum(); + } + { + let mut tcp = TcpPacket::new_unchecked(&mut raw[smoltcp::wire::IPV4_HEADER_LEN..]); + tcp.set_src_port(src.port()); + tcp.set_dst_port(dst.port()); + tcp.set_header_len(smoltcp::wire::TCP_HEADER_LEN as u8); + tcp.set_syn(syn); + tcp.set_ack(ack); + tcp.fill_checksum(&IpAddress::Ipv4(*src.ip()), &IpAddress::Ipv4(*dst.ip())); + } + + let mut packet = ZCPacket::new_with_payload(&raw); + packet.fill_peer_manager_hdr(1, 2, PacketType::Data as u8); + packet + } + + fn tcp_engine() -> TcpProxyEngine { + TcpProxyEngine::new(Arc::new(ProxyCidrTable::from_snapshot(ProxyCidrSnapshot { + rules: vec![ProxyCidrRule { + cidr: "127.0.0.0/24".parse().unwrap(), + mapped_cidr: Some("10.10.10.0/24".parse().unwrap()), + }], + }))) + } + + fn peer_ctx() -> TcpProxyPeerContext { + TcpProxyPeerContext { + local_inet: Some("10.144.144.204/24".parse().unwrap()), + virtual_ipv4: Some("10.144.144.204".parse().unwrap()), + local_port: 8899, + enable_exit_node: false, + no_tun: false, + smoltcp_enabled: false, + } + } + + #[test] + fn peer_syn_creates_entry_and_rewrites_to_local_stack() { + let engine = tcp_engine(); + let src = SocketAddrV4::new("10.144.144.206".parse().unwrap(), 50000); + let mapped_dst = SocketAddrV4::new("10.10.10.42".parse().unwrap(), 80); + let mut packet = build_tcp_packet(src, mapped_dst, true, false); + + assert_eq!( + engine.try_handle_peer_packet(TcpProxyMode::Tcp, &mut packet, peer_ctx()), + TcpProxyPacketAction::Handled { new_syn: true } + ); + + let ipv4 = Ipv4Packet::new_checked(packet.payload()).unwrap(); + assert_eq!(ipv4.src_addr(), *src.ip()); + assert_eq!( + ipv4.dst_addr(), + "10.144.144.204".parse::().unwrap() + ); + let tcp = TcpPacket::new_checked(ipv4.payload()).unwrap(); + assert_eq!(tcp.src_port(), src.port()); + assert_eq!(tcp.dst_port(), 8899); + + let entries = engine.list_entries(); + assert_eq!(entries.len(), 1); + assert_eq!(entries[0].src, SocketAddr::V4(src)); + assert_eq!( + entries[0].dst, + SocketAddr::V4(SocketAddrV4::new("127.0.0.42".parse().unwrap(), 80)) + ); + assert_eq!(entries[0].mapped_dst, SocketAddr::V4(mapped_dst)); + } + + #[test] + fn nic_response_rewrites_back_to_mapped_destination() { + let engine = tcp_engine(); + let src = SocketAddrV4::new("10.144.144.206".parse().unwrap(), 50000); + let mapped_dst = SocketAddrV4::new("10.10.10.42".parse().unwrap(), 80); + let mut request = build_tcp_packet(src, mapped_dst, true, false); + assert!(matches!( + engine.try_handle_peer_packet(TcpProxyMode::Tcp, &mut request, peer_ctx()), + TcpProxyPacketAction::Handled { new_syn: true } + )); + let entry = engine + .accept_connection( + SocketAddr::V4(src), + Some("10.144.144.204/24".parse().unwrap()), + ) + .unwrap(); + assert_eq!(entry.state(), TcpNatEntryState::ConnectingDst); + + let local = SocketAddrV4::new("10.144.144.204".parse().unwrap(), 8899); + let mut response = build_tcp_packet(local, src, false, true); + assert!(engine.try_process_packet_from_nic( + &mut response, + TcpProxyNicContext { + local_inet: Some("10.144.144.204/24".parse().unwrap()), + local_port: 8899, + my_peer_id: 2, + smoltcp_enabled: false, + }, + )); + + let hdr: &PeerManagerHeader = response.peer_manager_header().unwrap(); + assert!(hdr.is_no_proxy()); + let ipv4 = Ipv4Packet::new_checked(response.payload()).unwrap(); + assert_eq!(ipv4.src_addr(), *mapped_dst.ip()); + assert_eq!(ipv4.dst_addr(), *src.ip()); + let tcp = TcpPacket::new_checked(ipv4.payload()).unwrap(); + assert_eq!(tcp.src_port(), mapped_dst.port()); + assert_eq!(tcp.dst_port(), src.port()); + } + + #[test] + fn accept_connection_does_not_resurrect_closed_syn_entry() { + let engine = tcp_engine(); + let src = SocketAddrV4::new("10.144.144.206".parse().unwrap(), 50000); + let mapped_dst = SocketAddrV4::new("10.10.10.42".parse().unwrap(), 80); + let mut request = build_tcp_packet(src, mapped_dst, true, false); + assert!(matches!( + engine.try_handle_peer_packet(TcpProxyMode::Tcp, &mut request, peer_ctx()), + TcpProxyPacketAction::Handled { new_syn: true } + )); + let entry = engine.syn_map.get(&SocketAddr::V4(src)).unwrap().clone(); + entry.set_state(TcpNatEntryState::Closed); + + assert!( + engine + .accept_connection( + SocketAddr::V4(src), + Some("10.144.144.204/24".parse().unwrap()), + ) + .is_none() + ); + assert!(engine.syn_map.get(&SocketAddr::V4(src)).is_none()); + assert!(engine.addr_conn_map.get(&SocketAddr::V4(src)).is_none()); + assert!(engine.conn_map.is_empty()); + } +} diff --git a/easytier-core/src/gateway/proxy/tcp_proxy_service.rs b/easytier-core/src/gateway/proxy/tcp_proxy_service.rs new file mode 100644 index 00000000..d41bff2b --- /dev/null +++ b/easytier-core/src/gateway/proxy/tcp_proxy_service.rs @@ -0,0 +1,631 @@ +use std::future::Future; +use std::sync::{Arc, Weak, atomic::Ordering}; +use std::time::Duration; + +use atomic_shim::AtomicU64; +use tokio::io::{AsyncWriteExt, copy}; +use tokio::task::JoinSet; + +use crate::{ + foundation::time::timeout, + packet::ZCPacket, + peers::{ + NicPacketFilter, PeerPacketFilter, + peer_manager::{PeerManagerCore, PipelineRegistrationGuard}, + }, + socket::{ + SocketContext, + tcp::{TcpBindOptions, TcpListenOptions, VirtualTcpListener, VirtualTcpListenerFactory}, + }, +}; + +use super::cidr_table::ProxyCidrTable; +use super::tcp_proxy_engine::{ + TcpNatEntry, TcpNatEntryState, TcpProxyEngine, TcpProxyMode, TcpProxyNicContext, + TcpProxyPacketAction, TcpProxyPeerContext, +}; +use super::traits::{ + ProxyRuntimeError, TcpProxyConnectContext, TcpProxyDestinationConnector, TcpProxyRuntime, + TcpProxyStream, +}; +#[cfg(feature = "proxy-smoltcp-stack")] +use crate::gateway::smoltcp::{SmolTcpStack, output_dst_ip}; + +fn spawn_tcp_proxy_task( + lifecycle: &AtomicU64, + expected_generation: u64, + tasks: &std::sync::Mutex>, + task: impl Future + Send + 'static, +) -> bool { + let mut tasks = tasks.lock().unwrap(); + if lifecycle.load(Ordering::Acquire) != expected_generation { + return false; + } + tasks.spawn(task); + true +} + +pub struct TcpProxyService< + R: TcpProxyRuntime + 'static, + F: VirtualTcpListenerFactory, + C: TcpProxyDestinationConnector, +> { + peer_manager: Arc, + runtime: Arc, + listener_factory: Arc, + socket_context: SocketContext, + connector: Arc, + engine: Arc, + mode: TcpProxyMode, + peer_pipeline_guard: std::sync::Mutex>, + nic_pipeline_guard: std::sync::Mutex>, + kernel_listener: std::sync::Mutex>>, + #[cfg(feature = "proxy-smoltcp-stack")] + smoltcp_stack: std::sync::Mutex>>, + tasks: std::sync::Mutex>, + lifecycle: AtomicU64, +} + +impl + TcpProxyService +{ + pub fn new_with_socket_context( + peer_manager: Arc, + runtime: Arc, + listener_factory: Arc, + connector: Arc, + cidr_table: Arc, + socket_context: SocketContext, + ) -> Arc { + let mode = connector.proxy_mode(); + Arc::new(Self { + peer_manager, + runtime, + listener_factory, + socket_context, + connector, + engine: Arc::new(TcpProxyEngine::new(cidr_table)), + mode, + peer_pipeline_guard: std::sync::Mutex::new(None), + nic_pipeline_guard: std::sync::Mutex::new(None), + kernel_listener: std::sync::Mutex::new(None), + #[cfg(feature = "proxy-smoltcp-stack")] + smoltcp_stack: std::sync::Mutex::new(None), + tasks: std::sync::Mutex::new(JoinSet::new()), + lifecycle: AtomicU64::new(0), + }) + } + + pub fn engine(&self) -> Arc { + self.engine.clone() + } + + pub fn is_started(&self) -> bool { + self.lifecycle.load(Ordering::Acquire) & 1 != 0 + } + + pub async fn start(self: &Arc, register_pipeline: bool) -> Result<(), ProxyRuntimeError> { + let generation = loop { + let stopped = self.lifecycle.load(Ordering::Acquire); + if stopped & 1 != 0 { + return Ok(()); + } + let active = stopped.wrapping_add(1); + if self + .lifecycle + .compare_exchange(stopped, active, Ordering::AcqRel, Ordering::Acquire) + .is_ok() + { + break active; + } + }; + + let snapshot = self.runtime.proxy_runtime_snapshot(); + let start_result = if snapshot.smoltcp_enabled { + #[cfg(feature = "proxy-smoltcp-stack")] + { + self.start_smoltcp(generation).await + } + + #[cfg(not(feature = "proxy-smoltcp-stack"))] + { + Err(ProxyRuntimeError::Other(anyhow::anyhow!( + "smoltcp proxy stack feature is disabled" + ))) + } + } else { + self.start_kernel_listener(generation).await + }; + + if let Err(err) = start_result { + let _ = self.lifecycle.compare_exchange( + generation, + generation.wrapping_add(1), + Ordering::AcqRel, + Ordering::Acquire, + ); + return Err(err); + } + + if register_pipeline { + self.register_pipeline().await; + } + self.spawn_syn_cleanup(generation); + + Ok(()) + } + + pub fn stop(&self) { + self.stop_resources(); + self.tasks.lock().unwrap().abort_all(); + } + + pub(crate) async fn stop_and_wait(&self) { + self.stop_resources(); + let mut tasks = std::mem::take(&mut *self.tasks.lock().unwrap()); + tasks.shutdown().await; + } + + fn stop_resources(&self) { + loop { + let generation = self.lifecycle.load(Ordering::Acquire); + if generation & 1 == 0 + || self + .lifecycle + .compare_exchange( + generation, + generation.wrapping_add(1), + Ordering::AcqRel, + Ordering::Acquire, + ) + .is_ok() + { + break; + } + } + if let Some(guard) = self.peer_pipeline_guard.lock().unwrap().take() { + guard.close(); + } + if let Some(guard) = self.nic_pipeline_guard.lock().unwrap().take() { + guard.close(); + } + if let Some(listener) = self.kernel_listener.lock().unwrap().take() { + drop(listener); + } + #[cfg(feature = "proxy-smoltcp-stack")] + if let Some(stack) = self.smoltcp_stack.lock().unwrap().take() { + drop(stack); + } + } + + async fn register_pipeline(self: &Arc) { + self.register_peer_pipeline().await; + self.register_nic_pipeline().await; + } + + pub async fn register_peer_pipeline(self: &Arc) { + if self.peer_pipeline_guard.lock().unwrap().is_some() { + return; + } + let peer_guard = self + .peer_manager + .add_managed_packet_process_pipeline(Box::new(TcpProxyServiceFilter { + service: Arc::downgrade(self), + })) + .await; + self.peer_pipeline_guard.lock().unwrap().replace(peer_guard); + } + + pub async fn register_nic_pipeline(self: &Arc) { + if self.nic_pipeline_guard.lock().unwrap().is_some() { + return; + } + let nic_guard = self + .peer_manager + .add_managed_nic_packet_process_pipeline(Box::new(TcpProxyServiceFilter { + service: Arc::downgrade(self), + })) + .await; + self.nic_pipeline_guard.lock().unwrap().replace(nic_guard); + } + + fn spawn_syn_cleanup(self: &Arc, generation: u64) { + let service = Arc::downgrade(self); + let _ = spawn_tcp_proxy_task(&self.lifecycle, generation, &self.tasks, async move { + loop { + crate::foundation::time::sleep(Duration::from_secs(10)).await; + let Some(service) = service.upgrade() else { + break; + }; + service.engine.cleanup_expired_syn(Duration::from_secs(30)); + service.drain_completed_tasks(); + } + }); + } + + fn drain_completed_tasks(&self) { + let mut tasks = self.tasks.lock().unwrap(); + while let Some(result) = tasks.try_join_next() { + if let Err(err) = result { + tracing::warn!(?err, "tcp proxy task finished with error"); + } + } + } + + async fn start_kernel_listener( + self: &Arc, + generation: u64, + ) -> Result<(), ProxyRuntimeError> { + let listen_addr = std::net::SocketAddr::new(std::net::Ipv4Addr::UNSPECIFIED.into(), 0); + let listener = self + .listener_factory + .bind_tcp( + TcpListenOptions::proxy_nat(listen_addr).with_bind( + TcpBindOptions::default() + .with_context( + self.socket_context + .clone() + .with_ip_version(crate::socket::IpVersion::V4), + ) + .with_local_addr(Some(listen_addr)), + ), + ) + .await?; + self.engine.set_local_port(listener.local_addr()?.port()); + self.kernel_listener + .lock() + .unwrap() + .replace(listener.clone()); + + let service = Arc::downgrade(self); + let _ = spawn_tcp_proxy_task(&self.lifecycle, generation, &self.tasks, async move { + loop { + let accept_ret = listener.accept().await; + let Ok((src_stream, socket_addr)) = accept_ret else { + tracing::error!( + error = ?accept_ret.err(), + "nat tcp listener accept failed" + ); + continue; + }; + let Some(service) = service.upgrade() else { + break; + }; + service + .handle_accept(generation, socket_addr, Box::new(src_stream)) + .await; + } + }); + + Ok(()) + } + + #[cfg(feature = "proxy-smoltcp-stack")] + async fn start_smoltcp(self: &Arc, generation: u64) -> Result<(), ProxyRuntimeError> { + let local_ip = self + .runtime + .proxy_runtime_snapshot() + .local_inet + .map(|inet| inet.address()) + .unwrap_or(std::net::Ipv4Addr::new(192, 88, 99, 254)); + let stack = SmolTcpStack::new(local_ip).await?; + self.engine.set_local_port(stack.local_port()); + + let mut output_rx = stack.take_output_rx().await?; + let peer_manager = self.peer_manager.clone(); + let _ = spawn_tcp_proxy_task(&self.lifecycle, generation, &self.tasks, async move { + while let Some(data) = output_rx.recv().await { + tracing::trace!(?data, "receive from smoltcp stack and send to peer manager"); + let dst = match output_dst_ip(&data) { + Ok(dst) => dst, + Err(err) => { + tracing::error!(?err, ?data, "invalid smoltcp output packet"); + continue; + } + }; + let packet = ZCPacket::new_with_payload(&data); + if let Err(err) = peer_manager.send_msg_by_ip(packet, dst, false).await { + tracing::error!(?err, "send to peer failed in smoltcp sender"); + } + } + tracing::error!("smoltcp stack stream exited"); + }); + + let service = Arc::downgrade(self); + let accept_stack = stack.clone(); + let _ = spawn_tcp_proxy_task(&self.lifecycle, generation, &self.tasks, async move { + loop { + let accept_ret = accept_stack.accept().await; + let Ok((socket_addr, src_stream)) = accept_ret else { + tracing::error!(error = ?accept_ret.err(), "smoltcp accept failed"); + continue; + }; + let Some(service) = service.upgrade() else { + break; + }; + service + .handle_accept(generation, socket_addr, src_stream) + .await; + } + }); + + self.smoltcp_stack.lock().unwrap().replace(stack); + Ok(()) + } + + async fn handle_accept( + self: Arc, + generation: u64, + socket_addr: std::net::SocketAddr, + src_stream: Box, + ) { + let snapshot = self.runtime.proxy_runtime_snapshot(); + let Some(entry) = self + .engine + .accept_connection(socket_addr, snapshot.local_inet) + else { + tracing::error!( + ?socket_addr, + "tcp connection from unknown source, ignore it" + ); + return; + }; + tracing::info!( + ?socket_addr, + "tcp connection accepted for proxy, nat dst: {:?}", + entry.real_dst() + ); + + if !spawn_tcp_proxy_task( + &self.lifecycle, + generation, + &self.tasks, + Self::connect_to_nat_dst(self.clone(), src_stream, entry.clone()), + ) { + entry.set_state(TcpNatEntryState::Closed); + self.engine.remove_entry(entry.id()); + } + } + + async fn connect_to_nat_dst( + service: Arc, + mut src_stream: Box, + entry: Arc, + ) { + if service.runtime.should_deny_tcp_proxy(entry.real_dst()) { + tracing::error!( + ?entry, + "nat dst port {} is in running listeners, ignore it", + entry.real_dst().port() + ); + entry.set_state(TcpNatEntryState::Closed); + service.engine.remove_entry(entry.id()); + return; + } + + let ctx = TcpProxyConnectContext { + src: entry.src(), + real_dst: entry.real_dst(), + mapped_dst: entry.mapped_dst(), + }; + let socket_dst = if service.runtime.is_ip_local_virtual_ip(&ctx.real_dst.ip()) { + std::net::SocketAddr::new(std::net::Ipv4Addr::LOCALHOST.into(), ctx.real_dst.port()) + } else { + ctx.real_dst + }; + service.runtime.record_tcp_proxy_connect(ctx, socket_dst); + + let Ok(dst_stream) = service.connector.connect(ctx.src, socket_dst).await else { + tracing::error!("connect to dst failed: {:?}", entry); + entry.set_state(TcpNatEntryState::Closed); + service.engine.remove_entry(entry.id()); + return; + }; + let mut dst_stream: Box = Box::new(dst_stream); + + tracing::info!(?entry, "tcp connection to dst established"); + if entry.state() == TcpNatEntryState::ConnectingDst { + entry.set_state(TcpNatEntryState::Connected); + } + + let ret = copy_bidirectional_no_shutdown(src_stream.as_mut(), dst_stream.as_mut()).await; + tracing::info!(nat_entry = ?entry, ret = ?ret, "nat tcp connection closed"); + + entry.set_state(TcpNatEntryState::ClosingSrc); + let ret = timeout(Duration::from_secs(10), src_stream.shutdown()).await; + tracing::info!(nat_entry = ?entry, ret = ?ret, "src tcp stream shutdown"); + + entry.set_state(TcpNatEntryState::ClosingDst); + let ret = timeout(Duration::from_secs(10), dst_stream.shutdown()).await; + tracing::info!(nat_entry = ?entry, ret = ?ret, "dst tcp stream shutdown"); + + drop(src_stream); + drop(dst_stream); + + entry.set_state(TcpNatEntryState::Closed); + crate::foundation::time::sleep(Duration::from_secs(10)).await; + service.engine.remove_entry(entry.id()); + } + + async fn handle_peer_packet(self: Arc, mut packet: ZCPacket) -> Option { + let snapshot = self.runtime.proxy_runtime_snapshot(); + let action = self.engine.try_handle_peer_packet( + self.mode, + &mut packet, + TcpProxyPeerContext { + local_inet: snapshot.local_inet, + virtual_ipv4: snapshot.virtual_ipv4, + local_port: self.engine.local_port(), + enable_exit_node: snapshot.enable_exit_node, + no_tun: snapshot.no_tun, + smoltcp_enabled: snapshot.smoltcp_enabled, + }, + ); + let TcpProxyPacketAction::Handled { new_syn: _new_syn } = action else { + return Some(packet); + }; + + if snapshot.smoltcp_enabled { + #[cfg(feature = "proxy-smoltcp-stack")] + self.handle_smoltcp_packet(packet, _new_syn).await; + + #[cfg(not(feature = "proxy-smoltcp-stack"))] + tracing::error!("smoltcp packet received but proxy-smoltcp-stack is disabled"); + } else if let Err(err) = self.peer_manager.get_nic_channel().send(packet).await { + tracing::error!(?err, "send to nic failed"); + } + + None + } + + #[cfg(feature = "proxy-smoltcp-stack")] + async fn handle_smoltcp_packet(&self, packet: ZCPacket, new_syn: bool) { + let stack = self.smoltcp_stack.lock().unwrap().clone(); + let Some(stack) = stack else { + tracing::error!("smoltcp stack is not started"); + return; + }; + if new_syn { + stack.add_listener().await; + } + if let Err(err) = stack.send_ingress(packet).await { + tracing::error!(?err, "send to smoltcp stack failed"); + } + } + + async fn handle_nic_packet(&self, packet: &mut ZCPacket) -> bool { + let snapshot = self.runtime.proxy_runtime_snapshot(); + self.engine.try_process_packet_from_nic( + packet, + TcpProxyNicContext { + local_inet: snapshot.local_inet, + local_port: self.engine.local_port(), + my_peer_id: self.peer_manager.my_peer_id(), + smoltcp_enabled: snapshot.smoltcp_enabled, + }, + ) + } +} + +async fn copy_bidirectional_no_shutdown( + src: &mut dyn TcpProxyStream, + dst: &mut dyn TcpProxyStream, +) -> Result<(), ProxyRuntimeError> { + let (mut src_reader, mut src_writer) = tokio::io::split(src); + let (mut dst_reader, mut dst_writer) = tokio::io::split(dst); + let src_to_dst = copy(&mut src_reader, &mut dst_writer); + let dst_to_src = copy(&mut dst_reader, &mut src_writer); + tokio::pin!(src_to_dst); + tokio::pin!(dst_to_src); + tokio::select! { + result = &mut src_to_dst => { + result?; + } + result = &mut dst_to_src => { + result?; + } + } + Ok(()) +} + +impl + Drop for TcpProxyService +{ + fn drop(&mut self) { + self.stop(); + } +} + +struct TcpProxyServiceFilter< + R: TcpProxyRuntime + 'static, + F: VirtualTcpListenerFactory, + C: TcpProxyDestinationConnector, +> { + service: Weak>, +} + +#[async_trait::async_trait] +impl + PeerPacketFilter for TcpProxyServiceFilter +{ + async fn try_process_packet_from_peer(&self, packet: ZCPacket) -> Option { + let Some(service) = self.service.upgrade() else { + return Some(packet); + }; + service.handle_peer_packet(packet).await + } +} + +#[async_trait::async_trait] +impl + NicPacketFilter for TcpProxyServiceFilter +{ + async fn try_process_packet_from_nic(&self, packet: &mut ZCPacket) -> bool { + let Some(service) = self.service.upgrade() else { + return false; + }; + service.handle_nic_packet(packet).await + } +} + +#[cfg(test)] +mod tests { + use super::*; + use std::sync::atomic::AtomicBool; + + struct DropSignal(Arc); + + impl Drop for DropSignal { + fn drop(&mut self) { + self.0.store(true, Ordering::Release); + } + } + + fn pending_task(dropped: Arc) -> impl Future + Send + 'static { + let signal = DropSignal(dropped); + async move { + let _signal = signal; + std::future::pending::<()>().await; + } + } + + #[tokio::test] + async fn stop_fence_linearizes_task_registration() { + let lifecycle = AtomicU64::new(1); + let tasks = std::sync::Mutex::new(JoinSet::new()); + let accepted_dropped = Arc::new(AtomicBool::new(false)); + + assert!(spawn_tcp_proxy_task( + &lifecycle, + 1, + &tasks, + pending_task(accepted_dropped.clone()), + )); + lifecycle.store(2, Ordering::Release); + let mut stopping = std::mem::take(&mut *tasks.lock().unwrap()); + stopping.shutdown().await; + assert!(accepted_dropped.load(Ordering::Acquire)); + + lifecycle.store(3, Ordering::Release); + + let rejected_dropped = Arc::new(AtomicBool::new(false)); + assert!(!spawn_tcp_proxy_task( + &lifecycle, + 1, + &tasks, + pending_task(rejected_dropped.clone()), + )); + assert!(rejected_dropped.load(Ordering::Acquire)); + + let current_dropped = Arc::new(AtomicBool::new(false)); + assert!(spawn_tcp_proxy_task( + &lifecycle, + 3, + &tasks, + pending_task(current_dropped.clone()), + )); + let mut current = std::mem::take(&mut *tasks.lock().unwrap()); + current.shutdown().await; + assert!(current_dropped.load(Ordering::Acquire)); + } +} diff --git a/easytier-core/src/gateway/proxy/tcp_socket_connector.rs b/easytier-core/src/gateway/proxy/tcp_socket_connector.rs new file mode 100644 index 00000000..f17ff9a0 --- /dev/null +++ b/easytier-core/src/gateway/proxy/tcp_socket_connector.rs @@ -0,0 +1,62 @@ +use std::{net::SocketAddr, sync::Arc, time::Duration}; + +use anyhow::Context; + +use crate::{ + foundation::time::timeout, + socket::{ + SocketContext, + tcp::{TcpConnectOptions, VirtualTcpSocketFactory}, + }, +}; + +use super::{tcp_proxy_engine::TcpProxyMode, traits::TcpProxyDestinationConnector}; + +pub struct TcpSocketProxyConnector { + socket_factory: Arc, + socket_context: SocketContext, +} + +impl TcpSocketProxyConnector { + pub fn new(socket_factory: Arc) -> Self { + Self { + socket_factory, + socket_context: SocketContext::default(), + } + } + + pub fn with_socket_context(mut self, socket_context: SocketContext) -> Self { + self.socket_context = socket_context; + self + } +} + +#[async_trait::async_trait] +impl TcpProxyDestinationConnector for TcpSocketProxyConnector { + type DstStream = F::Socket; + + async fn connect(&self, _src: SocketAddr, dst: SocketAddr) -> anyhow::Result { + timeout( + Duration::from_secs(10), + self.socket_factory.connect_tcp( + TcpConnectOptions::proxy_nat(dst).with_bind( + crate::socket::tcp::TcpBindOptions::default().with_context( + self.socket_context + .clone() + .with_ip_version(if dst.is_ipv4() { + crate::socket::IpVersion::V4 + } else { + crate::socket::IpVersion::V6 + }), + ), + ), + ), + ) + .await? + .with_context(|| format!("connect to nat dst failed: {dst:?}")) + } + + fn proxy_mode(&self) -> TcpProxyMode { + TcpProxyMode::Tcp + } +} diff --git a/easytier-core/src/gateway/proxy/traits.rs b/easytier-core/src/gateway/proxy/traits.rs new file mode 100644 index 00000000..4aee5d0d --- /dev/null +++ b/easytier-core/src/gateway/proxy/traits.rs @@ -0,0 +1,98 @@ +use std::net::{IpAddr, Ipv4Addr, SocketAddr}; +use std::sync::{Arc, Weak}; + +use bytes::Bytes; +use cidr::Ipv4Inet; +use tokio::io::{AsyncRead, AsyncWrite}; + +use super::tcp_proxy_engine::TcpProxyMode; +use super::udp_proxy_engine::UdpNatEntryId; + +pub use super::icmp_host::{IcmpProxyHost, IcmpProxySocket, ProxyRuntimeError}; + +#[derive(Clone, Copy, Debug, Default)] +pub(crate) struct ProxyRuntimeSnapshot { + pub local_inet: Option, + pub virtual_ipv4: Option, + pub no_tun: bool, + pub enable_exit_node: bool, + pub smoltcp_enabled: bool, + pub latency_first: bool, +} + +pub(crate) trait ProxyRuntimeInfo: Send + Sync { + fn proxy_runtime_snapshot(&self) -> ProxyRuntimeSnapshot; + fn is_ip_local_virtual_ip(&self, ip: &IpAddr) -> bool; +} + +pub(crate) trait WrappedTcpDestinationRuntime: Send + Sync { + fn is_ip_local_virtual_ip(&self, ip: &IpAddr) -> bool; + fn no_tun(&self) -> bool; + fn should_deny_tcp_proxy(&self, dst: SocketAddr) -> bool; +} + +#[async_trait::async_trait] +pub(crate) trait IcmpProxyRuntime: ProxyRuntimeInfo { + type Socket: IcmpProxySocket + ?Sized; + + async fn start_icmp(&self) -> Result, ProxyRuntimeError>; + + fn stop_icmp(&self); +} + +#[async_trait::async_trait] +pub(crate) trait UdpProxyResponseSink: Send + Sync { + async fn handle_socket_response( + &self, + entry_id: UdpNatEntryId, + src: SocketAddr, + payload: Bytes, + ); +} + +pub(crate) trait UdpProxyPolicy: ProxyRuntimeInfo { + fn should_deny_udp_proxy(&self, dst: SocketAddr) -> bool; + fn udp_response_ipv4_mtu(&self) -> usize; +} + +#[async_trait::async_trait] +pub(crate) trait UdpProxyRuntime: ProxyRuntimeInfo { + fn should_deny_udp_proxy(&self, dst: SocketAddr) -> bool; + fn udp_response_ipv4_mtu(&self) -> usize; + + async fn send_udp_to_socket( + &self, + entry_id: UdpNatEntryId, + dst: SocketAddr, + payload: Bytes, + response_sink: Weak, + ) -> Result<(), ProxyRuntimeError>; + + fn close_udp_socket(&self, entry_id: UdpNatEntryId); +} + +pub trait TcpProxyStream: AsyncRead + AsyncWrite + Unpin + Send {} + +impl TcpProxyStream for T where T: AsyncRead + AsyncWrite + Unpin + Send {} + +#[derive(Debug, Clone, Copy)] +pub(crate) struct TcpProxyConnectContext { + pub src: SocketAddr, + pub real_dst: SocketAddr, + pub mapped_dst: SocketAddr, +} + +pub(crate) trait TcpProxyRuntime: ProxyRuntimeInfo { + fn should_deny_tcp_proxy(&self, dst: SocketAddr) -> bool; + + fn record_tcp_proxy_connect(&self, ctx: TcpProxyConnectContext, socket_dst: SocketAddr); +} + +#[async_trait::async_trait] +pub(crate) trait TcpProxyDestinationConnector: Send + Sync + 'static { + type DstStream: TcpProxyStream + 'static; + + async fn connect(&self, src: SocketAddr, dst: SocketAddr) -> anyhow::Result; + + fn proxy_mode(&self) -> TcpProxyMode; +} diff --git a/easytier-core/src/gateway/proxy/udp_proxy_engine.rs b/easytier-core/src/gateway/proxy/udp_proxy_engine.rs new file mode 100644 index 00000000..6873ecd9 --- /dev/null +++ b/easytier-core/src/gateway/proxy/udp_proxy_engine.rs @@ -0,0 +1,507 @@ +use std::{ + net::{IpAddr, Ipv4Addr, SocketAddr, SocketAddrV4}, + sync::{ + Arc, + atomic::{AtomicBool, AtomicU16, Ordering}, + }, + time::{Duration, Instant}, +}; + +use bytes::Bytes; +use dashmap::{DashMap, mapref::entry::Entry}; +use smoltcp::wire::{IpAddress, IpProtocol, Ipv4Packet, UdpPacket}; + +use crate::{ + config::PeerId, + packet::{PacketType, ZCPacket}, +}; + +use super::{ + cidr_table::ProxyCidrTable, + ip_reassembler::{ComposeIpv4PacketArgs, IpReassembler, compose_ipv4_packet}, + traits::UdpProxyRuntime, +}; + +#[derive(Debug, Hash, Eq, PartialEq, Clone, Copy)] +pub struct UdpNatKey { + pub src_socket: SocketAddr, + pub dst_socket: SocketAddr, +} + +impl UdpNatKey { + pub fn new(src_socket: SocketAddr, dst_socket: SocketAddr) -> Self { + Self { + src_socket, + dst_socket, + } + } +} + +#[derive(Debug, Clone, Copy, Hash, Eq, PartialEq)] +pub struct UdpNatEntryId(uuid::Uuid); + +impl UdpNatEntryId { + pub fn new() -> Self { + Self(uuid::Uuid::new_v4()) + } +} + +impl Default for UdpNatEntryId { + fn default() -> Self { + Self::new() + } +} + +#[derive(Debug)] +pub struct UdpNatEntry { + id: UdpNatEntryId, + src_peer_id: PeerId, + my_peer_id: PeerId, + src_socket: SocketAddr, + real_dst_ip: Ipv4Addr, + mapped_dst_ip: Ipv4Addr, + virtual_ipv4: Ipv4Addr, + last_active_time: parking_lot::Mutex, + stopped: AtomicBool, + denied: bool, +} + +impl UdpNatEntry { + fn new( + src_peer_id: PeerId, + my_peer_id: PeerId, + src_socket: SocketAddr, + real_dst_ip: Ipv4Addr, + mapped_dst_ip: Ipv4Addr, + virtual_ipv4: Ipv4Addr, + denied: bool, + ) -> Self { + Self { + id: UdpNatEntryId::new(), + src_peer_id, + my_peer_id, + src_socket, + real_dst_ip, + mapped_dst_ip, + virtual_ipv4, + last_active_time: parking_lot::Mutex::new(Instant::now()), + stopped: AtomicBool::new(false), + denied, + } + } + + pub fn id(&self) -> UdpNatEntryId { + self.id + } + + pub fn is_denied(&self) -> bool { + self.denied + } + + pub fn stop(&self) { + self.stopped.store(true, Ordering::Relaxed); + } + + pub fn is_stopped(&self) -> bool { + self.stopped.load(Ordering::Relaxed) + } + + fn mark_active(&self) { + *self.last_active_time.lock() = Instant::now(); + } + + fn is_active(&self, ttl: Duration) -> bool { + self.last_active_time.lock().elapsed() < ttl && !self.is_stopped() + } +} + +#[derive(Clone, Copy, Debug)] +pub struct UdpProxyPeerContext { + pub virtual_ipv4: Option, + pub enable_exit_node: bool, + pub no_tun: bool, +} + +#[derive(Debug)] +pub enum UdpProxyAction { + ForwardToSocket { + entry_id: UdpNatEntryId, + dst: SocketAddr, + payload: Bytes, + }, + Drop, + Pass, +} + +#[derive(Debug)] +pub struct UdpProxyEngine { + cidr_table: Arc, + nat_table: DashMap>, + nat_ids: DashMap, + ip_reassembler: IpReassembler, + entry_ttl: Duration, + next_ip_id: AtomicU16, +} + +impl UdpProxyEngine { + pub fn new(cidr_table: Arc, fragment_timeout: Duration) -> Self { + Self { + cidr_table, + nat_table: DashMap::new(), + nat_ids: DashMap::new(), + ip_reassembler: IpReassembler::new(fragment_timeout), + entry_ttl: Duration::from_secs(180), + next_ip_id: AtomicU16::new(1), + } + } + + pub fn entry_ids(&self) -> Vec { + self.nat_ids.iter().map(|entry| *entry.key()).collect() + } + + pub fn remove_expired_entries(&self) -> Vec { + let mut removed = Vec::new(); + self.nat_table.retain(|_, entry| { + if entry.is_active(self.entry_ttl) { + true + } else { + tracing::info!(?entry, "udp nat table entry removed"); + entry.stop(); + self.nat_ids.remove(&entry.id()); + removed.push(entry.id()); + false + } + }); + self.nat_table.shrink_to_fit(); + self.nat_ids.shrink_to_fit(); + removed + } + + pub fn remove_entry(&self, entry_id: UdpNatEntryId) { + if let Some((_, key)) = self.nat_ids.remove(&entry_id) + && let Some((_, entry)) = self.nat_table.remove(&key) + { + entry.stop(); + } + } + + pub fn remove_expired_fragments(&self) { + self.ip_reassembler.remove_expired_packets(); + } + + pub fn handle_peer_packet( + &self, + packet: &ZCPacket, + ctx: UdpProxyPeerContext, + runtime: &impl UdpProxyRuntime, + ) -> UdpProxyAction { + if self.cidr_table.is_empty() && !ctx.enable_exit_node && !ctx.no_tun { + return UdpProxyAction::Pass; + } + + let Some(virtual_ipv4) = ctx.virtual_ipv4 else { + return UdpProxyAction::Pass; + }; + let Some(hdr) = packet.peer_manager_header() else { + return UdpProxyAction::Pass; + }; + let is_exit_node = hdr.is_exit_node(); + if hdr.packet_type != PacketType::Data as u8 || hdr.is_no_proxy() { + return UdpProxyAction::Pass; + }; + + let Ok(ipv4) = Ipv4Packet::new_checked(packet.payload()) else { + return UdpProxyAction::Pass; + }; + if ipv4.version() != 4 || ipv4.next_header() != IpProtocol::Udp { + return UdpProxyAction::Pass; + } + + let origin_dst_ip = ipv4.dst_addr(); + let mut real_dst_ip = origin_dst_ip; + let no_tun_local_virtual_ip = + ctx.no_tun && Some(origin_dst_ip) == ctx.virtual_ipv4.as_ref().copied(); + if let Some(mapped_real_ip) = self.cidr_table.lookup_v4(origin_dst_ip) { + real_dst_ip = mapped_real_ip; + } else if !is_exit_node && !no_tun_local_virtual_ip { + return UdpProxyAction::Pass; + } + + let reassembled_buf; + let udp_packet = if IpReassembler::is_packet_fragmented(&ipv4) { + let Some(buf) = self.ip_reassembler.add_fragment(&ipv4) else { + return UdpProxyAction::Drop; + }; + reassembled_buf = buf; + let Ok(udp_packet) = UdpPacket::new_checked(reassembled_buf.as_slice()) else { + return UdpProxyAction::Pass; + }; + udp_packet + } else { + let Ok(udp_packet) = UdpPacket::new_checked(ipv4.payload()) else { + return UdpProxyAction::Pass; + }; + udp_packet + }; + + let dst_socket = if runtime.is_ip_local_virtual_ip(&IpAddr::V4(real_dst_ip)) { + SocketAddr::new(Ipv4Addr::LOCALHOST.into(), udp_packet.dst_port()) + } else { + SocketAddr::new(real_dst_ip.into(), udp_packet.dst_port()) + }; + + tracing::trace!( + ?packet, + ?ipv4, + ?udp_packet, + "udp nat packet request received" + ); + + let nat_key = UdpNatKey::new( + SocketAddr::new(ipv4.src_addr().into(), udp_packet.src_port()), + SocketAddr::new(origin_dst_ip.into(), udp_packet.dst_port()), + ); + let deny_dst = SocketAddr::new(real_dst_ip.into(), udp_packet.dst_port()); + let nat_entry = match self.nat_table.entry(nat_key) { + Entry::Occupied(entry) => entry.get().clone(), + Entry::Vacant(entry) => { + let denied = runtime.should_deny_udp_proxy(deny_dst); + let nat_entry = Arc::new(UdpNatEntry::new( + hdr.from_peer_id.get(), + hdr.to_peer_id.get(), + nat_key.src_socket, + real_dst_ip, + origin_dst_ip, + virtual_ipv4, + denied, + )); + self.nat_ids.insert(nat_entry.id(), nat_key); + tracing::info!(?packet, ?ipv4, ?udp_packet, "udp nat table entry created"); + entry.insert(nat_entry).clone() + } + }; + + if nat_entry.is_denied() { + tracing::debug!( + dst_port = udp_packet.dst_port(), + "dst socket is in running listeners, ignore it" + ); + return UdpProxyAction::Drop; + } + + nat_entry.mark_active(); + UdpProxyAction::ForwardToSocket { + entry_id: nat_entry.id(), + dst: dst_socket, + payload: Bytes::copy_from_slice(udp_packet.payload()), + } + } + + pub fn handle_socket_response( + &self, + entry_id: UdpNatEntryId, + src_socket: SocketAddr, + payload: &[u8], + ipv4_mtu: usize, + ) -> anyhow::Result> { + let Some(key) = self.nat_ids.get(&entry_id).map(|entry| *entry.value()) else { + return Ok(Vec::new()); + }; + let Some(entry) = self.nat_table.get(&key).map(|entry| entry.clone()) else { + return Ok(Vec::new()); + }; + entry.mark_active(); + + let SocketAddr::V4(mut src_v4) = src_socket else { + return Ok(Vec::new()); + }; + let SocketAddr::V4(nat_src_v4) = entry.src_socket else { + return Ok(Vec::new()); + }; + + let has_mapped_dst = entry.real_dst_ip != entry.mapped_dst_ip; + let mut reply_src_ip = *src_v4.ip(); + if has_mapped_dst && reply_src_ip == entry.real_dst_ip { + reply_src_ip = entry.mapped_dst_ip; + } else if reply_src_ip.is_loopback() { + reply_src_ip = entry.virtual_ipv4; + } + if has_mapped_dst && reply_src_ip == entry.real_dst_ip { + reply_src_ip = entry.mapped_dst_ip; + } + src_v4.set_ip(reply_src_ip); + + let payload_mtu = ipv4_mtu + .saturating_sub(smoltcp::wire::IPV4_HEADER_LEN) + .max(8); + let payload_mtu = payload_mtu - (payload_mtu % 8); + let ip_id = self.next_ip_id.fetch_add(1, Ordering::Relaxed); + compose_udp_ipv4_response(&entry, &src_v4, &nat_src_v4, payload, payload_mtu, ip_id) + } +} + +fn compose_udp_ipv4_response( + entry: &UdpNatEntry, + src_v4: &SocketAddrV4, + nat_src_v4: &SocketAddrV4, + payload: &[u8], + payload_mtu: usize, + ip_id: u16, +) -> anyhow::Result> { + assert_eq!(0, payload_mtu % 8); + + let mut buf = + vec![0; smoltcp::wire::IPV4_HEADER_LEN + smoltcp::wire::UDP_HEADER_LEN + payload.len()]; + let udp_start = smoltcp::wire::IPV4_HEADER_LEN; + let udp_len = smoltcp::wire::UDP_HEADER_LEN + payload.len(); + { + let mut udp_packet = UdpPacket::new_unchecked(&mut buf[udp_start..udp_start + udp_len]); + udp_packet.set_src_port(src_v4.port()); + udp_packet.set_dst_port(nat_src_v4.port()); + udp_packet.set_len(udp_len as u16); + udp_packet.payload_mut().copy_from_slice(payload); + udp_packet.fill_checksum( + &IpAddress::Ipv4(*src_v4.ip()), + &IpAddress::Ipv4(*nat_src_v4.ip()), + ); + } + + let mut packets = Vec::new(); + compose_ipv4_packet( + ComposeIpv4PacketArgs { + buf: &mut buf, + src_v4: src_v4.ip(), + dst_v4: nat_src_v4.ip(), + next_protocol: IpProtocol::Udp, + payload_len: udp_len, + payload_mtu, + ip_id, + }, + |buf| { + let mut packet = ZCPacket::new_with_payload(buf); + packet.fill_peer_manager_hdr( + entry.my_peer_id, + entry.src_peer_id, + PacketType::Data as u8, + ); + packet + .mut_peer_manager_header() + .expect("peer manager header") + .set_no_proxy(true); + packets.push(packet); + Ok(()) + }, + )?; + + Ok(packets) +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::gateway::proxy::{ + cidr_table::{ProxyCidrRule, ProxyCidrSnapshot}, + traits::{ProxyRuntimeError, ProxyRuntimeInfo, ProxyRuntimeSnapshot, UdpProxyResponseSink}, + }; + + struct TestRuntime; + + impl ProxyRuntimeInfo for TestRuntime { + fn proxy_runtime_snapshot(&self) -> ProxyRuntimeSnapshot { + ProxyRuntimeSnapshot::default() + } + + fn is_ip_local_virtual_ip(&self, ip: &IpAddr) -> bool { + matches!(ip, IpAddr::V4(ip) if *ip == Ipv4Addr::new(10, 144, 144, 204)) + } + } + + #[async_trait::async_trait] + impl UdpProxyRuntime for TestRuntime { + fn should_deny_udp_proxy(&self, _dst_socket: SocketAddr) -> bool { + false + } + + fn udp_response_ipv4_mtu(&self) -> usize { + 1280 + } + + async fn send_udp_to_socket( + &self, + _entry_id: UdpNatEntryId, + _dst: SocketAddr, + _payload: bytes::Bytes, + _response_sink: std::sync::Weak, + ) -> Result<(), ProxyRuntimeError> { + Ok(()) + } + + fn close_udp_socket(&self, _entry_id: UdpNatEntryId) {} + } + + #[test] + fn socket_response_uses_mapped_source_for_mapped_destination() { + let table = Arc::new(ProxyCidrTable::from_snapshot(ProxyCidrSnapshot { + rules: vec![ProxyCidrRule { + cidr: "10.144.144.204/32".parse().unwrap(), + mapped_cidr: Some("10.10.10.3/32".parse().unwrap()), + }], + })); + let engine = UdpProxyEngine::new(table, Duration::from_secs(10)); + + let mut request = + vec![0; smoltcp::wire::IPV4_HEADER_LEN + smoltcp::wire::UDP_HEADER_LEN + 7]; + { + let mut ipv4 = Ipv4Packet::new_unchecked(&mut request); + ipv4.set_version(4); + ipv4.set_header_len(smoltcp::wire::IPV4_HEADER_LEN as u8); + ipv4.set_total_len( + (smoltcp::wire::IPV4_HEADER_LEN + smoltcp::wire::UDP_HEADER_LEN + 7) as u16, + ); + ipv4.set_hop_limit(64); + ipv4.set_next_header(IpProtocol::Udp); + ipv4.set_src_addr("10.144.144.206".parse().unwrap()); + ipv4.set_dst_addr("10.10.10.3".parse().unwrap()); + ipv4.fill_checksum(); + } + { + let mut udp = UdpPacket::new_unchecked(&mut request[smoltcp::wire::IPV4_HEADER_LEN..]); + udp.set_src_port(53864); + udp.set_dst_port(12345); + udp.set_len((smoltcp::wire::UDP_HEADER_LEN + 7) as u16); + udp.payload_mut().copy_from_slice(b"request"); + udp.fill_checksum( + &IpAddress::Ipv4("10.144.144.206".parse().unwrap()), + &IpAddress::Ipv4("10.10.10.3".parse().unwrap()), + ); + } + + let mut zc = ZCPacket::new_with_payload(&request); + zc.fill_peer_manager_hdr(1, 2, PacketType::Data as u8); + + let action = engine.handle_peer_packet( + &zc, + UdpProxyPeerContext { + virtual_ipv4: Some("10.144.144.204".parse().unwrap()), + enable_exit_node: false, + no_tun: false, + }, + &TestRuntime, + ); + let UdpProxyAction::ForwardToSocket { entry_id, .. } = action else { + panic!("expected forward action"); + }; + + let packets = engine + .handle_socket_response(entry_id, "127.0.0.1:12345".parse().unwrap(), b"reply", 1280) + .unwrap(); + assert_eq!(packets.len(), 1); + let ipv4 = Ipv4Packet::new_checked(packets[0].payload()).unwrap(); + assert_eq!(ipv4.src_addr(), Ipv4Addr::new(10, 10, 10, 3)); + assert_eq!(ipv4.dst_addr(), Ipv4Addr::new(10, 144, 144, 206)); + let udp = UdpPacket::new_checked(ipv4.payload()).unwrap(); + assert_eq!(udp.src_port(), 12345); + assert_eq!(udp.dst_port(), 53864); + assert_eq!(udp.payload(), b"reply"); + } +} diff --git a/easytier-core/src/gateway/proxy/udp_proxy_service.rs b/easytier-core/src/gateway/proxy/udp_proxy_service.rs new file mode 100644 index 00000000..e06d4e1a --- /dev/null +++ b/easytier-core/src/gateway/proxy/udp_proxy_service.rs @@ -0,0 +1,217 @@ +use std::sync::{ + Arc, Weak, + atomic::{AtomicBool, Ordering}, +}; +use std::time::Duration; + +use bytes::Bytes; +use tokio::sync::mpsc::{self, Receiver, Sender}; +use tokio::task::JoinSet; + +use crate::packet::ZCPacket; +use crate::peers::PeerPacketFilter; +use crate::peers::peer_manager::{PeerManagerCore, PipelineRegistrationGuard}; + +use super::cidr_table::ProxyCidrTable; +use super::traits::{UdpProxyResponseSink, UdpProxyRuntime}; +use super::udp_proxy_engine::{UdpNatEntryId, UdpProxyAction, UdpProxyEngine, UdpProxyPeerContext}; + +pub struct UdpProxyService { + peer_manager: Arc, + runtime: Arc, + engine: Arc, + response_tx: Sender, + response_rx: std::sync::Mutex>>, + pipeline_guard: std::sync::Mutex>, + tasks: std::sync::Mutex>, + started: AtomicBool, +} + +impl UdpProxyService { + pub fn new( + peer_manager: Arc, + runtime: Arc, + cidr_table: Arc, + fragment_timeout: Duration, + ) -> Arc { + let (response_tx, response_rx) = mpsc::channel(1024); + Arc::new(Self { + peer_manager, + runtime, + engine: Arc::new(UdpProxyEngine::new(cidr_table, fragment_timeout)), + response_tx, + response_rx: std::sync::Mutex::new(Some(response_rx)), + pipeline_guard: std::sync::Mutex::new(None), + tasks: std::sync::Mutex::new(JoinSet::new()), + started: AtomicBool::new(false), + }) + } + + pub async fn start(self: &Arc) { + if self.started.swap(true, Ordering::AcqRel) { + return; + } + + let guard = self + .peer_manager + .add_managed_packet_process_pipeline(Box::new(UdpProxyServiceFilter { + service: Arc::downgrade(self), + })) + .await; + self.pipeline_guard.lock().unwrap().replace(guard); + + if let Some(mut response_rx) = self.response_rx.lock().unwrap().take() { + let service = Arc::downgrade(self); + self.tasks.lock().unwrap().spawn(async move { + while let Some(mut packet) = response_rx.recv().await { + let Some(service) = service.upgrade() else { + break; + }; + let latency_first = service.runtime.proxy_runtime_snapshot().latency_first; + let Some(hdr) = packet.mut_peer_manager_header() else { + continue; + }; + hdr.set_latency_first(latency_first); + let dst_peer_id = hdr.to_peer_id.into(); + tracing::trace!(?packet, ?dst_peer_id, "udp nat packet response send"); + if let Err(err) = service + .peer_manager + .send_msg_for_proxy(packet, dst_peer_id) + .await + { + tracing::error!(?err, "send udp proxy response to peer failed"); + } + } + }); + } + + let service = Arc::downgrade(self); + self.tasks.lock().unwrap().spawn(async move { + loop { + crate::foundation::time::sleep(Duration::from_secs(15)).await; + let Some(service) = service.upgrade() else { + break; + }; + for entry_id in service.engine.remove_expired_entries() { + service.runtime.close_udp_socket(entry_id); + } + } + }); + + let service = Arc::downgrade(self); + self.tasks.lock().unwrap().spawn(async move { + loop { + crate::foundation::time::sleep(Duration::from_secs(1)).await; + let Some(service) = service.upgrade() else { + break; + }; + service.engine.remove_expired_fragments(); + } + }); + } + + pub fn stop(&self) { + if !self.started.swap(false, Ordering::AcqRel) { + return; + } + if let Some(guard) = self.pipeline_guard.lock().unwrap().take() { + guard.close(); + } + self.tasks.lock().unwrap().abort_all(); + for entry_id in self.engine.entry_ids() { + self.engine.remove_entry(entry_id); + self.runtime.close_udp_socket(entry_id); + } + } + + async fn handle_peer_packet(self: Arc, packet: ZCPacket) -> Option { + let snapshot = self.runtime.proxy_runtime_snapshot(); + let action = self.engine.handle_peer_packet( + &packet, + UdpProxyPeerContext { + virtual_ipv4: snapshot.virtual_ipv4, + enable_exit_node: snapshot.enable_exit_node, + no_tun: snapshot.no_tun, + }, + self.runtime.as_ref(), + ); + + let UdpProxyAction::ForwardToSocket { + entry_id, + dst, + payload, + } = action + else { + return matches!(action, UdpProxyAction::Pass).then_some(packet); + }; + + let sink: Arc = self.clone(); + if let Err(err) = self + .runtime + .send_udp_to_socket(entry_id, dst, payload, Arc::downgrade(&sink)) + .await + { + tracing::error!(?err, ?entry_id, "udp proxy runtime send failed"); + self.engine.remove_entry(entry_id); + self.runtime.close_udp_socket(entry_id); + } + + None + } +} + +impl Drop for UdpProxyService { + fn drop(&mut self) { + self.stop(); + } +} + +#[async_trait::async_trait] +impl UdpProxyResponseSink for UdpProxyService { + async fn handle_socket_response( + &self, + entry_id: UdpNatEntryId, + src: std::net::SocketAddr, + payload: Bytes, + ) { + let packets = match self.engine.handle_socket_response( + entry_id, + src, + payload.as_ref(), + self.runtime.udp_response_ipv4_mtu(), + ) { + Ok(packets) => packets, + Err(err) => { + tracing::error!(?err, ?entry_id, "compose udp response packet failed"); + self.engine.remove_entry(entry_id); + self.runtime.close_udp_socket(entry_id); + return; + } + }; + + for mut packet in packets { + let Some(hdr) = packet.mut_peer_manager_header() else { + continue; + }; + let dst_peer_id: crate::config::PeerId = hdr.to_peer_id.into(); + tracing::trace!(?packet, ?dst_peer_id, "udp nat packet response queued"); + if let Err(err) = self.response_tx.try_send(packet) { + tracing::error!(?err, ?dst_peer_id, "queue udp proxy response failed"); + } + } + } +} + +struct UdpProxyServiceFilter { + service: Weak>, +} + +#[async_trait::async_trait] +impl PeerPacketFilter for UdpProxyServiceFilter { + async fn try_process_packet_from_peer(&self, packet: ZCPacket) -> Option { + let Some(service) = self.service.upgrade() else { + return Some(packet); + }; + service.handle_peer_packet(packet).await + } +} diff --git a/easytier-core/src/gateway/proxy/udp_socket_runtime.rs b/easytier-core/src/gateway/proxy/udp_socket_runtime.rs new file mode 100644 index 00000000..56c2f2a3 --- /dev/null +++ b/easytier-core/src/gateway/proxy/udp_socket_runtime.rs @@ -0,0 +1,647 @@ +use std::{ + net::{IpAddr, SocketAddr}, + sync::{ + Arc, Weak, + atomic::{AtomicBool, Ordering}, + }, + time::Duration, +}; + +use bytes::Bytes; +use dashmap::{DashMap, mapref::entry::Entry}; +use tokio::sync::Notify; +use tokio_util::task::AbortOnDropHandle; + +use crate::{ + foundation::time, + socket::udp::{UdpBindOptions, VirtualUdpSocket, VirtualUdpSocketFactory}, +}; + +use super::{ + traits::{ + ProxyRuntimeError, ProxyRuntimeInfo, ProxyRuntimeSnapshot, UdpProxyPolicy, + UdpProxyResponseSink, UdpProxyRuntime, + }, + udp_proxy_engine::UdpNatEntryId, +}; + +const UDP_PROXY_RECEIVE_BUFFER_SIZE: usize = 64 * 1024; + +struct UdpSocketNatEntry +where + S: VirtualUdpSocket, +{ + socket: Arc, + receive_task: std::sync::Mutex>>, + closed: AtomicBool, +} + +struct UdpSocketEntryState +where + S: VirtualUdpSocket, +{ + entry: Option>>, + error: Option, +} + +struct UdpSocketEntrySlot +where + S: VirtualUdpSocket, +{ + state: std::sync::Mutex>, + changed: Notify, + closed: AtomicBool, +} + +impl UdpSocketEntrySlot +where + S: VirtualUdpSocket, +{ + fn new() -> Arc { + Arc::new(Self { + state: std::sync::Mutex::new(UdpSocketEntryState { + entry: None, + error: None, + }), + changed: Notify::new(), + closed: AtomicBool::new(false), + }) + } + + async fn wait_for_entry(&self) -> Result>, ProxyRuntimeError> { + loop { + let changed = self.changed.notified(); + { + let state = self.state.lock().unwrap(); + if let Some(entry) = &state.entry { + return Ok(entry.clone()); + } + if let Some(error) = &state.error { + return Err(ProxyRuntimeError::Other(anyhow::anyhow!(error.clone()))); + } + if self.closed.load(Ordering::Acquire) { + return Err(ProxyRuntimeError::Other(anyhow::anyhow!( + "UDP proxy socket entry was closed while being created" + ))); + } + } + changed.await; + } + } + + fn fail(&self, error: &ProxyRuntimeError) { + self.state.lock().unwrap().error = Some(error.to_string()); + self.changed.notify_waiters(); + } + + fn cancel_creation(&self) { + self.closed.store(true, Ordering::Release); + self.state.lock().unwrap().error = + Some("UDP proxy socket creation was cancelled".to_owned()); + self.changed.notify_waiters(); + } + + fn close(&self) { + self.closed.store(true, Ordering::Release); + if let Some(entry) = self.state.lock().unwrap().entry.take() { + entry.stop(); + } + self.changed.notify_waiters(); + } +} + +fn remove_entry_slot( + entries: &DashMap>>, + entry_id: UdpNatEntryId, + slot: &Arc>, +) where + S: VirtualUdpSocket, +{ + if let Entry::Occupied(entry) = entries.entry(entry_id) + && Arc::ptr_eq(entry.get(), slot) + { + entry.remove(); + } +} + +struct UdpSocketCreationGuard +where + S: VirtualUdpSocket, +{ + entries: Arc>>>, + entry_id: UdpNatEntryId, + slot: Arc>, + armed: bool, +} + +impl UdpSocketCreationGuard +where + S: VirtualUdpSocket, +{ + fn new( + entries: Arc>>>, + entry_id: UdpNatEntryId, + slot: Arc>, + ) -> Self { + Self { + entries, + entry_id, + slot, + armed: true, + } + } + + fn disarm(&mut self) { + self.armed = false; + } +} + +impl Drop for UdpSocketCreationGuard +where + S: VirtualUdpSocket, +{ + fn drop(&mut self) { + if !self.armed { + return; + } + self.slot.cancel_creation(); + remove_entry_slot(&self.entries, self.entry_id, &self.slot); + } +} + +impl UdpSocketNatEntry +where + S: VirtualUdpSocket, +{ + fn new(socket: Arc) -> Arc { + Arc::new(Self { + socket, + receive_task: std::sync::Mutex::new(None), + closed: AtomicBool::new(false), + }) + } + + fn start_receive_task( + self: &Arc, + entry_id: UdpNatEntryId, + response_sink: Weak, + receive_timeout: Duration, + ) { + if self.closed.load(Ordering::Acquire) { + return; + } + + let socket = self.socket.clone(); + let task = AbortOnDropHandle::new(tokio::spawn(async move { + loop { + let mut buffer = vec![0; UDP_PROXY_RECEIVE_BUFFER_SIZE]; + let (length, source) = + match time::timeout(receive_timeout, socket.recv_from(&mut buffer)).await { + Ok(Ok(received)) => received, + Ok(Err(error)) => { + tracing::error!(?error, ?entry_id, "UDP proxy receive failed"); + break; + } + Err(error) => { + tracing::error!(?error, ?entry_id, "UDP proxy receive timed out"); + break; + } + }; + + let Some(response_sink) = response_sink.upgrade() else { + break; + }; + response_sink + .handle_socket_response( + entry_id, + source, + Bytes::copy_from_slice(&buffer[..length]), + ) + .await; + } + })); + + let mut receive_task = self.receive_task.lock().unwrap(); + if self.closed.load(Ordering::Acquire) { + drop(task); + } else { + receive_task.replace(task); + } + } + + fn stop(&self) { + self.closed.store(true, Ordering::Release); + self.receive_task.lock().unwrap().take(); + } +} + +pub struct UdpSocketProxyRuntime +where + F: VirtualUdpSocketFactory, + P: UdpProxyPolicy, +{ + factory: Arc, + policy: Arc

, + bind_options: UdpBindOptions, + receive_timeout: Duration, + entries: Arc>>>, + closing: AtomicBool, +} + +impl UdpSocketProxyRuntime +where + F: VirtualUdpSocketFactory, + P: UdpProxyPolicy, +{ + pub fn new( + factory: Arc, + policy: Arc

, + bind_options: UdpBindOptions, + receive_timeout: Duration, + ) -> Self { + Self { + factory, + policy, + bind_options, + receive_timeout, + entries: Arc::new(DashMap::new()), + closing: AtomicBool::new(false), + } + } + + async fn ensure_socket_entry( + &self, + entry_id: UdpNatEntryId, + response_sink: Weak, + ) -> Result>, ProxyRuntimeError> { + if self.closing.load(Ordering::Acquire) { + return Err(ProxyRuntimeError::Other(anyhow::anyhow!( + "UDP proxy runtime is closing" + ))); + } + + let (slot, create) = match self.entries.entry(entry_id) { + Entry::Occupied(entry) => (entry.get().clone(), false), + Entry::Vacant(entry) => { + let slot = UdpSocketEntrySlot::new(); + entry.insert(slot.clone()); + (slot, true) + } + }; + if self.closing.load(Ordering::Acquire) { + slot.close(); + self.remove_slot(entry_id, &slot); + return Err(ProxyRuntimeError::Other(anyhow::anyhow!( + "UDP proxy runtime is closing" + ))); + } + if !create { + return slot.wait_for_entry().await; + } + + let mut creation = + UdpSocketCreationGuard::new(self.entries.clone(), entry_id, slot.clone()); + + let socket = match self.factory.bind_udp(self.bind_options.clone()).await { + Ok(socket) => socket, + Err(error) => { + let error = ProxyRuntimeError::Other(error); + slot.fail(&error); + self.remove_slot(entry_id, &slot); + creation.disarm(); + return Err(error); + } + }; + let candidate = UdpSocketNatEntry::new(socket); + { + let mut state = slot.state.lock().unwrap(); + if slot.closed.load(Ordering::Acquire) || self.closing.load(Ordering::Acquire) { + candidate.stop(); + drop(state); + slot.changed.notify_waiters(); + self.remove_slot(entry_id, &slot); + creation.disarm(); + return Err(ProxyRuntimeError::Other(anyhow::anyhow!( + "UDP proxy socket entry was closed while being created" + ))); + } + candidate.start_receive_task(entry_id, response_sink, self.receive_timeout); + state.entry = Some(candidate.clone()); + } + slot.changed.notify_waiters(); + creation.disarm(); + Ok(candidate) + } + + fn remove_slot(&self, entry_id: UdpNatEntryId, slot: &Arc>) { + remove_entry_slot(&self.entries, entry_id, slot); + } + + pub fn close_all(&self) { + self.closing.store(true, Ordering::Release); + for slot in self.entries.iter() { + slot.close(); + } + self.entries.clear(); + self.entries.shrink_to_fit(); + } +} + +impl ProxyRuntimeInfo for UdpSocketProxyRuntime +where + F: VirtualUdpSocketFactory, + P: UdpProxyPolicy, +{ + fn proxy_runtime_snapshot(&self) -> ProxyRuntimeSnapshot { + self.policy.proxy_runtime_snapshot() + } + + fn is_ip_local_virtual_ip(&self, ip: &IpAddr) -> bool { + self.policy.is_ip_local_virtual_ip(ip) + } +} + +#[async_trait::async_trait] +impl UdpProxyRuntime for UdpSocketProxyRuntime +where + F: VirtualUdpSocketFactory, + P: UdpProxyPolicy, +{ + fn should_deny_udp_proxy(&self, dst: SocketAddr) -> bool { + self.policy.should_deny_udp_proxy(dst) + } + + fn udp_response_ipv4_mtu(&self) -> usize { + self.policy.udp_response_ipv4_mtu() + } + + async fn send_udp_to_socket( + &self, + entry_id: UdpNatEntryId, + dst: SocketAddr, + payload: Bytes, + response_sink: Weak, + ) -> Result<(), ProxyRuntimeError> { + let entry = self.ensure_socket_entry(entry_id, response_sink).await?; + if entry.closed.load(Ordering::Acquire) { + return Err(ProxyRuntimeError::Other(anyhow::anyhow!( + "UDP proxy socket entry is closed" + ))); + } + entry.socket.send_to(&payload, dst).await?; + Ok(()) + } + + fn close_udp_socket(&self, entry_id: UdpNatEntryId) { + if let Some((_, slot)) = self.entries.remove(&entry_id) { + slot.close(); + } + self.entries.shrink_to_fit(); + } +} + +impl Drop for UdpSocketProxyRuntime +where + F: VirtualUdpSocketFactory, + P: UdpProxyPolicy, +{ + fn drop(&mut self) { + self.close_all(); + } +} + +#[cfg(test)] +mod tests { + use std::{ + io, + net::Ipv4Addr, + sync::atomic::{AtomicUsize, Ordering as AtomicOrdering}, + }; + + use super::*; + use crate::socket::udp::UdpSocketPurpose; + + #[derive(Default)] + struct RecordingSocket { + sent: std::sync::Mutex, SocketAddr)>>, + } + + #[async_trait::async_trait] + impl VirtualUdpSocket for RecordingSocket { + fn local_addr(&self) -> io::Result { + Ok("127.0.0.1:40000".parse().unwrap()) + } + + async fn send_to(&self, data: &[u8], addr: SocketAddr) -> io::Result { + self.sent.lock().unwrap().push((data.to_vec(), addr)); + Ok(data.len()) + } + + async fn recv_from(&self, _buf: &mut [u8]) -> io::Result<(usize, SocketAddr)> { + std::future::pending().await + } + } + + #[derive(Default)] + struct RecordingFactory { + options: std::sync::Mutex>, + socket: Arc, + } + + #[async_trait::async_trait] + impl VirtualUdpSocketFactory for RecordingFactory { + type Socket = RecordingSocket; + + async fn bind_udp(&self, options: UdpBindOptions) -> anyhow::Result> { + self.options.lock().unwrap().push(options); + Ok(self.socket.clone()) + } + } + + #[derive(Default)] + struct DelayedFactory { + socket: Arc, + bind_calls: AtomicUsize, + bind_started: Notify, + release_bind: Notify, + } + + #[async_trait::async_trait] + impl VirtualUdpSocketFactory for DelayedFactory { + type Socket = RecordingSocket; + + async fn bind_udp(&self, _options: UdpBindOptions) -> anyhow::Result> { + self.bind_calls.fetch_add(1, AtomicOrdering::AcqRel); + self.bind_started.notify_one(); + self.release_bind.notified().await; + Ok(self.socket.clone()) + } + } + + struct TestPolicy; + + impl ProxyRuntimeInfo for TestPolicy { + fn proxy_runtime_snapshot(&self) -> ProxyRuntimeSnapshot { + ProxyRuntimeSnapshot::default() + } + + fn is_ip_local_virtual_ip(&self, _ip: &IpAddr) -> bool { + false + } + } + + impl UdpProxyPolicy for TestPolicy { + fn should_deny_udp_proxy(&self, _dst: SocketAddr) -> bool { + false + } + + fn udp_response_ipv4_mtu(&self) -> usize { + 1280 + } + } + + struct NoopResponseSink; + + #[async_trait::async_trait] + impl UdpProxyResponseSink for NoopResponseSink { + async fn handle_socket_response( + &self, + _entry_id: UdpNatEntryId, + _src: SocketAddr, + _payload: Bytes, + ) { + } + } + + #[tokio::test] + async fn reuses_one_host_socket_per_nat_entry_and_recreates_after_close() { + let factory = Arc::new(RecordingFactory::default()); + let runtime = UdpSocketProxyRuntime::new( + factory.clone(), + Arc::new(TestPolicy), + UdpBindOptions::proxy_nat(), + Duration::from_secs(120), + ); + let sink: Arc = Arc::new(NoopResponseSink); + let entry_id = UdpNatEntryId::new(); + let destination = SocketAddr::from((Ipv4Addr::LOCALHOST, 53)); + + runtime + .send_udp_to_socket( + entry_id, + destination, + Bytes::from_static(b"first"), + Arc::downgrade(&sink), + ) + .await + .unwrap(); + runtime + .send_udp_to_socket( + entry_id, + destination, + Bytes::from_static(b"second"), + Arc::downgrade(&sink), + ) + .await + .unwrap(); + + assert_eq!(factory.options.lock().unwrap().len(), 1); + assert_eq!( + factory.options.lock().unwrap()[0].purpose, + UdpSocketPurpose::ProxyNat + ); + assert_eq!(factory.socket.sent.lock().unwrap().len(), 2); + + runtime.close_udp_socket(entry_id); + runtime + .send_udp_to_socket( + entry_id, + destination, + Bytes::from_static(b"third"), + Arc::downgrade(&sink), + ) + .await + .unwrap(); + assert_eq!(factory.options.lock().unwrap().len(), 2); + } + + #[tokio::test] + async fn close_all_cancels_an_inflight_socket_creation() { + let factory = Arc::new(DelayedFactory::default()); + let runtime = Arc::new(UdpSocketProxyRuntime::new( + factory.clone(), + Arc::new(TestPolicy), + UdpBindOptions::proxy_nat(), + Duration::from_secs(120), + )); + let sink: Arc = Arc::new(NoopResponseSink); + let entry_id = UdpNatEntryId::new(); + let destination = SocketAddr::from((Ipv4Addr::LOCALHOST, 53)); + let send_task = tokio::spawn({ + let runtime = runtime.clone(); + let sink = Arc::downgrade(&sink); + async move { + runtime + .send_udp_to_socket(entry_id, destination, Bytes::from_static(b"request"), sink) + .await + } + }); + + factory.bind_started.notified().await; + runtime.close_all(); + factory.release_bind.notify_one(); + + let error = send_task.await.unwrap().unwrap_err(); + assert!(error.to_string().contains("closed while being created")); + assert_eq!(factory.bind_calls.load(AtomicOrdering::Acquire), 1); + assert!(factory.socket.sent.lock().unwrap().is_empty()); + assert!(runtime.entries.is_empty()); + } + + #[tokio::test] + async fn cancelled_creator_releases_the_slot_for_retry() { + let factory = Arc::new(DelayedFactory::default()); + let runtime = Arc::new(UdpSocketProxyRuntime::new( + factory.clone(), + Arc::new(TestPolicy), + UdpBindOptions::proxy_nat(), + Duration::from_secs(120), + )); + let sink: Arc = Arc::new(NoopResponseSink); + let entry_id = UdpNatEntryId::new(); + let destination = SocketAddr::from((Ipv4Addr::LOCALHOST, 53)); + let first_send = tokio::spawn({ + let runtime = runtime.clone(); + let sink = Arc::downgrade(&sink); + async move { + runtime + .send_udp_to_socket(entry_id, destination, Bytes::from_static(b"first"), sink) + .await + } + }); + + factory.bind_started.notified().await; + first_send.abort(); + assert!(first_send.await.unwrap_err().is_cancelled()); + + let retry = tokio::spawn({ + let runtime = runtime.clone(); + let sink = Arc::downgrade(&sink); + async move { + runtime + .send_udp_to_socket(entry_id, destination, Bytes::from_static(b"retry"), sink) + .await + } + }); + factory.bind_started.notified().await; + factory.release_bind.notify_one(); + + time::timeout(Duration::from_secs(1), retry) + .await + .unwrap() + .unwrap() + .unwrap(); + assert_eq!(factory.bind_calls.load(AtomicOrdering::Acquire), 2); + assert_eq!(factory.socket.sent.lock().unwrap().len(), 1); + } +} diff --git a/easytier-core/src/gateway/proxy/wrapped_tcp_proxy.rs b/easytier-core/src/gateway/proxy/wrapped_tcp_proxy.rs new file mode 100644 index 00000000..e1ce8c5c --- /dev/null +++ b/easytier-core/src/gateway/proxy/wrapped_tcp_proxy.rs @@ -0,0 +1,450 @@ +use std::{ + future::Future, + net::{IpAddr, Ipv4Addr, SocketAddr, SocketAddrV4}, + sync::Arc, +}; + +use easytier_proto::acl::{ChainType, Protocol}; +use smoltcp::wire::{IpProtocol, Ipv4Packet, TcpPacket}; + +use crate::{ + gateway::proxy::{ + cidr_table::ProxyCidrTable, proxy_acl::ProxyAclHandler, + traits::WrappedTcpDestinationRuntime, + }, + packet::{PacketType, ZCPacket}, + peers::{ + acl::{filter::AclFilter, processor::PacketInfo}, + route::Route, + }, +}; + +#[async_trait::async_trait] +pub trait WrappedTcpPeerGroupResolver: Send + Sync { + async fn get_peer_groups_by_ip(&self, ip: &IpAddr) -> Arc>; +} + +#[async_trait::async_trait] +impl WrappedTcpPeerGroupResolver for T +where + T: Route + Send + Sync + ?Sized, +{ + async fn get_peer_groups_by_ip(&self, ip: &IpAddr) -> Arc> { + Route::get_peer_groups_by_ip(self, ip).await + } +} + +#[derive(Debug, Clone, Copy)] +pub struct WrappedTcpDestinationRequest { + pub src: SocketAddr, + pub dst: SocketAddr, + pub initial_packet_size: usize, +} + +#[derive(Clone)] +pub struct WrappedTcpDestinationPlan { + pub socket_dst: SocketAddr, + pub acl_handler: ProxyAclHandler, +} + +pub async fn plan_wrapped_tcp_destination( + request: WrappedTcpDestinationRequest, + cidr_table: &ProxyCidrTable, + runtime: &dyn WrappedTcpDestinationRuntime, + group_resolver: &GroupResolver, + acl_filter: Arc, +) -> anyhow::Result +where + GroupResolver: WrappedTcpPeerGroupResolver + ?Sized, +{ + let mut mapped_dst = request.dst; + if let IpAddr::V4(dst_ip) = mapped_dst.ip() + && let Some(real_ip) = cidr_table.lookup_v4(dst_ip) + { + mapped_dst.set_ip(real_ip.into()); + } + + let src_ip = request.src.ip(); + let dst_ip = mapped_dst.ip(); + let (src_groups, dst_groups) = tokio::join!( + group_resolver.get_peer_groups_by_ip(&src_ip), + group_resolver.get_peer_groups_by_ip(&dst_ip), + ); + + if runtime.should_deny_tcp_proxy(mapped_dst) { + anyhow::bail!( + "dst socket {:?} is in running listeners, ignore it", + mapped_dst + ); + } + + let send_to_self = runtime.is_ip_local_virtual_ip(&dst_ip); + let socket_dst = if send_to_self && runtime.no_tun() { + SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), mapped_dst.port()) + } else { + mapped_dst + }; + + let acl_handler = ProxyAclHandler { + acl_filter, + packet_info: PacketInfo { + src_ip, + dst_ip, + src_port: Some(request.src.port()), + dst_port: Some(socket_dst.port()), + protocol: Protocol::Tcp, + packet_size: request.initial_packet_size, + src_groups, + dst_groups, + }, + chain_type: if send_to_self { + ChainType::Inbound + } else { + ChainType::Forward + }, + }; + acl_handler.handle_packet_size(request.initial_packet_size)?; + + Ok(WrappedTcpDestinationPlan { + socket_dst, + acl_handler, + }) +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum WrappedTcpProxyTransport { + Kcp, + Quic, +} + +#[derive(Debug, Clone, Copy)] +pub struct WrappedTcpProxyNicContext { + pub transport: WrappedTcpProxyTransport, + pub my_peer_id: u32, + pub local_ipv4: Option, + pub smoltcp_enabled: bool, +} + +pub async fn try_process_wrapped_tcp_packet_from_nic( + zc_packet: &mut ZCPacket, + ctx: WrappedTcpProxyNicContext, + is_tcp_proxy_connection: ConnectionLookup, + check_dst_allowed: AllowCheck, +) -> bool +where + ConnectionLookup: Fn(SocketAddr) -> bool, + AllowCheck: FnOnce(Ipv4Addr) -> AllowCheckFut, + AllowCheckFut: Future, +{ + let Some(hdr) = zc_packet.peer_manager_header() else { + return false; + }; + if hdr.packet_type != PacketType::Data as u8 { + return false; + } + + let Ok(ip_packet) = Ipv4Packet::new_checked(zc_packet.payload()) else { + return false; + }; + if ip_packet.version() != 4 || ip_packet.next_header() != IpProtocol::Tcp { + return false; + } + + let Ok(tcp_packet) = TcpPacket::new_checked(ip_packet.payload()) else { + return false; + }; + let src_ip = ip_packet.src_addr(); + let dst_ip = ip_packet.dst_addr(); + let src_port = tcp_packet.src_port(); + let is_syn = tcp_packet.syn() && !tcp_packet.ack(); + + if is_syn { + if !check_dst_allowed(dst_ip).await { + tracing::warn!( + ?ctx.transport, + dst = %dst_ip, + "wrapped tcp proxy src dst is not allowed" + ); + return false; + } + } else if !is_tcp_proxy_connection(SocketAddr::V4(SocketAddrV4::new(src_ip, src_port))) { + return false; + } + + if let Some(local_ipv4) = ctx.local_ipv4 + && src_ip != local_ipv4 + && !ctx.smoltcp_enabled + { + tracing::warn!( + ?ctx.transport, + src = %src_ip, + dst = %dst_ip, + "wrapped tcp proxy net-to-net input is not allowed without smoltcp" + ); + return false; + } + + let hdr = zc_packet + .mut_peer_manager_header() + .expect("peer manager header"); + hdr.to_peer_id = ctx.my_peer_id.into(); + match ctx.transport { + WrappedTcpProxyTransport::Kcp => { + hdr.mark_kcp_src_modified(); + } + WrappedTcpProxyTransport::Quic => { + hdr.mark_quic_src_modified(); + } + } + true +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::gateway::proxy::cidr_table::{ProxyCidrRule, ProxyCidrSnapshot}; + use smoltcp::wire::{IpAddress, IpProtocol, Ipv4Packet, TcpPacket}; + + struct TestDestinationRuntime { + local_ip: IpAddr, + no_tun: bool, + denied: Option, + } + + impl WrappedTcpDestinationRuntime for TestDestinationRuntime { + fn is_ip_local_virtual_ip(&self, ip: &IpAddr) -> bool { + *ip == self.local_ip + } + + fn no_tun(&self) -> bool { + self.no_tun + } + + fn should_deny_tcp_proxy(&self, dst: SocketAddr) -> bool { + self.denied == Some(dst) + } + } + + struct TestGroupResolver; + + #[async_trait::async_trait] + impl WrappedTcpPeerGroupResolver for TestGroupResolver { + async fn get_peer_groups_by_ip(&self, ip: &IpAddr) -> Arc> { + Arc::new(vec![format!("group-{ip}")]) + } + } + + fn destination_runtime(local_ip: &str, no_tun: bool) -> TestDestinationRuntime { + TestDestinationRuntime { + local_ip: local_ip.parse().unwrap(), + no_tun, + denied: None, + } + } + + fn mapped_cidr_table() -> ProxyCidrTable { + ProxyCidrTable::from_snapshot(ProxyCidrSnapshot { + rules: vec![ProxyCidrRule { + cidr: "10.10.0.0/16".parse().unwrap(), + mapped_cidr: Some("100.64.0.0/16".parse().unwrap()), + }], + }) + } + + async fn destination_plan( + runtime: &TestDestinationRuntime, + dst: &str, + ) -> anyhow::Result { + plan_wrapped_tcp_destination( + WrappedTcpDestinationRequest { + src: "10.20.0.2:40000".parse().unwrap(), + dst: dst.parse().unwrap(), + initial_packet_size: b"header".len(), + }, + &mapped_cidr_table(), + runtime, + &TestGroupResolver, + Arc::new(AclFilter::new()), + ) + .await + } + + #[tokio::test] + async fn destination_plan_maps_cidr_and_builds_forward_acl_context() { + let plan = destination_plan(&destination_runtime("10.30.0.1", false), "100.64.2.3:443") + .await + .unwrap(); + + assert_eq!(plan.socket_dst, "10.10.2.3:443".parse().unwrap()); + assert_eq!(plan.acl_handler.chain_type, ChainType::Forward); + assert_eq!( + plan.acl_handler.packet_info.dst_ip, + "10.10.2.3".parse::().unwrap() + ); + assert_eq!( + plan.acl_handler.packet_info.dst_groups.as_ref(), + &["group-10.10.2.3"] + ); + } + + #[tokio::test] + async fn destination_plan_rewrites_local_no_tun_socket_after_acl_identity() { + let plan = destination_plan(&destination_runtime("10.10.2.3", true), "100.64.2.3:443") + .await + .unwrap(); + + assert_eq!(plan.socket_dst, "127.0.0.1:443".parse().unwrap()); + assert_eq!(plan.acl_handler.chain_type, ChainType::Inbound); + assert_eq!( + plan.acl_handler.packet_info.dst_ip, + "10.10.2.3".parse::().unwrap() + ); + } + + #[tokio::test] + async fn destination_plan_denies_mapped_running_listener() { + let mut runtime = destination_runtime("10.30.0.1", false); + runtime.denied = Some("10.10.2.3:443".parse().unwrap()); + + let err = destination_plan(&runtime, "100.64.2.3:443") + .await + .err() + .expect("mapped listener must be denied"); + assert!(err.to_string().contains("running listeners")); + } + + fn build_tcp_packet(src: SocketAddrV4, dst: SocketAddrV4, syn: bool, ack: bool) -> ZCPacket { + let mut raw = vec![0; smoltcp::wire::IPV4_HEADER_LEN + smoltcp::wire::TCP_HEADER_LEN]; + { + let mut ipv4 = Ipv4Packet::new_unchecked(&mut raw); + ipv4.set_version(4); + ipv4.set_header_len(smoltcp::wire::IPV4_HEADER_LEN as u8); + ipv4.set_total_len( + (smoltcp::wire::IPV4_HEADER_LEN + smoltcp::wire::TCP_HEADER_LEN) as u16, + ); + ipv4.set_hop_limit(64); + ipv4.set_next_header(IpProtocol::Tcp); + ipv4.set_src_addr(*src.ip()); + ipv4.set_dst_addr(*dst.ip()); + ipv4.fill_checksum(); + } + { + let mut tcp = TcpPacket::new_unchecked(&mut raw[smoltcp::wire::IPV4_HEADER_LEN..]); + tcp.set_src_port(src.port()); + tcp.set_dst_port(dst.port()); + tcp.set_header_len(smoltcp::wire::TCP_HEADER_LEN as u8); + tcp.set_syn(syn); + tcp.set_ack(ack); + tcp.fill_checksum(&IpAddress::Ipv4(*src.ip()), &IpAddress::Ipv4(*dst.ip())); + } + + let mut packet = ZCPacket::new_with_payload(&raw); + packet.fill_peer_manager_hdr(1, 2, PacketType::Data as u8); + packet + } + + fn context(transport: WrappedTcpProxyTransport) -> WrappedTcpProxyNicContext { + WrappedTcpProxyNicContext { + transport, + my_peer_id: 42, + local_ipv4: Some("10.144.144.204".parse().unwrap()), + smoltcp_enabled: false, + } + } + + #[tokio::test] + async fn allowed_syn_is_marked_for_kcp() { + let src = SocketAddrV4::new("10.144.144.204".parse().unwrap(), 50000); + let dst = SocketAddrV4::new("10.10.10.10".parse().unwrap(), 80); + let mut packet = build_tcp_packet(src, dst, true, false); + + assert!( + try_process_wrapped_tcp_packet_from_nic( + &mut packet, + context(WrappedTcpProxyTransport::Kcp), + |_| false, + |_| async { true }, + ) + .await + ); + + let hdr = packet.peer_manager_header().unwrap(); + assert_eq!(hdr.to_peer_id.get(), 42); + assert!(hdr.is_kcp_src_modified()); + } + + #[tokio::test] + async fn denied_syn_is_not_marked() { + let src = SocketAddrV4::new("10.144.144.204".parse().unwrap(), 50000); + let dst = SocketAddrV4::new("10.10.10.10".parse().unwrap(), 80); + let mut packet = build_tcp_packet(src, dst, true, false); + + assert!( + !try_process_wrapped_tcp_packet_from_nic( + &mut packet, + context(WrappedTcpProxyTransport::Kcp), + |_| false, + |_| async { false }, + ) + .await + ); + + let hdr = packet.peer_manager_header().unwrap(); + assert_eq!(hdr.packet_type, PacketType::Data as u8); + } + + #[tokio::test] + async fn established_non_syn_is_marked_for_quic() { + let src = SocketAddrV4::new("10.144.144.204".parse().unwrap(), 50000); + let dst = SocketAddrV4::new("10.10.10.10".parse().unwrap(), 80); + let mut packet = build_tcp_packet(src, dst, false, true); + + assert!( + try_process_wrapped_tcp_packet_from_nic( + &mut packet, + context(WrappedTcpProxyTransport::Quic), + |addr| addr == SocketAddr::V4(src), + |_| async { false }, + ) + .await + ); + + let hdr = packet.peer_manager_header().unwrap(); + assert_eq!(hdr.to_peer_id.get(), 42); + assert!(hdr.is_quic_src_modified()); + } + + #[tokio::test] + async fn non_syn_without_connection_is_rejected() { + let src = SocketAddrV4::new("10.144.144.204".parse().unwrap(), 50000); + let dst = SocketAddrV4::new("10.10.10.10".parse().unwrap(), 80); + let mut packet = build_tcp_packet(src, dst, false, true); + + assert!( + !try_process_wrapped_tcp_packet_from_nic( + &mut packet, + context(WrappedTcpProxyTransport::Quic), + |_| false, + |_| async { true }, + ) + .await + ); + } + + #[tokio::test] + async fn net_to_net_without_smoltcp_is_rejected() { + let src = SocketAddrV4::new("10.144.144.205".parse().unwrap(), 50000); + let dst = SocketAddrV4::new("10.10.10.10".parse().unwrap(), 80); + let mut packet = build_tcp_packet(src, dst, true, false); + + assert!( + !try_process_wrapped_tcp_packet_from_nic( + &mut packet, + context(WrappedTcpProxyTransport::Kcp), + |_| false, + |_| async { true }, + ) + .await + ); + } +} diff --git a/easytier-core/src/gateway/proxy/wrapped_transport.rs b/easytier-core/src/gateway/proxy/wrapped_transport.rs new file mode 100644 index 00000000..cc1ab698 --- /dev/null +++ b/easytier-core/src/gateway/proxy/wrapped_transport.rs @@ -0,0 +1,1142 @@ +use std::sync::{Arc, Weak}; + +use async_trait::async_trait; +use bytes::BytesMut; +use tokio::{sync::Mutex, task::JoinSet}; + +use crate::{ + config::runtime::CoreRuntimeConfigStore, + connectivity::direct::DirectConnectorHost, + connectivity::hole_punch::tcp::TcpHolePunchHost, + gateway::proxy::cidr_table::ProxyCidrTable, + listener::RunningListenerRegistry, + packet::{PacketType, ZCPacket, ZCPacketType}, + peers::{ + PeerPacketFilter, + peer_manager::{PeerManagerCore, PipelineRegistrationGuard}, + }, + process_runtime::ProtectedTcpPortRegistry, + socket::SocketContext, +}; + +#[cfg(all(feature = "proxy-packet", any(feature = "proxy-smoltcp-stack", test)))] +mod connect_api; +mod engine_api; +#[cfg(feature = "proxy-packet")] +mod packet_api; +#[cfg(feature = "proxy-packet")] +#[path = "wrapped_transport/packet_plane.rs"] +mod packet_plane; +#[cfg(not(feature = "proxy-packet"))] +#[path = "wrapped_transport/packet_plane_disabled.rs"] +mod packet_plane; + +#[cfg(feature = "proxy-packet")] +pub use engine_api::{WrappedTransportAcceptedStream, WrappedTransportConnect}; +pub use engine_api::{WrappedTransportEngine, WrappedTransportEngineStart}; +#[cfg(feature = "proxy-packet")] +pub use packet_plane::WrappedTransportDestinationIngress; +use packet_plane::{WrappedTransportPacketPlane, WrappedTransportPacketState}; + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub struct WrappedTransportDirections { + pub source: bool, + pub destination: bool, +} + +impl WrappedTransportDirections { + fn enabled(self) -> bool { + self.source || self.destination + } +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum WrappedTransportKind { + Kcp, + Quic, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum WrappedTransportRole { + Source, + Destination, +} + +#[derive(Debug, Clone)] +pub struct WrappedTransportDatagram { + pub transport: WrappedTransportKind, + pub role: WrappedTransportRole, + pub peer_id: u32, + pub buffer: WrappedTransportDatagramBuffer, +} + +#[derive(Debug, Clone)] +pub struct WrappedTransportDatagramBuffer(ZCPacket); + +impl WrappedTransportDatagramBuffer { + pub fn copy_from_payload(payload: &[u8]) -> Self { + Self(ZCPacket::new_with_payload(payload)) + } + + pub fn from_packet_buffer(buffer: BytesMut, packet_type: ZCPacketType) -> Self { + Self(ZCPacket::new_from_buf(buffer, packet_type)) + } +} + +impl WrappedTransportDatagram { + fn into_packet(self, my_peer_id: u32) -> ZCPacket { + let mut packet = self.buffer.0; + packet.fill_peer_manager_hdr( + my_peer_id, + self.peer_id, + self.transport.packet_type(self.role) as u8, + ); + packet + } +} + +#[derive(Default)] +pub struct WrappedTransportEngines { + pub kcp: Option>, + pub quic: Option>, +} + +#[derive(Default)] +struct WrappedTransportProxyState { + active: bool, + kcp_started: bool, + quic_started: bool, + packet: WrappedTransportPacketState, + pipeline_guards: Vec, + tasks: JoinSet<()>, +} + +impl WrappedTransportProxyState { + fn has_partial_start(&self, packet_plane: &WrappedTransportPacketPlane) -> bool { + self.kcp_started + || self.quic_started + || !self.pipeline_guards.is_empty() + || !self.tasks.is_empty() + || packet_plane.has_partial_start(&self.packet) + } +} + +impl WrappedTransportKind { + fn packet_type(self, role: WrappedTransportRole) -> PacketType { + match (self, role) { + (Self::Kcp, WrappedTransportRole::Source) => PacketType::KcpSrc, + (Self::Kcp, WrappedTransportRole::Destination) => PacketType::KcpDst, + (Self::Quic, WrappedTransportRole::Source) => PacketType::QuicSrc, + (Self::Quic, WrappedTransportRole::Destination) => PacketType::QuicDst, + } + } + + fn incoming_packet_type(self, role: WrappedTransportRole) -> PacketType { + match role { + WrappedTransportRole::Source => self.packet_type(WrappedTransportRole::Destination), + WrappedTransportRole::Destination => self.packet_type(WrappedTransportRole::Source), + } + } +} + +struct WrappedTransportPeerFilter { + engine: Weak, + transport: WrappedTransportKind, + role: WrappedTransportRole, +} + +#[async_trait] +impl PeerPacketFilter for WrappedTransportPeerFilter { + async fn try_process_packet_from_peer(&self, packet: ZCPacket) -> Option { + let Some(header) = packet.peer_manager_header() else { + return Some(packet); + }; + if header.packet_type != self.transport.incoming_packet_type(self.role) as u8 { + return Some(packet); + } + let Some(engine) = self.engine.upgrade() else { + return Some(packet); + }; + let from_peer_id = header.from_peer_id.get(); + if let Err(error) = engine + .inject_peer_datagram(self.role, from_peer_id, packet.payload_bytes().freeze()) + .await + { + tracing::debug!( + ?error, + transport = ?self.transport, + role = ?self.role, + "failed to inject wrapped transport packet" + ); + } + None + } +} + +pub(crate) struct WrappedTransportProxyModule { + peer_manager: Arc, + runtime_config: CoreRuntimeConfigStore, + kcp: Option>, + quic: Option>, + packet_plane: WrappedTransportPacketPlane, + state: Mutex, +} + +impl WrappedTransportProxyModule { + const DATAGRAM_QUEUE_CAPACITY: usize = 1024; + + #[allow(clippy::too_many_arguments)] + pub(crate) fn new( + peer_manager: Arc, + runtime_config: CoreRuntimeConfigStore, + kcp: Option>, + quic: Option>, + host: Arc, + protected_tcp_ports: Arc, + running_listeners: Arc, + cidr_table: Arc, + socket_context: SocketContext, + ) -> Option> + where + H: DirectConnectorHost + TcpHolePunchHost, + { + if kcp.is_none() && quic.is_none() { + return None; + } + let packet_plane = WrappedTransportPacketPlane::new( + peer_manager.clone(), + runtime_config.clone(), + &kcp, + &quic, + host, + protected_tcp_ports, + running_listeners, + cidr_table, + socket_context, + ); + Some(Arc::new(Self { + peer_manager, + runtime_config, + kcp, + quic, + packet_plane, + state: Mutex::new(WrappedTransportProxyState::default()), + })) + } + + fn directions(&self) -> (WrappedTransportDirections, WrappedTransportDirections) { + let snapshot = self.runtime_config.snapshot(); + let flags = &snapshot.peer.flags; + ( + WrappedTransportDirections { + source: flags.enable_kcp_proxy, + destination: !flags.disable_kcp_input, + }, + WrappedTransportDirections { + source: flags.enable_quic_proxy, + destination: !flags.disable_quic_input, + }, + ) + } + + fn spawn_datagram_egress( + &self, + state: &mut WrappedTransportProxyState, + ) -> tokio::sync::mpsc::Sender { + let (tx, mut rx) = + tokio::sync::mpsc::channel::(Self::DATAGRAM_QUEUE_CAPACITY); + let peer_manager = self.peer_manager.clone(); + state.tasks.spawn(async move { + while let Some(datagram) = rx.recv().await { + let peer_id = datagram.peer_id; + let transport = datagram.transport; + let role = datagram.role; + let packet = datagram.into_packet(peer_manager.my_peer_id()); + if let Err(error) = peer_manager.send_msg_for_proxy(packet, peer_id).await { + tracing::error!( + ?error, + ?transport, + ?role, + peer_id, + "failed to send wrapped transport packet" + ); + } + } + }); + tx + } + + async fn register_peer_filters( + &self, + state: &mut WrappedTransportProxyState, + engine: &Arc, + transport: WrappedTransportKind, + directions: WrappedTransportDirections, + ) { + if directions.source { + let guard = self + .peer_manager + .add_managed_packet_process_pipeline(Box::new(WrappedTransportPeerFilter { + engine: Arc::downgrade(engine), + transport, + role: WrappedTransportRole::Source, + })) + .await; + state.pipeline_guards.push(guard); + } + if directions.destination { + let guard = self + .peer_manager + .add_managed_packet_process_pipeline(Box::new(WrappedTransportPeerFilter { + engine: Arc::downgrade(engine), + transport, + role: WrappedTransportRole::Destination, + })) + .await; + state.pipeline_guards.push(guard); + } + } + + async fn stop_started(&self, state: &mut WrappedTransportProxyState) { + self.packet_plane.clear_connect_ready(&mut state.packet); + for guard in state.pipeline_guards.drain(..).rev() { + guard.close(); + } + if state.quic_started { + self.packet_plane + .stop_source(&mut state.packet, WrappedTransportKind::Quic) + .await; + if let Some(quic) = &self.quic { + quic.stop().await; + } + state.quic_started = false; + } + if state.kcp_started { + self.packet_plane + .stop_source(&mut state.packet, WrappedTransportKind::Kcp) + .await; + if let Some(kcp) = &self.kcp { + kcp.stop().await; + } + state.kcp_started = false; + } + self.packet_plane.stop_destination(&mut state.packet).await; + state.tasks.shutdown().await; + state.active = false; + } + + pub(crate) async fn start(&self) -> anyhow::Result<()> { + let mut state = self.state.lock().await; + if state.active { + return Ok(()); + } + if state.has_partial_start(&self.packet_plane) { + self.stop_started(&mut state).await; + } + let (kcp_directions, quic_directions) = self.directions(); + if let Err(error) = self + .packet_plane + .start_destinations( + &mut state.packet, + kcp_directions, + quic_directions, + self.kcp.is_some(), + self.quic.is_some(), + ) + .await + { + self.stop_started(&mut state).await; + return Err(error); + } + + if let Some(kcp) = &self.kcp + && kcp_directions.enabled() + { + let datagrams = self.spawn_datagram_egress(&mut state); + state.kcp_started = true; + let options = self.packet_plane.engine_start( + &state.packet, + WrappedTransportKind::Kcp, + kcp_directions, + self.peer_manager.my_peer_id(), + datagrams, + ); + if let Err(error) = kcp.prepare(options).await { + self.stop_started(&mut state).await; + return Err(error); + } + self.register_peer_filters(&mut state, kcp, WrappedTransportKind::Kcp, kcp_directions) + .await; + if let Err(error) = self + .packet_plane + .start_source(&mut state.packet, WrappedTransportKind::Kcp, kcp_directions) + .await + { + self.stop_started(&mut state).await; + return Err(error); + } + if let Err(error) = kcp.activate().await { + self.stop_started(&mut state).await; + return Err(error); + } + self.packet_plane.mark_source_connect_ready( + &mut state.packet, + WrappedTransportKind::Kcp, + kcp_directions, + ); + } + if let Some(quic) = &self.quic + && quic_directions.enabled() + { + let datagrams = self.spawn_datagram_egress(&mut state); + state.quic_started = true; + let options = self.packet_plane.engine_start( + &state.packet, + WrappedTransportKind::Quic, + quic_directions, + self.peer_manager.my_peer_id(), + datagrams, + ); + if let Err(error) = quic.prepare(options).await { + self.stop_started(&mut state).await; + return Err(error); + } + self.register_peer_filters( + &mut state, + quic, + WrappedTransportKind::Quic, + quic_directions, + ) + .await; + if let Err(error) = self + .packet_plane + .start_source( + &mut state.packet, + WrappedTransportKind::Quic, + quic_directions, + ) + .await + { + self.stop_started(&mut state).await; + return Err(error); + } + if let Err(error) = quic.activate().await { + self.stop_started(&mut state).await; + return Err(error); + } + self.packet_plane.mark_source_connect_ready( + &mut state.packet, + WrappedTransportKind::Quic, + quic_directions, + ); + } + state.active = true; + Ok(()) + } + + pub(crate) async fn stop(&self) { + let mut state = self.state.lock().await; + self.stop_started(&mut state).await; + } +} + +#[cfg(test)] +mod tests { + use std::sync::{ + Arc, Mutex, + atomic::{AtomicUsize, Ordering}, + }; + + use bytes::Bytes; + use tokio::sync::Notify; + + #[cfg(feature = "proxy-packet")] + use crate::gateway::proxy::{ + tcp_proxy_engine::TcpNatEntrySnapshot, + traits::TcpProxyStream, + wrapped_transport_destination::{ + WrappedTransportDestinationIngresses, WrappedTransportDestinationLifecycle, + }, + }; + use crate::{ + config::peers::{HostRoutingPolicy, PeerRuntimeConfig, PeerRuntimeSnapshot}, + config::runtime::CoreRuntimeConfig, + config::{CoreConfig, NetworkIdentity, NodeConfig}, + peers::{create_packet_recv_chan, peer_manager::PortablePeerManagerConfig}, + }; + + use super::*; + + #[cfg(feature = "proxy-packet")] + #[derive(Default)] + struct TestWrappedTransportDestination { + active: std::sync::atomic::AtomicBool, + enabled: std::sync::atomic::AtomicU8, + } + + #[cfg(feature = "proxy-packet")] + #[async_trait] + impl WrappedTransportDestinationLifecycle for TestWrappedTransportDestination { + async fn start( + self: Arc, + kcp: bool, + quic: bool, + ) -> anyhow::Result { + self.active + .store(true, std::sync::atomic::Ordering::Release); + self.enabled.store( + (kcp as u8) | ((quic as u8) << 1), + std::sync::atomic::Ordering::Release, + ); + Ok(WrappedTransportDestinationIngresses::default()) + } + + async fn stop(&self) { + self.active + .store(false, std::sync::atomic::Ordering::Release); + self.enabled.store(0, std::sync::atomic::Ordering::Release); + } + + fn entry_snapshots(&self, _transport: WrappedTransportKind) -> Vec { + Vec::new() + } + + fn is_started(&self, transport: WrappedTransportKind) -> bool { + let mask = match transport { + WrappedTransportKind::Kcp => 1, + WrappedTransportKind::Quic => 2, + }; + self.active.load(std::sync::atomic::Ordering::Acquire) + && self.enabled.load(std::sync::atomic::Ordering::Acquire) & mask != 0 + } + } + + impl WrappedTransportProxyModule { + fn new_without_sources( + peer_manager: Arc, + runtime_config: CoreRuntimeConfigStore, + kcp: Option>, + quic: Option>, + ) -> Option> { + if kcp.is_none() && quic.is_none() { + return None; + } + + #[cfg(feature = "proxy-packet")] + let packet_plane = WrappedTransportPacketPlane { + kcp_source: None, + quic_source: None, + destination: Some(Arc::new(TestWrappedTransportDestination::default())), + }; + #[cfg(not(feature = "proxy-packet"))] + let packet_plane = WrappedTransportPacketPlane::default(); + + Some(Arc::new(Self { + peer_manager, + runtime_config, + kcp, + quic, + packet_plane, + state: tokio::sync::Mutex::new(WrappedTransportProxyState::default()), + })) + } + } + + struct RecordingWrappedTransportEngine { + name: &'static str, + fail_start: bool, + events: Arc>>, + } + + struct CancelOnceStopEngine { + stop_calls: AtomicUsize, + stop_entered: Notify, + release_first_stop: Notify, + events: Arc>>, + } + + #[derive(Default)] + struct AbortOncePrepareEngine { + prepare_calls: AtomicUsize, + activate_calls: AtomicUsize, + stop_calls: AtomicUsize, + first_prepare_entered: Notify, + } + + #[derive(Default)] + struct RecordingDatagramEngine { + injections: Mutex>, + } + + #[async_trait] + impl WrappedTransportEngine for RecordingDatagramEngine { + async fn prepare(&self, _options: WrappedTransportEngineStart) -> anyhow::Result<()> { + Ok(()) + } + + async fn activate(&self) -> anyhow::Result<()> { + Ok(()) + } + + async fn inject_peer_datagram( + &self, + role: WrappedTransportRole, + from_peer_id: u32, + payload: Bytes, + ) -> anyhow::Result<()> { + self.injections + .lock() + .unwrap() + .push((role, from_peer_id, payload)); + Ok(()) + } + + #[cfg(feature = "proxy-packet")] + async fn connect_source( + &self, + _request: WrappedTransportConnect, + ) -> anyhow::Result> { + anyhow::bail!("recording engine does not open streams") + } + + async fn stop(&self) {} + } + + #[async_trait] + impl WrappedTransportEngine for CancelOnceStopEngine { + async fn prepare(&self, options: WrappedTransportEngineStart) -> anyhow::Result<()> { + self.events.lock().unwrap().push("prepare:blocking".into()); + assert_eq!( + options.directions, + WrappedTransportDirections { + source: true, + destination: true, + } + ); + Ok(()) + } + + async fn activate(&self) -> anyhow::Result<()> { + self.events.lock().unwrap().push("activate:blocking".into()); + Ok(()) + } + + async fn inject_peer_datagram( + &self, + _role: WrappedTransportRole, + _from_peer_id: u32, + _payload: Bytes, + ) -> anyhow::Result<()> { + Ok(()) + } + + #[cfg(feature = "proxy-packet")] + async fn connect_source( + &self, + _request: WrappedTransportConnect, + ) -> anyhow::Result> { + anyhow::bail!("blocking engine does not open streams") + } + + async fn stop(&self) { + let call = self.stop_calls.fetch_add(1, Ordering::AcqRel); + self.events.lock().unwrap().push("stop:blocking".into()); + if call == 0 { + self.stop_entered.notify_one(); + self.release_first_stop.notified().await; + } + } + } + + #[async_trait] + impl WrappedTransportEngine for AbortOncePrepareEngine { + async fn prepare(&self, _options: WrappedTransportEngineStart) -> anyhow::Result<()> { + if self.prepare_calls.fetch_add(1, Ordering::AcqRel) == 0 { + self.first_prepare_entered.notify_one(); + std::future::pending::<()>().await; + } + Ok(()) + } + + async fn activate(&self) -> anyhow::Result<()> { + self.activate_calls.fetch_add(1, Ordering::AcqRel); + Ok(()) + } + + async fn inject_peer_datagram( + &self, + _role: WrappedTransportRole, + _from_peer_id: u32, + _payload: Bytes, + ) -> anyhow::Result<()> { + Ok(()) + } + + #[cfg(feature = "proxy-packet")] + async fn connect_source( + &self, + _request: WrappedTransportConnect, + ) -> anyhow::Result> { + anyhow::bail!("blocking engine does not open streams") + } + + async fn stop(&self) { + self.stop_calls.fetch_add(1, Ordering::AcqRel); + } + } + + #[async_trait] + impl WrappedTransportEngine for RecordingWrappedTransportEngine { + async fn prepare(&self, options: WrappedTransportEngineStart) -> anyhow::Result<()> { + let directions = options.directions; + self.events.lock().unwrap().push(format!( + "prepare:{}:{}:{}", + self.name, directions.source, directions.destination + )); + if self.fail_start { + anyhow::bail!("{} start failed", self.name); + } + Ok(()) + } + + async fn activate(&self) -> anyhow::Result<()> { + self.events + .lock() + .unwrap() + .push(format!("activate:{}", self.name)); + Ok(()) + } + + async fn inject_peer_datagram( + &self, + _role: WrappedTransportRole, + _from_peer_id: u32, + _payload: Bytes, + ) -> anyhow::Result<()> { + Ok(()) + } + + #[cfg(feature = "proxy-packet")] + async fn connect_source( + &self, + _request: WrappedTransportConnect, + ) -> anyhow::Result> { + anyhow::bail!("recording engine does not open streams") + } + + async fn stop(&self) { + self.events + .lock() + .unwrap() + .push(format!("stop:{}", self.name)); + } + } + + fn wrapped_transport_engine( + name: &'static str, + fail_start: bool, + events: &Arc>>, + ) -> Arc { + Arc::new(RecordingWrappedTransportEngine { + name, + fail_start, + events: events.clone(), + }) + } + + fn wrapped_transport_peer_manager() -> Arc { + let (packet_tx, _packet_rx) = create_packet_recv_chan(); + Arc::new( + PeerManagerCore::new_portable_for_test( + PortablePeerManagerConfig::new(PeerRuntimeConfig { + core: CoreConfig { + node: NodeConfig { + peer_id: Some(1), + network_name: "wrapped-transport-test".to_owned(), + ..Default::default() + }, + ..Default::default() + }, + network_identity: NetworkIdentity { + network_name: "wrapped-transport-test".to_owned(), + network_secret: Some("secret".to_owned()), + network_secret_digest: None, + }, + stun_info: Default::default(), + feature_flags: Default::default(), + secure_mode: None, + host_routing: HostRoutingPolicy::default(), + }), + packet_tx, + ) + .unwrap(), + ) + } + + fn wrapped_transport_runtime( + kcp: WrappedTransportDirections, + quic: WrappedTransportDirections, + ) -> CoreRuntimeConfigStore { + let mut peer = PeerRuntimeSnapshot::default(); + peer.flags.enable_kcp_proxy = kcp.source; + peer.flags.disable_kcp_input = !kcp.destination; + peer.flags.enable_quic_proxy = quic.source; + peer.flags.disable_quic_input = !quic.destination; + CoreRuntimeConfigStore::new(CoreRuntimeConfig::default(), Arc::new(peer)) + } + + #[tokio::test] + async fn peer_filter_maps_packet_role_and_releases_stale_engine() { + let engine = Arc::new(RecordingDatagramEngine::default()); + let dyn_engine: Arc = engine.clone(); + let filter = WrappedTransportPeerFilter { + engine: Arc::downgrade(&dyn_engine), + transport: WrappedTransportKind::Kcp, + role: WrappedTransportRole::Source, + }; + let mut packet = ZCPacket::new_with_payload(b"payload"); + packet.fill_peer_manager_hdr(7, 1, PacketType::KcpDst as u8); + + assert!(filter.try_process_packet_from_peer(packet).await.is_none()); + assert_eq!( + *engine.injections.lock().unwrap(), + [( + WrappedTransportRole::Source, + 7, + Bytes::from_static(b"payload"), + )] + ); + + drop(dyn_engine); + drop(engine); + let mut packet = ZCPacket::new_with_payload(b"stale"); + packet.fill_peer_manager_hdr(7, 1, PacketType::KcpDst as u8); + assert!(filter.try_process_packet_from_peer(packet).await.is_some()); + } + + #[test] + fn datagram_preserves_wire_packet_types() { + for (transport, role, packet_type) in [ + ( + WrappedTransportKind::Kcp, + WrappedTransportRole::Source, + PacketType::KcpSrc, + ), + ( + WrappedTransportKind::Kcp, + WrappedTransportRole::Destination, + PacketType::KcpDst, + ), + ( + WrappedTransportKind::Quic, + WrappedTransportRole::Source, + PacketType::QuicSrc, + ), + ( + WrappedTransportKind::Quic, + WrappedTransportRole::Destination, + PacketType::QuicDst, + ), + ] { + let packet = WrappedTransportDatagram { + transport, + role, + peer_id: 9, + buffer: WrappedTransportDatagramBuffer::copy_from_payload(b"wire"), + } + .into_packet(7); + let header = packet.peer_manager_header().unwrap(); + + assert_eq!(header.from_peer_id.get(), 7); + assert_eq!(header.to_peer_id.get(), 9); + assert_eq!(header.packet_type, packet_type as u8); + assert_eq!(packet.payload(), b"wire"); + } + } + + #[test] + fn packet_buffer_preserves_headroom_allocation() { + let packet = ZCPacket::new_with_payload(b"wire"); + let packet_type = packet.packet_type(); + let buffer = packet.inner(); + let allocation = buffer.as_ptr(); + + let mut packet = WrappedTransportDatagram { + transport: WrappedTransportKind::Quic, + role: WrappedTransportRole::Source, + peer_id: 9, + buffer: WrappedTransportDatagramBuffer::from_packet_buffer(buffer, packet_type), + } + .into_packet(7); + + assert_eq!(packet.mut_inner().as_ptr(), allocation); + assert_eq!(packet.payload(), b"wire"); + } + + #[tokio::test] + async fn module_reads_activation_flags_and_stops_in_reverse() { + let events = Arc::new(Mutex::new(Vec::new())); + let runtime = wrapped_transport_runtime( + WrappedTransportDirections { + source: false, + destination: false, + }, + WrappedTransportDirections { + source: false, + destination: false, + }, + ); + let module = WrappedTransportProxyModule::new_without_sources( + wrapped_transport_peer_manager(), + runtime.clone(), + Some(wrapped_transport_engine("kcp", false, &events)), + Some(wrapped_transport_engine("quic", false, &events)), + ) + .unwrap(); + + runtime.update_peer(Arc::new({ + let mut peer = PeerRuntimeSnapshot::default(); + peer.flags.enable_kcp_proxy = true; + peer.flags.disable_kcp_input = true; + peer.flags.enable_quic_proxy = false; + peer.flags.disable_quic_input = false; + peer + })); + + module.start().await.unwrap(); + module.start().await.unwrap(); + module.stop().await; + module.stop().await; + + assert_eq!( + *events.lock().unwrap(), + [ + "prepare:kcp:true:false", + "activate:kcp", + "prepare:quic:false:true", + "activate:quic", + "stop:quic", + "stop:kcp", + ] + ); + } + + #[cfg(feature = "proxy-packet")] + #[tokio::test] + async fn destination_state_only_reports_available_engines() { + let events = Arc::new(Mutex::new(Vec::new())); + let runtime = wrapped_transport_runtime( + WrappedTransportDirections { + source: false, + destination: true, + }, + WrappedTransportDirections { + source: false, + destination: true, + }, + ); + let module = WrappedTransportProxyModule::new_without_sources( + wrapped_transport_peer_manager(), + runtime, + Some(wrapped_transport_engine("kcp", false, &events)), + None, + ) + .unwrap(); + + module.start().await.unwrap(); + + assert!(module.destination_is_started(WrappedTransportKind::Kcp)); + assert!(!module.destination_is_started(WrappedTransportKind::Quic)); + module.stop().await; + } + + #[cfg(feature = "proxy-packet")] + #[tokio::test] + async fn source_connect_readiness_tracks_active_direction() { + let events = Arc::new(Mutex::new(Vec::new())); + let destination_only = WrappedTransportProxyModule::new_without_sources( + wrapped_transport_peer_manager(), + wrapped_transport_runtime( + WrappedTransportDirections { + source: false, + destination: true, + }, + WrappedTransportDirections { + source: false, + destination: false, + }, + ), + Some(wrapped_transport_engine("kcp", false, &events)), + None, + ) + .unwrap(); + destination_only.start().await.unwrap(); + assert!( + !destination_only + .source_connect_ready(WrappedTransportKind::Kcp) + .await + ); + destination_only.stop().await; + + let source = WrappedTransportProxyModule::new_without_sources( + wrapped_transport_peer_manager(), + wrapped_transport_runtime( + WrappedTransportDirections { + source: true, + destination: false, + }, + WrappedTransportDirections { + source: false, + destination: false, + }, + ), + Some(wrapped_transport_engine("kcp", false, &events)), + None, + ) + .unwrap(); + source.start().await.unwrap(); + assert!(source.source_connect_ready(WrappedTransportKind::Kcp).await); + source.stop().await; + assert!(!source.source_connect_ready(WrappedTransportKind::Kcp).await); + } + + #[tokio::test] + async fn module_rolls_back_failing_engine_and_predecessor() { + let events = Arc::new(Mutex::new(Vec::new())); + let runtime = wrapped_transport_runtime( + WrappedTransportDirections { + source: true, + destination: true, + }, + WrappedTransportDirections { + source: true, + destination: true, + }, + ); + let module = WrappedTransportProxyModule::new_without_sources( + wrapped_transport_peer_manager(), + runtime, + Some(wrapped_transport_engine("kcp", false, &events)), + Some(wrapped_transport_engine("quic", true, &events)), + ) + .unwrap(); + + assert!(module.start().await.is_err()); + + assert_eq!( + *events.lock().unwrap(), + [ + "prepare:kcp:true:true", + "activate:kcp", + "prepare:quic:true:true", + "stop:quic", + "stop:kcp", + ] + ); + } + + #[tokio::test] + async fn module_retries_cleanup_after_stop_is_cancelled() { + let events = Arc::new(Mutex::new(Vec::new())); + let blocking = Arc::new(CancelOnceStopEngine { + stop_calls: AtomicUsize::new(0), + stop_entered: Notify::new(), + release_first_stop: Notify::new(), + events: events.clone(), + }); + let runtime = wrapped_transport_runtime( + WrappedTransportDirections { + source: true, + destination: true, + }, + WrappedTransportDirections { + source: true, + destination: true, + }, + ); + let module = WrappedTransportProxyModule::new_without_sources( + wrapped_transport_peer_manager(), + runtime, + Some(wrapped_transport_engine("kcp", false, &events)), + Some(blocking.clone()), + ) + .unwrap(); + module.start().await.unwrap(); + + let stop_task = tokio::spawn({ + let module = module.clone(); + async move { module.stop().await } + }); + blocking.stop_entered.notified().await; + stop_task.abort(); + assert!(stop_task.await.unwrap_err().is_cancelled()); + + module.stop().await; + + assert_eq!(blocking.stop_calls.load(Ordering::Acquire), 2); + assert_eq!( + *events.lock().unwrap(), + [ + "prepare:kcp:true:true", + "activate:kcp", + "prepare:blocking", + "activate:blocking", + "stop:blocking", + "stop:blocking", + "stop:kcp", + ] + ); + } + + #[tokio::test] + async fn module_cleans_partial_start_before_retry() { + let events = Arc::new(Mutex::new(Vec::new())); + let blocking = Arc::new(AbortOncePrepareEngine::default()); + let runtime = wrapped_transport_runtime( + WrappedTransportDirections { + source: true, + destination: true, + }, + WrappedTransportDirections { + source: true, + destination: true, + }, + ); + let module = WrappedTransportProxyModule::new_without_sources( + wrapped_transport_peer_manager(), + runtime, + Some(wrapped_transport_engine("kcp", false, &events)), + Some(blocking.clone()), + ) + .unwrap(); + + let first_start = tokio::spawn({ + let module = module.clone(); + async move { module.start().await } + }); + blocking.first_prepare_entered.notified().await; + first_start.abort(); + assert!(first_start.await.unwrap_err().is_cancelled()); + #[cfg(feature = "proxy-packet")] + assert!(!module.source_connect_ready(WrappedTransportKind::Kcp).await); + + module.start().await.unwrap(); + + assert_eq!(blocking.prepare_calls.load(Ordering::Acquire), 2); + assert_eq!(blocking.activate_calls.load(Ordering::Acquire), 1); + assert_eq!(blocking.stop_calls.load(Ordering::Acquire), 1); + assert_eq!( + *events.lock().unwrap(), + [ + "prepare:kcp:true:true", + "activate:kcp", + "stop:kcp", + "prepare:kcp:true:true", + "activate:kcp", + ] + ); + + module.stop().await; + assert_eq!(blocking.stop_calls.load(Ordering::Acquire), 2); + } +} diff --git a/easytier-core/src/gateway/proxy/wrapped_transport/connect_api.rs b/easytier-core/src/gateway/proxy/wrapped_transport/connect_api.rs new file mode 100644 index 00000000..143cc43d --- /dev/null +++ b/easytier-core/src/gateway/proxy/wrapped_transport/connect_api.rs @@ -0,0 +1,56 @@ +use std::net::SocketAddr; + +use crate::gateway::proxy::traits::TcpProxyStream; + +use super::{ + WrappedTransportKind, WrappedTransportPacketPlane, WrappedTransportPacketState, + WrappedTransportProxyModule, packet_plane::connect_wrapped_transport_source, +}; + +impl WrappedTransportPacketPlane { + fn source_connect_ready( + &self, + state: &WrappedTransportPacketState, + transport: WrappedTransportKind, + ) -> bool { + match transport { + WrappedTransportKind::Kcp => state.kcp_source_connect_ready, + WrappedTransportKind::Quic => state.quic_source_connect_ready, + } + } +} + +impl WrappedTransportProxyModule { + pub(crate) async fn source_connect_ready(&self, transport: WrappedTransportKind) -> bool { + let state = self.state.lock().await; + state.active + && self + .packet_plane + .source_connect_ready(&state.packet, transport) + } + + pub(crate) async fn connect_source( + &self, + transport: WrappedTransportKind, + src: SocketAddr, + dst: SocketAddr, + ) -> anyhow::Result> { + let engine = { + let state = self.state.lock().await; + let ready = state.active + && self + .packet_plane + .source_connect_ready(&state.packet, transport); + if !ready { + anyhow::bail!("{transport:?} source is not ready"); + } + match transport { + WrappedTransportKind::Kcp => self.kcp.clone(), + WrappedTransportKind::Quic => self.quic.clone(), + } + } + .ok_or_else(|| anyhow::anyhow!("{transport:?} engine is not available"))?; + + connect_wrapped_transport_source(&self.peer_manager, engine, src, dst).await + } +} diff --git a/easytier-core/src/gateway/proxy/wrapped_transport/engine_api.rs b/easytier-core/src/gateway/proxy/wrapped_transport/engine_api.rs new file mode 100644 index 00000000..1de1c48c --- /dev/null +++ b/easytier-core/src/gateway/proxy/wrapped_transport/engine_api.rs @@ -0,0 +1,56 @@ +#[cfg(feature = "proxy-packet")] +use std::net::SocketAddr; + +use async_trait::async_trait; +use bytes::Bytes; + +#[cfg(feature = "proxy-packet")] +use crate::gateway::proxy::{ + traits::TcpProxyStream, wrapped_transport_destination::WrappedTransportDestinationIngress, +}; + +use super::{WrappedTransportDatagram, WrappedTransportDirections, WrappedTransportRole}; + +#[derive(Clone)] +pub struct WrappedTransportEngineStart { + pub directions: WrappedTransportDirections, + pub my_peer_id: u32, + pub datagrams: tokio::sync::mpsc::Sender, + #[cfg(feature = "proxy-packet")] + pub destination_ingress: Option, +} + +#[cfg(feature = "proxy-packet")] +#[derive(Debug, Clone, Copy)] +pub struct WrappedTransportConnect { + pub my_peer_id: u32, + pub dst_peer_id: u32, + pub src: SocketAddr, + pub dst: SocketAddr, +} + +#[cfg(feature = "proxy-packet")] +pub struct WrappedTransportAcceptedStream { + pub src: SocketAddr, + pub dst: SocketAddr, + pub initial_acl_packet_size: usize, + pub stream: Box, +} + +#[async_trait] +pub trait WrappedTransportEngine: Send + Sync + 'static { + async fn prepare(&self, options: WrappedTransportEngineStart) -> anyhow::Result<()>; + async fn activate(&self) -> anyhow::Result<()>; + async fn inject_peer_datagram( + &self, + role: WrappedTransportRole, + from_peer_id: u32, + payload: Bytes, + ) -> anyhow::Result<()>; + #[cfg(feature = "proxy-packet")] + async fn connect_source( + &self, + request: WrappedTransportConnect, + ) -> anyhow::Result>; + async fn stop(&self); +} diff --git a/easytier-core/src/gateway/proxy/wrapped_transport/packet_api.rs b/easytier-core/src/gateway/proxy/wrapped_transport/packet_api.rs new file mode 100644 index 00000000..d9a00a2e --- /dev/null +++ b/easytier-core/src/gateway/proxy/wrapped_transport/packet_api.rs @@ -0,0 +1,43 @@ +use crate::gateway::proxy::tcp_proxy_engine::TcpNatEntrySnapshot; + +use super::{WrappedTransportKind, WrappedTransportProxyModule}; + +impl WrappedTransportProxyModule { + pub(crate) fn source_entry_snapshots( + &self, + transport: WrappedTransportKind, + ) -> Vec { + match transport { + WrappedTransportKind::Kcp => self.packet_plane.kcp_source.as_ref(), + WrappedTransportKind::Quic => self.packet_plane.quic_source.as_ref(), + } + .map_or_else(Vec::new, |source| source.entry_snapshots()) + } + + pub(crate) fn source_is_started(&self, transport: WrappedTransportKind) -> bool { + match transport { + WrappedTransportKind::Kcp => self.packet_plane.kcp_source.as_ref(), + WrappedTransportKind::Quic => self.packet_plane.quic_source.as_ref(), + } + .is_some_and(|source| source.is_started()) + } + + pub(crate) fn destination_entry_snapshots( + &self, + transport: WrappedTransportKind, + ) -> Vec { + self.packet_plane + .destination + .as_ref() + .map_or_else(Vec::new, |destination| { + destination.entry_snapshots(transport) + }) + } + + pub(crate) fn destination_is_started(&self, transport: WrappedTransportKind) -> bool { + self.packet_plane + .destination + .as_ref() + .is_some_and(|destination| destination.is_started(transport)) + } +} diff --git a/easytier-core/src/gateway/proxy/wrapped_transport/packet_plane.rs b/easytier-core/src/gateway/proxy/wrapped_transport/packet_plane.rs new file mode 100644 index 00000000..5432bd3a --- /dev/null +++ b/easytier-core/src/gateway/proxy/wrapped_transport/packet_plane.rs @@ -0,0 +1,531 @@ +use std::{ + net::{IpAddr, SocketAddr}, + sync::{Arc, Mutex as StdMutex, Weak}, +}; + +use async_trait::async_trait; + +use crate::{ + config::runtime::CoreRuntimeConfigStore, + connectivity::direct::DirectConnectorHost, + connectivity::hole_punch::tcp::TcpHolePunchHost, + gateway::proxy::{ + cidr_table::ProxyCidrTable, + service::CoreProxyRuntime, + tcp_proxy_engine::{TcpNatEntrySnapshot, TcpProxyMode, TcpProxyNicContext}, + tcp_proxy_service::TcpProxyService, + traits::{ProxyRuntimeInfo, TcpProxyDestinationConnector, TcpProxyStream}, + wrapped_tcp_proxy::{ + WrappedTcpProxyNicContext, WrappedTcpProxyTransport, + try_process_wrapped_tcp_packet_from_nic, + }, + wrapped_transport_destination::{ + WrappedTransportDestinationIngresses, WrappedTransportDestinationLifecycle, + WrappedTransportDestinationModule, + }, + }, + listener::RunningListenerRegistry, + packet::ZCPacket, + peers::peer_manager::{PeerManagerCore, PipelineRegistrationGuard}, + process_runtime::ProtectedTcpPortRegistry, + socket::SocketContext, +}; + +pub use crate::gateway::proxy::wrapped_transport_destination::WrappedTransportDestinationIngress; + +use super::{ + WrappedTransportConnect, WrappedTransportDatagram, WrappedTransportDirections, + WrappedTransportEngine, WrappedTransportEngineStart, WrappedTransportKind, +}; + +#[derive(Default)] +pub(super) struct WrappedTransportPacketState { + pub(super) kcp_source_connect_ready: bool, + pub(super) quic_source_connect_ready: bool, + pub(super) kcp_source_started: bool, + pub(super) quic_source_started: bool, + pub(super) destination_started: bool, + pub(super) kcp_source_guard: Option, + pub(super) quic_source_guard: Option, + destination_ingresses: WrappedTransportDestinationIngresses, +} + +impl WrappedTransportPacketState { + fn has_partial_start(&self) -> bool { + self.kcp_source_started || self.quic_source_started || self.destination_started + } +} + +#[derive(Default)] +pub(super) struct WrappedTransportPacketPlane { + pub(super) kcp_source: Option>, + pub(super) quic_source: Option>, + pub(super) destination: Option>, +} + +#[derive(Clone)] +struct WrappedTransportSourceConnector { + peer_manager: Arc, + engine: Weak, + transport: WrappedTransportKind, +} + +pub(super) async fn connect_wrapped_transport_source( + peer_manager: &PeerManagerCore, + engine: Arc, + src: SocketAddr, + dst: SocketAddr, +) -> anyhow::Result> { + let SocketAddr::V4(dst_v4) = dst else { + anyhow::bail!("IPv6 is not supported by wrapped TCP proxy"); + }; + let dst_peer_id = peer_manager + .get_peer_map() + .get_peer_id_by_ipv4(dst_v4.ip()) + .await + .ok_or_else(|| anyhow::anyhow!("no peer found for wrapped TCP dst: {dst}"))?; + engine + .connect_source(WrappedTransportConnect { + my_peer_id: peer_manager.my_peer_id(), + dst_peer_id, + src, + dst, + }) + .await +} + +#[async_trait] +impl TcpProxyDestinationConnector for WrappedTransportSourceConnector { + type DstStream = Box; + + async fn connect(&self, src: SocketAddr, dst: SocketAddr) -> anyhow::Result { + let engine = self + .engine + .upgrade() + .ok_or_else(|| anyhow::anyhow!("wrapped transport engine is not available"))?; + connect_wrapped_transport_source(&self.peer_manager, engine, src, dst).await + } + + fn proxy_mode(&self) -> TcpProxyMode { + match self.transport { + WrappedTransportKind::Kcp => TcpProxyMode::KcpSrc, + WrappedTransportKind::Quic => TcpProxyMode::QuicSrc, + } + } +} + +type WrappedTransportSourceService = + TcpProxyService, H, WrappedTransportSourceConnector>; + +pub(super) struct WrappedTransportSource +where + H: DirectConnectorHost + TcpHolePunchHost, +{ + peer_manager: Arc, + runtime: Arc>, + host: Arc, + connector: Arc, + cidr_table: Arc, + socket_context: SocketContext, + service: StdMutex>>>, + transport: WrappedTransportKind, +} + +impl WrappedTransportSource +where + H: DirectConnectorHost + TcpHolePunchHost, +{ + #[allow(clippy::too_many_arguments)] + pub(super) fn new( + peer_manager: Arc, + host: Arc, + protected_tcp_ports: Arc, + running_listeners: Arc, + runtime_config: CoreRuntimeConfigStore, + cidr_table: Arc, + socket_context: SocketContext, + engine: &Arc, + transport: WrappedTransportKind, + ) -> Arc { + let protocol_label = match transport { + WrappedTransportKind::Kcp => "KCP", + WrappedTransportKind::Quic => "QUIC", + }; + let runtime = CoreProxyRuntime::new( + peer_manager.clone(), + host.clone(), + protected_tcp_ports, + running_listeners, + runtime_config, + protocol_label, + ); + let connector = Arc::new(WrappedTransportSourceConnector { + peer_manager: peer_manager.clone(), + engine: Arc::downgrade(engine), + transport, + }); + Arc::new(Self { + peer_manager, + runtime, + host, + connector, + cidr_table, + socket_context, + service: StdMutex::new(None), + transport, + }) + } + + fn build_service(&self) -> Arc> { + TcpProxyService::new_with_socket_context( + self.peer_manager.clone(), + self.runtime.clone(), + self.host.clone(), + self.connector.clone(), + self.cidr_table.clone(), + self.socket_context.clone(), + ) + } + + async fn start_service(self: &Arc) -> Result { + let service = { + let mut active = self.service.lock().unwrap(); + if active.is_some() { + anyhow::bail!("wrapped transport source is already started"); + } + let service = self.build_service(); + active.replace(service.clone()); + service + }; + self.runtime.latch_smoltcp(); + let guard = self + .peer_manager + .add_managed_nic_packet_process_pipeline(Box::new(WrappedTransportSourceFilter { + source: Arc::downgrade(self), + })) + .await; + service.register_peer_pipeline().await; + if let Err(error) = service.start(false).await { + guard.close(); + self.stop_service().await; + return Err(anyhow::Error::new(error)); + } + Ok(guard) + } + + async fn stop_service(&self) { + let service = self.service.lock().unwrap().take(); + if let Some(service) = service { + service.stop_and_wait().await; + } + } + + fn source_entry_snapshots(&self) -> Vec { + self.service + .lock() + .unwrap() + .as_ref() + .map_or_else(Vec::new, |service| service.engine().list_entries()) + } + + async fn try_process_nic_packet(&self, packet: &mut ZCPacket) -> bool { + let Some(service) = self.service.lock().unwrap().clone() else { + return false; + }; + let snapshot = self.runtime.proxy_runtime_snapshot(); + let engine = service.engine(); + if engine.try_process_packet_from_nic( + packet, + TcpProxyNicContext { + local_inet: snapshot.local_inet, + local_port: engine.local_port(), + my_peer_id: self.peer_manager.my_peer_id(), + smoltcp_enabled: snapshot.smoltcp_enabled, + }, + ) { + return true; + } + + let connection_engine = engine.clone(); + let peer_manager = self.peer_manager.clone(); + let transport = self.transport; + try_process_wrapped_tcp_packet_from_nic( + packet, + WrappedTcpProxyNicContext { + transport: match transport { + WrappedTransportKind::Kcp => WrappedTcpProxyTransport::Kcp, + WrappedTransportKind::Quic => WrappedTcpProxyTransport::Quic, + }, + my_peer_id: self.peer_manager.my_peer_id(), + local_ipv4: snapshot.local_inet.map(|inet| inet.address()), + smoltcp_enabled: snapshot.smoltcp_enabled, + }, + move |src| connection_engine.is_tcp_proxy_connection(src), + move |dst_ip| async move { + match transport { + WrappedTransportKind::Kcp => { + peer_manager + .check_allow_kcp_to_dst(&IpAddr::V4(dst_ip)) + .await + } + WrappedTransportKind::Quic => { + peer_manager + .check_allow_quic_to_dst(&IpAddr::V4(dst_ip)) + .await + } + } + }, + ) + .await + } +} + +#[async_trait] +pub(super) trait WrappedTransportSourceLifecycle: Send + Sync { + async fn start(self: Arc) -> anyhow::Result; + async fn stop(&self); + fn entry_snapshots(&self) -> Vec; + fn is_started(&self) -> bool; +} + +#[async_trait] +impl WrappedTransportSourceLifecycle for WrappedTransportSource +where + H: DirectConnectorHost + TcpHolePunchHost, +{ + async fn start(self: Arc) -> anyhow::Result { + self.start_service().await + } + + async fn stop(&self) { + self.stop_service().await; + } + + fn entry_snapshots(&self) -> Vec { + self.source_entry_snapshots() + } + + fn is_started(&self) -> bool { + self.service + .lock() + .unwrap() + .as_ref() + .is_some_and(|service| service.is_started()) + } +} + +struct WrappedTransportSourceFilter +where + H: DirectConnectorHost + TcpHolePunchHost, +{ + source: Weak>, +} + +#[async_trait] +impl crate::peers::NicPacketFilter for WrappedTransportSourceFilter +where + H: DirectConnectorHost + TcpHolePunchHost, +{ + async fn try_process_packet_from_nic(&self, packet: &mut ZCPacket) -> bool { + let Some(source) = self.source.upgrade() else { + return false; + }; + source.try_process_nic_packet(packet).await + } +} + +impl WrappedTransportPacketPlane { + pub(super) fn has_partial_start(&self, state: &WrappedTransportPacketState) -> bool { + state.has_partial_start() + } + + pub(super) fn clear_connect_ready(&self, state: &mut WrappedTransportPacketState) { + state.kcp_source_connect_ready = false; + state.quic_source_connect_ready = false; + } + + pub(super) async fn stop_source( + &self, + state: &mut WrappedTransportPacketState, + transport: WrappedTransportKind, + ) { + let (source, started, guard) = match transport { + WrappedTransportKind::Kcp => ( + &self.kcp_source, + &mut state.kcp_source_started, + &mut state.kcp_source_guard, + ), + WrappedTransportKind::Quic => ( + &self.quic_source, + &mut state.quic_source_started, + &mut state.quic_source_guard, + ), + }; + if let Some(guard) = guard.take() { + guard.close(); + } + if *started { + if let Some(source) = source { + source.stop().await; + } + *started = false; + } + } + + pub(super) async fn stop_destination(&self, state: &mut WrappedTransportPacketState) { + if state.destination_started { + if let Some(destination) = &self.destination { + destination.stop().await; + } + state.destination_started = false; + } + state.destination_ingresses = WrappedTransportDestinationIngresses::default(); + } + + pub(super) async fn start_destinations( + &self, + state: &mut WrappedTransportPacketState, + kcp_directions: WrappedTransportDirections, + quic_directions: WrappedTransportDirections, + kcp_available: bool, + quic_available: bool, + ) -> anyhow::Result<()> { + let kcp = kcp_directions.destination && kcp_available; + let quic = quic_directions.destination && quic_available; + state.destination_ingresses = if kcp || quic { + state.destination_started = true; + self.destination + .as_ref() + .expect("wrapped destination owner must match its engines") + .clone() + .start(kcp, quic) + .await? + } else { + WrappedTransportDestinationIngresses::default() + }; + Ok(()) + } + + pub(super) fn engine_start( + &self, + state: &WrappedTransportPacketState, + transport: WrappedTransportKind, + directions: WrappedTransportDirections, + my_peer_id: u32, + datagrams: tokio::sync::mpsc::Sender, + ) -> WrappedTransportEngineStart { + let destination_ingress = match transport { + WrappedTransportKind::Kcp => state.destination_ingresses.kcp.clone(), + WrappedTransportKind::Quic => state.destination_ingresses.quic.clone(), + }; + WrappedTransportEngineStart { + directions, + my_peer_id, + datagrams, + destination_ingress, + } + } + + pub(super) async fn start_source( + &self, + state: &mut WrappedTransportPacketState, + transport: WrappedTransportKind, + directions: WrappedTransportDirections, + ) -> anyhow::Result<()> { + if !directions.source { + return Ok(()); + } + let (source, started, guard) = match transport { + WrappedTransportKind::Kcp => ( + &self.kcp_source, + &mut state.kcp_source_started, + &mut state.kcp_source_guard, + ), + WrappedTransportKind::Quic => ( + &self.quic_source, + &mut state.quic_source_started, + &mut state.quic_source_guard, + ), + }; + let Some(source) = source else { + return Ok(()); + }; + *started = true; + *guard = Some(source.clone().start().await?); + Ok(()) + } + + pub(super) fn mark_source_connect_ready( + &self, + state: &mut WrappedTransportPacketState, + transport: WrappedTransportKind, + directions: WrappedTransportDirections, + ) { + if !directions.source { + return; + } + match transport { + WrappedTransportKind::Kcp => state.kcp_source_connect_ready = true, + WrappedTransportKind::Quic => state.quic_source_connect_ready = true, + } + } + + #[allow(clippy::too_many_arguments)] + pub(super) fn new( + peer_manager: Arc, + runtime_config: CoreRuntimeConfigStore, + kcp: &Option>, + quic: &Option>, + host: Arc, + protected_tcp_ports: Arc, + running_listeners: Arc, + cidr_table: Arc, + socket_context: SocketContext, + ) -> Self + where + H: DirectConnectorHost + TcpHolePunchHost, + { + let kcp_source = kcp.as_ref().map(|engine| { + WrappedTransportSource::new( + peer_manager.clone(), + host.clone(), + protected_tcp_ports.clone(), + running_listeners.clone(), + runtime_config.clone(), + cidr_table.clone(), + socket_context.clone(), + engine, + WrappedTransportKind::Kcp, + ) as Arc + }); + let quic_source = quic.as_ref().map(|engine| { + WrappedTransportSource::new( + peer_manager.clone(), + host.clone(), + protected_tcp_ports.clone(), + running_listeners.clone(), + runtime_config.clone(), + cidr_table.clone(), + socket_context.clone(), + engine, + WrappedTransportKind::Quic, + ) as Arc + }); + let destination = (kcp.is_some() || quic.is_some()).then(|| { + WrappedTransportDestinationModule::new( + peer_manager, + host, + protected_tcp_ports, + running_listeners, + runtime_config, + cidr_table, + socket_context, + ) as Arc + }); + Self { + kcp_source, + quic_source, + destination, + } + } +} diff --git a/easytier-core/src/gateway/proxy/wrapped_transport/packet_plane_disabled.rs b/easytier-core/src/gateway/proxy/wrapped_transport/packet_plane_disabled.rs new file mode 100644 index 00000000..9b5a12a8 --- /dev/null +++ b/easytier-core/src/gateway/proxy/wrapped_transport/packet_plane_disabled.rs @@ -0,0 +1,97 @@ +use std::sync::Arc; + +use crate::{ + config::runtime::CoreRuntimeConfigStore, connectivity::direct::DirectConnectorHost, + connectivity::hole_punch::tcp::TcpHolePunchHost, gateway::proxy::cidr_table::ProxyCidrTable, + listener::RunningListenerRegistry, peers::peer_manager::PeerManagerCore, + process_runtime::ProtectedTcpPortRegistry, socket::SocketContext, +}; + +use super::{ + WrappedTransportDatagram, WrappedTransportDirections, WrappedTransportEngine, + WrappedTransportEngineStart, WrappedTransportKind, +}; + +#[derive(Default)] +pub(super) struct WrappedTransportPacketState; + +#[derive(Default)] +pub(super) struct WrappedTransportPacketPlane; + +impl WrappedTransportPacketPlane { + pub(super) fn has_partial_start(&self, _state: &WrappedTransportPacketState) -> bool { + false + } + + pub(super) fn clear_connect_ready(&self, _state: &mut WrappedTransportPacketState) {} + + pub(super) async fn stop_source( + &self, + _state: &mut WrappedTransportPacketState, + _transport: WrappedTransportKind, + ) { + } + + pub(super) async fn stop_destination(&self, _state: &mut WrappedTransportPacketState) {} + + pub(super) async fn start_destinations( + &self, + _state: &mut WrappedTransportPacketState, + _kcp_directions: WrappedTransportDirections, + _quic_directions: WrappedTransportDirections, + _kcp_available: bool, + _quic_available: bool, + ) -> anyhow::Result<()> { + Ok(()) + } + + pub(super) fn engine_start( + &self, + _state: &WrappedTransportPacketState, + _transport: WrappedTransportKind, + directions: WrappedTransportDirections, + my_peer_id: u32, + datagrams: tokio::sync::mpsc::Sender, + ) -> WrappedTransportEngineStart { + WrappedTransportEngineStart { + directions, + my_peer_id, + datagrams, + } + } + + pub(super) async fn start_source( + &self, + _state: &mut WrappedTransportPacketState, + _transport: WrappedTransportKind, + _directions: WrappedTransportDirections, + ) -> anyhow::Result<()> { + Ok(()) + } + + pub(super) fn mark_source_connect_ready( + &self, + _state: &mut WrappedTransportPacketState, + _transport: WrappedTransportKind, + _directions: WrappedTransportDirections, + ) { + } + + #[allow(clippy::too_many_arguments)] + pub(super) fn new( + _peer_manager: Arc, + _runtime_config: CoreRuntimeConfigStore, + _kcp: &Option>, + _quic: &Option>, + _host: Arc, + _protected_tcp_ports: Arc, + _running_listeners: Arc, + _cidr_table: Arc, + _socket_context: SocketContext, + ) -> Self + where + H: DirectConnectorHost + TcpHolePunchHost, + { + Self + } +} diff --git a/easytier-core/src/gateway/proxy/wrapped_transport_destination.rs b/easytier-core/src/gateway/proxy/wrapped_transport_destination.rs new file mode 100644 index 00000000..1c8c4c5e --- /dev/null +++ b/easytier-core/src/gateway/proxy/wrapped_transport_destination.rs @@ -0,0 +1,375 @@ +use std::{ + net::SocketAddr, + sync::{ + Arc, + atomic::{AtomicBool, AtomicU8, Ordering}, + }, + time::{SystemTime, UNIX_EPOCH}, +}; + +use dashmap::DashMap; +use guarden::defer; +use tokio::{sync::Mutex, task::JoinSet}; +use tokio_util::sync::CancellationToken; + +use crate::{ + config::runtime::CoreRuntimeConfigStore, connectivity::direct::DirectConnectorHost, + connectivity::hole_punch::tcp::TcpHolePunchHost, listener::RunningListenerRegistry, + peers::peer_manager::PeerManagerCore, process_runtime::ProtectedTcpPortRegistry, + socket::SocketContext, +}; + +use super::{ + cidr_table::ProxyCidrTable, + service::CoreProxyRuntime, + tcp_proxy_engine::{TcpNatEntrySnapshot, TcpNatEntryState}, + tcp_socket_connector::TcpSocketProxyConnector, + traits::TcpProxyDestinationConnector, + wrapped_tcp_proxy::{WrappedTcpDestinationRequest, plan_wrapped_tcp_destination}, + wrapped_transport::{WrappedTransportAcceptedStream, WrappedTransportKind}, +}; + +struct DestinationEntry { + transport: WrappedTransportKind, + src: SocketAddr, + dst: SocketAddr, + mapped_dst: std::sync::RwLock, + start_time: u64, + state: AtomicU8, +} + +impl DestinationEntry { + fn new(transport: WrappedTransportKind, src: SocketAddr, dst: SocketAddr) -> Self { + Self { + transport, + src, + dst, + mapped_dst: std::sync::RwLock::new(dst), + start_time: SystemTime::now() + .duration_since(UNIX_EPOCH) + .map(|duration| duration.as_secs()) + .unwrap_or_default(), + state: AtomicU8::new(TcpNatEntryState::ConnectingDst as u8), + } + } + + fn set_mapped_dst(&self, mapped_dst: SocketAddr) { + *self + .mapped_dst + .write() + .expect("wrapped destination entry mutex poisoned") = mapped_dst; + } + + fn set_state(&self, state: TcpNatEntryState) { + self.state.store(state as u8, Ordering::Release); + } + + fn snapshot(&self) -> TcpNatEntrySnapshot { + let state = match self.state.load(Ordering::Acquire) { + value if value == TcpNatEntryState::Connected as u8 => TcpNatEntryState::Connected, + _ => TcpNatEntryState::ConnectingDst, + }; + TcpNatEntrySnapshot { + src: self.src, + dst: self.dst, + mapped_dst: *self + .mapped_dst + .read() + .expect("wrapped destination entry mutex poisoned"), + start_time: self.start_time, + state, + } + } +} + +struct AcceptedDestination { + transport: WrappedTransportKind, + stream: WrappedTransportAcceptedStream, +} + +#[derive(Clone)] +pub struct WrappedTransportDestinationIngress { + transport: WrappedTransportKind, + accepted: tokio::sync::mpsc::Sender, +} + +impl WrappedTransportDestinationIngress { + pub async fn submit(&self, stream: WrappedTransportAcceptedStream) -> anyhow::Result<()> { + self.accepted + .send(AcceptedDestination { + transport: self.transport, + stream, + }) + .await + .map_err(|_| anyhow::anyhow!("wrapped destination ingress is stopped")) + } +} + +#[derive(Default)] +pub(crate) struct WrappedTransportDestinationIngresses { + pub(crate) kcp: Option, + pub(crate) quic: Option, +} + +struct DestinationRun { + cancel: CancellationToken, + supervisor: tokio::task::JoinHandle<()>, +} + +async fn stop_destination_run(run: &Mutex>) { + let mut run = run.lock().await; + if let Some(current) = run.as_mut() { + current.cancel.cancel(); + let _ = (&mut current.supervisor).await; + run.take(); + } +} + +pub(crate) struct WrappedTransportDestinationModule +where + H: DirectConnectorHost + TcpHolePunchHost, +{ + runtime: Arc>, + connector: TcpSocketProxyConnector, + cidr_table: Arc, + peer_manager: Arc, + entries: DashMap>, + run: Mutex>, + active: AtomicBool, + enabled: AtomicU8, +} + +#[async_trait::async_trait] +pub(crate) trait WrappedTransportDestinationLifecycle: Send + Sync { + async fn start( + self: Arc, + kcp: bool, + quic: bool, + ) -> anyhow::Result; + async fn stop(&self); + fn entry_snapshots(&self, transport: WrappedTransportKind) -> Vec; + fn is_started(&self, transport: WrappedTransportKind) -> bool; +} + +impl WrappedTransportDestinationModule +where + H: DirectConnectorHost + TcpHolePunchHost, +{ + const INGRESS_CAPACITY: usize = 128; + + pub(crate) fn new( + peer_manager: Arc, + host: Arc, + protected_tcp_ports: Arc, + running_listeners: Arc, + runtime_config: CoreRuntimeConfigStore, + cidr_table: Arc, + socket_context: SocketContext, + ) -> Arc { + Arc::new(Self { + runtime: CoreProxyRuntime::new( + peer_manager.clone(), + host.clone(), + protected_tcp_ports, + running_listeners, + runtime_config, + "TCP", + ), + connector: TcpSocketProxyConnector::new(host).with_socket_context(socket_context), + cidr_table, + peer_manager, + entries: DashMap::new(), + run: Mutex::new(None), + active: AtomicBool::new(false), + enabled: AtomicU8::new(0), + }) + } + + pub(crate) async fn start( + self: &Arc, + kcp: bool, + quic: bool, + ) -> anyhow::Result { + let mut run = self.run.lock().await; + if run.is_some() { + anyhow::bail!("wrapped transport destination is already started"); + } + + let (accepted, mut receiver) = tokio::sync::mpsc::channel(Self::INGRESS_CAPACITY); + let cancel = CancellationToken::new(); + let task_cancel = cancel.clone(); + let owner = Arc::downgrade(self); + let supervisor = tokio::spawn(async move { + let mut sessions = JoinSet::new(); + loop { + tokio::select! { + biased; + _ = task_cancel.cancelled() => break, + accepted = receiver.recv() => { + let Some(accepted) = accepted else { break }; + let Some(owner) = owner.upgrade() else { break }; + sessions.spawn(async move { + if let Err(error) = owner.handle_destination(accepted).await { + tracing::debug!(?error, "wrapped destination session failed"); + } + }); + } + _ = sessions.join_next(), if !sessions.is_empty() => {} + } + } + sessions.shutdown().await; + }); + self.active.store(true, Ordering::Release); + self.enabled + .store((kcp as u8) | ((quic as u8) << 1), Ordering::Release); + *run = Some(DestinationRun { cancel, supervisor }); + + let ingress = |transport| WrappedTransportDestinationIngress { + transport, + accepted: accepted.clone(), + }; + Ok(WrappedTransportDestinationIngresses { + kcp: kcp.then(|| ingress(WrappedTransportKind::Kcp)), + quic: quic.then(|| ingress(WrappedTransportKind::Quic)), + }) + } + + pub(crate) async fn stop(&self) { + self.active.store(false, Ordering::Release); + self.enabled.store(0, Ordering::Release); + stop_destination_run(&self.run).await; + self.entries.clear(); + } + + pub(crate) fn entry_snapshots( + &self, + transport: WrappedTransportKind, + ) -> Vec { + self.entries + .iter() + .filter(|entry| entry.value().transport == transport) + .map(|entry| entry.value().snapshot()) + .collect() + } + + pub(crate) fn is_started(&self, transport: WrappedTransportKind) -> bool { + let mask = match transport { + WrappedTransportKind::Kcp => 1, + WrappedTransportKind::Quic => 2, + }; + self.active.load(Ordering::Acquire) && self.enabled.load(Ordering::Acquire) & mask != 0 + } + + async fn handle_destination(&self, accepted: AcceptedDestination) -> anyhow::Result<()> { + if !self.active.load(Ordering::Acquire) { + anyhow::bail!("wrapped transport destination is not active"); + } + + let entry_id = uuid::Uuid::new_v4(); + let entry = Arc::new(DestinationEntry::new( + accepted.transport, + accepted.stream.src, + accepted.stream.dst, + )); + self.entries.insert(entry_id, entry.clone()); + defer! { + self.remove_entry(entry_id); + } + + let route = self.peer_manager.get_route(); + let plan = plan_wrapped_tcp_destination( + WrappedTcpDestinationRequest { + src: accepted.stream.src, + dst: accepted.stream.dst, + initial_packet_size: accepted.stream.initial_acl_packet_size, + }, + self.cidr_table.as_ref(), + self.runtime.as_ref(), + route.as_ref(), + self.peer_manager.acl_filter(), + ) + .await?; + entry.set_mapped_dst(plan.socket_dst); + + tracing::debug!(dst = ?plan.socket_dst, "wrapped transport connect to destination"); + let destination = self + .connector + .connect("0.0.0.0:0".parse().unwrap(), plan.socket_dst) + .await?; + entry.set_state(TcpNatEntryState::Connected); + + plan.acl_handler + .copy_bidirection_with_acl(accepted.stream.stream, destination) + .await + } + + fn remove_entry(&self, id: uuid::Uuid) { + self.entries.remove(&id); + if self.entries.capacity() - self.entries.len() > 16 { + self.entries.shrink_to_fit(); + } + } +} + +#[async_trait::async_trait] +impl WrappedTransportDestinationLifecycle for WrappedTransportDestinationModule +where + H: DirectConnectorHost + TcpHolePunchHost, +{ + async fn start( + self: Arc, + kcp: bool, + quic: bool, + ) -> anyhow::Result { + WrappedTransportDestinationModule::start(&self, kcp, quic).await + } + + async fn stop(&self) { + WrappedTransportDestinationModule::stop(self).await; + } + + fn entry_snapshots(&self, transport: WrappedTransportKind) -> Vec { + WrappedTransportDestinationModule::entry_snapshots(self, transport) + } + + fn is_started(&self, transport: WrappedTransportKind) -> bool { + WrappedTransportDestinationModule::is_started(self, transport) + } +} + +#[cfg(test)] +mod tests { + use super::*; + use tokio::sync::Notify; + + #[tokio::test] + async fn cancelled_stop_retains_supervisor_for_retry() { + let cancel = CancellationToken::new(); + let stop_entered = Arc::new(Notify::new()); + let release_stop = Arc::new(Notify::new()); + let supervisor = tokio::spawn({ + let cancel = cancel.clone(); + let stop_entered = stop_entered.clone(); + let release_stop = release_stop.clone(); + async move { + cancel.cancelled().await; + stop_entered.notify_one(); + release_stop.notified().await; + } + }); + let run = Arc::new(Mutex::new(Some(DestinationRun { cancel, supervisor }))); + + let first_stop = tokio::spawn({ + let run = run.clone(); + async move { stop_destination_run(run.as_ref()).await } + }); + stop_entered.notified().await; + first_stop.abort(); + assert!(first_stop.await.unwrap_err().is_cancelled()); + assert!(run.lock().await.is_some()); + + release_stop.notify_one(); + stop_destination_run(run.as_ref()).await; + assert!(run.lock().await.is_none()); + } +} diff --git a/easytier-core/src/gateway/smoltcp/mod.rs b/easytier-core/src/gateway/smoltcp/mod.rs new file mode 100644 index 00000000..9e8128f4 --- /dev/null +++ b/easytier-core/src/gateway/smoltcp/mod.rs @@ -0,0 +1,7 @@ +mod stack; +mod tokio_smoltcp; + +pub(super) use stack::{SmolTcpStack, output_dst_ip}; +pub(super) use tokio_smoltcp::{ + BufferSize, Net, NetConfig, TcpListener, UdpSocket, channel_device, +}; diff --git a/easytier-core/src/gateway/smoltcp/stack.rs b/easytier-core/src/gateway/smoltcp/stack.rs new file mode 100644 index 00000000..e5b373a0 --- /dev/null +++ b/easytier-core/src/gateway/smoltcp/stack.rs @@ -0,0 +1,158 @@ +use std::net::{IpAddr, Ipv4Addr, SocketAddr}; +use std::sync::Arc; + +use tokio::sync::{Mutex, mpsc}; +use tokio::task::JoinSet; + +use crate::{ + foundation::time::{Duration, timeout}, + packet::ZCPacket, +}; + +use super::tokio_smoltcp::{BufferSize, Net, NetConfig, channel_device}; +use crate::gateway::proxy::traits::TcpProxyStream; + +type SmolTcpAcceptResult = anyhow::Result<(super::tokio_smoltcp::TcpStream, SocketAddr)>; + +pub struct SmolTcpStack { + ingress_tx: mpsc::Sender, + output_rx: Mutex>>>, + net: Arc>>, + listener_tx: mpsc::UnboundedSender, + listener_rx: Mutex>, + tasks: Arc>>, +} + +impl SmolTcpStack { + pub async fn new(local_ip: Ipv4Addr) -> anyhow::Result> { + let tasks = Arc::new(std::sync::Mutex::new(JoinSet::new())); + let mut cap = smoltcp::phy::DeviceCapabilities::default(); + cap.max_transmission_unit = 1280; + cap.medium = smoltcp::phy::Medium::Ip; + let (dev, stack_sink, stack_stream) = channel_device::ChannelDevice::new(cap); + + let (ingress_tx, mut ingress_rx) = mpsc::channel::(1000); + tasks.lock().unwrap().spawn(async move { + while let Some(packet) = ingress_rx.recv().await { + tracing::trace!( + target: "easytier_core::gateway::stack", + ?packet, + "receive from peer send to smoltcp packet" + ); + if let Err(err) = stack_sink.send(Ok(packet.payload().to_vec())).await { + tracing::error!( + target: "easytier_core::gateway::stack", + ?err, + "send to smoltcp stack failed" + ); + } + } + tracing::error!( + target: "easytier_core::gateway::stack", + "smoltcp stack sink exited" + ); + }); + + let interface_config = smoltcp::iface::Config::new(smoltcp::wire::HardwareAddress::Ip); + let net = Net::new( + dev, + NetConfig::new( + interface_config, + format!("{local_ip}/24").parse().unwrap(), + vec![format!("{local_ip}").parse().unwrap()], + Some(BufferSize { + tcp_rx_size: 1024 * 16, + tcp_tx_size: 1024 * 16, + ..Default::default() + }), + ), + ); + net.set_any_ip(true); + + let (listener_tx, listener_rx) = mpsc::unbounded_channel(); + Ok(Arc::new(Self { + ingress_tx, + output_rx: Mutex::new(Some(stack_stream)), + net: Arc::new(Mutex::new(Some(net))), + listener_tx, + listener_rx: Mutex::new(listener_rx), + tasks, + })) + } + + pub fn local_port(&self) -> u16 { + 8899 + } + + pub async fn send_ingress(&self, packet: ZCPacket) -> anyhow::Result<()> { + self.ingress_tx + .send(packet) + .await + .map_err(|err| anyhow::anyhow!("send to smoltcp ingress failed: {:?}", err)) + } + + pub async fn take_output_rx(&self) -> anyhow::Result>> { + self.output_rx + .lock() + .await + .take() + .ok_or_else(|| anyhow::anyhow!("smoltcp output receiver already taken")) + } + + pub async fn add_listener(&self) { + let tx = self.listener_tx.clone(); + let locked_net = self.net.lock().await; + let mut tcp = locked_net + .as_ref() + .expect("smoltcp net initialized") + .tcp_bind("0.0.0.0:8899".parse().unwrap()) + .await + .unwrap(); + self.tasks.lock().unwrap().spawn(async move { + let ret = timeout(Duration::from_secs(10), tcp.accept()).await; + if let Ok(accept_ret) = ret { + let _ = + tx.send(accept_ret.map_err(|err| { + anyhow::anyhow!("smol tcp listener accept failed: {:?}", err) + })); + } else { + tracing::error!( + target: "easytier_core::gateway::stack", + "smol tcp listener accept timeout" + ); + } + }); + tracing::info!( + target: "easytier_core::gateway::stack", + "smol tcp listener added" + ); + } + + pub async fn accept(&self) -> anyhow::Result<(SocketAddr, Box)> { + let (stream, src) = self + .listener_rx + .lock() + .await + .recv() + .await + .ok_or_else(|| anyhow::anyhow!("smoltcp listener closed"))??; + tracing::info!( + target: "easytier_core::gateway::stack", + ?src, + "smol tcp listener accepted" + ); + Ok((src, Box::new(stream))) + } +} + +impl Drop for SmolTcpStack { + fn drop(&mut self) { + self.tasks.lock().unwrap().abort_all(); + } +} + +pub fn output_dst_ip(data: &[u8]) -> anyhow::Result { + let ipv4 = smoltcp::wire::Ipv4Packet::new_checked(data) + .map_err(|err| anyhow::anyhow!("smoltcp output is not an IPv4 packet: {:?}", err))?; + Ok(IpAddr::V4(ipv4.dst_addr())) +} diff --git a/easytier/src/gateway/tokio_smoltcp/channel_device.rs b/easytier-core/src/gateway/smoltcp/tokio_smoltcp/channel_device.rs similarity index 100% rename from easytier/src/gateway/tokio_smoltcp/channel_device.rs rename to easytier-core/src/gateway/smoltcp/tokio_smoltcp/channel_device.rs diff --git a/easytier/src/gateway/tokio_smoltcp/device.rs b/easytier-core/src/gateway/smoltcp/tokio_smoltcp/device.rs similarity index 100% rename from easytier/src/gateway/tokio_smoltcp/device.rs rename to easytier-core/src/gateway/smoltcp/tokio_smoltcp/device.rs diff --git a/easytier/src/gateway/tokio_smoltcp/mod.rs b/easytier-core/src/gateway/smoltcp/tokio_smoltcp/mod.rs similarity index 77% rename from easytier/src/gateway/tokio_smoltcp/mod.rs rename to easytier-core/src/gateway/smoltcp/tokio_smoltcp/mod.rs index adb2d2f3..642f6e2a 100644 --- a/easytier/src/gateway/tokio_smoltcp/mod.rs +++ b/easytier-core/src/gateway/smoltcp/tokio_smoltcp/mod.rs @@ -4,7 +4,7 @@ use std::{ io, - net::{IpAddr, SocketAddr}, + net::SocketAddr, sync::{ Arc, atomic::{AtomicU16, Ordering}, @@ -15,9 +15,9 @@ use device::BufferDevice; use reactor::Reactor; pub use smoltcp; use smoltcp::{ - iface::{Config, Interface, Routes}, - time::{Duration, Instant}, - wire::{HardwareAddress, IpAddress, IpCidr}, + iface::{Config, Interface}, + time::Instant, + wire::{IpAddress, IpCidr}, }; pub use socket::{TcpListener, TcpStream, UdpSocket}; pub use socket_allocator::BufferSize; @@ -31,17 +31,6 @@ mod reactor; mod socket; mod socket_allocator; -/// Can be used to create a forever timestamp in neighbor. -// The 60_000 is the same as NeighborCache::ENTRY_LIFETIME. -pub const FOREVER: Instant = - Instant::from_micros_const(i64::MAX - Duration::from_millis(60_000).micros() as i64); - -pub struct Neighbor { - pub protocol_addr: IpAddress, - pub hardware_addr: HardwareAddress, - pub timestamp: Instant, -} - /// A config for a `Net`. /// /// This is used to configure the `Net`. @@ -78,7 +67,7 @@ pub struct Net { ip_addr: IpCidr, from_port: AtomicU16, stopper: Arc, - fut: AbortOnDropHandle>, + _fut: AbortOnDropHandle>, } impl std::fmt::Debug for Net { @@ -130,12 +119,9 @@ impl Net { ip_addr: config.ip_addr, from_port: AtomicU16::new(10001), stopper, - fut: AbortOnDropHandle::new(tokio::spawn(fut)), + _fut: AbortOnDropHandle::new(tokio::spawn(fut)), } } - pub fn get_address(&self) -> IpAddr { - self.ip_addr.address().into() - } pub fn get_port(&self) -> u16 { self.from_port .fetch_update(Ordering::SeqCst, Ordering::SeqCst, |x| { @@ -146,14 +132,14 @@ impl Net { /// Creates a new TcpListener, which will be bound to the specified address. pub async fn tcp_bind(&self, addr: SocketAddr) -> io::Result { let addr = self.set_address(addr); - TcpListener::new(self.reactor.clone(), addr.into()).await + TcpListener::new(self.reactor.clone(), socket_addr_to_endpoint(addr)).await } /// Opens a TCP connection to a remote host. pub async fn tcp_connect(&self, addr: SocketAddr, local_port: u16) -> io::Result { TcpStream::connect( self.reactor.clone(), (self.ip_addr.address(), local_port).into(), - addr.into(), + socket_addr_to_endpoint(addr), ) .await } @@ -161,7 +147,7 @@ impl Net { /// This function will create a new UDP socket and attempt to bind it to the `addr` provided. pub async fn udp_bind(&self, addr: SocketAddr) -> io::Result { let addr = self.set_address(addr); - UdpSocket::new(self.reactor.clone(), addr.into()).await + UdpSocket::new(self.reactor.clone(), socket_addr_to_endpoint(addr)).await } fn set_address(&self, mut addr: SocketAddr) -> SocketAddr { @@ -186,26 +172,12 @@ impl Net { iface.lock(); iface.set_any_ip(any_ip); } +} - /// Get whether AnyIP is enabled. - pub fn any_ip(&self) -> bool { - let iface = self.reactor.iface().clone(); - let iface = iface.lock(); - iface.any_ip() - } - - pub fn routes(&self, f: F) { - let iface = self.reactor.iface().clone(); - let iface = iface.lock(); - let routes = iface.routes(); - f(routes) - } - - pub fn routes_mut(&self, f: F) { - let iface = self.reactor.iface().clone(); - let mut iface = iface.lock(); - let routes = iface.routes_mut(); - f(routes) +fn socket_addr_to_endpoint(addr: SocketAddr) -> smoltcp::wire::IpEndpoint { + match addr { + SocketAddr::V4(addr) => addr.into(), + SocketAddr::V6(addr) => addr.into(), } } diff --git a/easytier/src/gateway/tokio_smoltcp/reactor.rs b/easytier-core/src/gateway/smoltcp/tokio_smoltcp/reactor.rs similarity index 91% rename from easytier/src/gateway/tokio_smoltcp/reactor.rs rename to easytier-core/src/gateway/smoltcp/tokio_smoltcp/reactor.rs index 21a735ef..f7371a0e 100644 --- a/easytier/src/gateway/tokio_smoltcp/reactor.rs +++ b/easytier-core/src/gateway/smoltcp/tokio_smoltcp/reactor.rs @@ -10,7 +10,9 @@ use smoltcp::{ time::{Duration, Instant}, }; use std::{collections::VecDeque, future::Future, io, sync::Arc}; -use tokio::{pin, select, sync::Notify, time::sleep}; +use tokio::{pin, select, sync::Notify}; + +use crate::foundation::time::sleep; pub(crate) type BufferInterface = Arc>; const MAX_BURST_SIZE: usize = 100; @@ -64,7 +66,7 @@ async fn run( timer .as_mut() - .reset(tokio::time::Instant::now() + deadline.into()); + .reset(crate::foundation::time::Instant::now() + deadline.into()); select! { _ = &mut timer => {}, _ = receive(&mut async_iface,&mut recv_buf) => {} @@ -148,7 +150,10 @@ impl Reactor { &self.socket_allocator } pub fn notify(&self) { - self.notify.notify_waiters(); + // The externally driven WASI runtime can park immediately after a + // socket enqueues work. Keep one permit when the reactor has not + // reached its select yet; notify_waiters() would lose that edge. + self.notify.notify_one(); } pub fn iface(&self) -> &BufferInterface { &self.iface diff --git a/easytier/src/gateway/tokio_smoltcp/socket.rs b/easytier-core/src/gateway/smoltcp/tokio_smoltcp/socket.rs similarity index 88% rename from easytier/src/gateway/tokio_smoltcp/socket.rs rename to easytier-core/src/gateway/smoltcp/tokio_smoltcp/socket.rs index f36e4bc3..663ee292 100644 --- a/easytier/src/gateway/tokio_smoltcp/socket.rs +++ b/easytier-core/src/gateway/smoltcp/tokio_smoltcp/socket.rs @@ -1,6 +1,5 @@ use super::{reactor::Reactor, socket_allocator::SocketHandle}; use futures::future::{self, poll_fn}; -use futures::{Stream, ready}; pub use smoltcp::socket::tcp; use smoltcp::socket::udp; use smoltcp::wire::{IpAddress, IpEndpoint}; @@ -59,35 +58,9 @@ impl TcpListener { pub async fn accept(&mut self) -> io::Result<(TcpStream, SocketAddr)> { poll_fn(|cx| self.poll_accept(cx)).await } - pub fn incoming(self) -> Incoming { - Incoming(self) - } pub fn local_addr(&self) -> io::Result { Ok(self.local_addr) } - - pub fn relisten(&self) { - let mut socket = self.reactor.get_socket::(*self.handle); - let local_endpoint = socket.local_endpoint().unwrap(); - socket.abort(); - socket.listen(local_endpoint).unwrap(); - self.reactor.notify(); - } - - pub fn is_listening(&self) -> bool { - let socket = self.reactor.get_socket::(*self.handle); - socket.is_listening() - } -} - -pub struct Incoming(TcpListener); - -impl Stream for Incoming { - type Item = io::Result; - fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { - let (tcp, _) = ready!(self.0.poll_accept(cx))?; - Poll::Ready(Some(Ok(tcp))) - } } fn ep2sa(ep: &IpEndpoint) -> SocketAddr { @@ -99,12 +72,18 @@ fn ep2sa(ep: &IpEndpoint) -> SocketAddr { } } +fn sa2ep(addr: SocketAddr) -> IpEndpoint { + match addr { + SocketAddr::V4(addr) => addr.into(), + SocketAddr::V6(addr) => addr.into(), + } +} + /// A TCP stream between a local and a remote socket. pub struct TcpStream { handle: SocketHandle, reactor: Arc, local_addr: SocketAddr, - peer_addr: SocketAddr, } impl TcpStream { @@ -130,12 +109,10 @@ impl TcpStream { connect_result.map_err(map_err)?; let local_addr = ep2sa(&local_endpoint); - let peer_addr = ep2sa(&remote_endpoint); let tcp = TcpStream { handle, reactor, local_addr, - peer_addr, }; tcp.reactor.notify(); @@ -149,7 +126,9 @@ impl TcpStream { let new_handle = reactor.socket_allocator().new_tcp_socket(); { let mut new_socket = reactor.get_socket::(*new_handle); - new_socket.listen(listener.local_addr).map_err(map_err)?; + new_socket + .listen(sa2ep(listener.local_addr)) + .map_err(map_err)?; } let (peer_addr, local_addr) = { let socket = reactor.get_socket::(*listener.handle); @@ -168,7 +147,6 @@ impl TcpStream { handle: replace(&mut listener.handle, new_handle), reactor, local_addr, - peer_addr, }, peer_addr, )) @@ -177,9 +155,6 @@ impl TcpStream { pub fn local_addr(&self) -> io::Result { Ok(self.local_addr) } - pub fn peer_addr(&self) -> io::Result { - Ok(self.peer_addr) - } pub fn poll_connected(&self, cx: &Context<'_>) -> Poll> { let mut socket = self.reactor.get_socket::(*self.handle); if socket.state() == tcp::State::Established { @@ -292,7 +267,7 @@ impl UdpSocket { target: SocketAddr, ) -> Poll> { let mut socket = self.reactor.get_socket::(*self.handle); - let target_ip: IpEndpoint = target.into(); + let target_ip: IpEndpoint = sa2ep(target); match socket.send_slice(buf, target_ip) { // the buffer is full @@ -336,6 +311,37 @@ impl UdpSocket { pub async fn recv_from(&self, buf: &mut [u8]) -> io::Result<(usize, SocketAddr)> { poll_fn(|cx| self.poll_recv_from(cx, buf)).await } + + pub fn poll_recv_from_limited( + &self, + cx: &Context<'_>, + max_len: usize, + ) -> Poll, SocketAddr, bool)>> { + let mut socket = self.reactor.get_socket::(*self.handle); + + match socket.recv() { + Err(udp::RecvError::Exhausted) => {} + Err(error) => return Poll::Ready(Err(map_err(error))), + Ok((payload, metadata)) => { + let copy_len = payload.len().min(max_len); + let truncated = copy_len < payload.len(); + let data = payload[..copy_len].to_vec(); + self.reactor.notify(); + return Poll::Ready(Ok((data, ep2sa(&metadata.endpoint), truncated))); + } + } + + socket.register_recv_waker(cx.waker()); + Poll::Pending + } + + pub async fn recv_from_limited( + &self, + max_len: usize, + ) -> io::Result<(Vec, SocketAddr, bool)> { + poll_fn(|cx| self.poll_recv_from_limited(cx, max_len)).await + } + pub fn local_addr(&self) -> io::Result { Ok(self.local_addr) } diff --git a/easytier/src/gateway/tokio_smoltcp/socket_allocator.rs b/easytier-core/src/gateway/smoltcp/tokio_smoltcp/socket_allocator.rs similarity index 100% rename from easytier/src/gateway/tokio_smoltcp/socket_allocator.rs rename to easytier-core/src/gateway/smoltcp/tokio_smoltcp/socket_allocator.rs diff --git a/easytier-core/src/gateway/socks5/adapter.rs b/easytier-core/src/gateway/socks5/adapter.rs new file mode 100644 index 00000000..9b578a0b --- /dev/null +++ b/easytier-core/src/gateway/socks5/adapter.rs @@ -0,0 +1,215 @@ +//! Host-listener adapter that translates SOCKS5 sessions into data-plane calls. + +use std::{ + net::SocketAddr, + sync::{ + Arc, + atomic::{AtomicBool, Ordering}, + }, + time::Duration, +}; + +use tokio::{sync::Mutex, task::JoinSet}; + +use crate::{ + config::runtime::CoreRuntimeConfigStore, + foundation::task::reap_joinset_background, + gateway::dataplane::{ + DataPlaneConsumerLease, DataPlaneError, DataPlaneErrorKind, DataPlaneRuntime, + DataPlaneTcpConnectOptions, DataPlaneTcpStream, + }, + host::dns::DnsResolver, + socket::{ + SocketContext, + tcp::{ + TcpListenOptions, TcpSocketPurpose, VirtualTcpListener, VirtualTcpListenerFactory, + VirtualTcpSocket, VirtualTcpSocketFactory, + }, + udp::VirtualUdpSocketFactory, + }, +}; + +use super::{ + AcceptAuthentication, AsyncTcpConnector, Config, HostSocks5ServerRuntime, Result, Socks5Socket, + SocksError, codec::ReplyError, +}; + +struct Socks5DataPlaneConnector +where + H: VirtualTcpSocketFactory + VirtualTcpListenerFactory + VirtualUdpSocketFactory, +{ + data_plane: Arc>, + source_addr: SocketAddr, +} + +#[async_trait::async_trait] +impl AsyncTcpConnector for Socks5DataPlaneConnector +where + H: VirtualTcpSocketFactory + VirtualTcpListenerFactory + VirtualUdpSocketFactory, +{ + type S = DataPlaneTcpStream; + + async fn tcp_connect(&self, addr: SocketAddr, timeout_s: u64) -> Result { + self.data_plane + .connect_tcp( + addr, + DataPlaneTcpConnectOptions::gateway( + Duration::from_secs(timeout_s), + TcpSocketPurpose::Socks5, + self.source_addr, + ), + ) + .await + .map_err(map_data_plane_error) + } +} + +fn map_data_plane_error(error: DataPlaneError) -> SocksError { + match error.kind() { + DataPlaneErrorKind::DeadlineExceeded => ReplyError::ConnectionTimeout.into(), + DataPlaneErrorKind::ConnectionRefused => ReplyError::ConnectionRefused.into(), + DataPlaneErrorKind::NoOverlayRoute | DataPlaneErrorKind::PathNotReady => { + ReplyError::NetworkUnreachable.into() + } + DataPlaneErrorKind::AddressFamilyUnsupported => ReplyError::AddressTypeNotSupported.into(), + DataPlaneErrorKind::Cancelled + | DataPlaneErrorKind::InstanceStopped + | DataPlaneErrorKind::HandleClosed + | DataPlaneErrorKind::NetworkChanged => ReplyError::ConnectionNotAllowed.into(), + DataPlaneErrorKind::AddressInUse + | DataPlaneErrorKind::ResourceLimit + | DataPlaneErrorKind::BufferTooSmall + | DataPlaneErrorKind::Io => SocksError::Other(anyhow::Error::new(error)), + } +} + +async fn handle_socks5_stream( + stream: S, + data_plane: Arc>, + source_addr: SocketAddr, + command_runtime: Arc>, +) where + H: VirtualTcpSocketFactory + VirtualTcpListenerFactory + VirtualUdpSocketFactory, + S: VirtualTcpSocket, +{ + let mut config = Config::::default(); + config.set_request_timeout(10); + config.set_skip_auth(false); + config.set_allow_no_auth(true); + + let connector = Socks5DataPlaneConnector { + data_plane, + source_addr, + }; + let socket = Socks5Socket::new(stream, Arc::new(config), connector, command_runtime); + if let Err(error) = socket.upgrade_to_socks5().await { + tracing::error!(?error, "SOCKS5 session failed"); + } +} + +pub(crate) struct Socks5GatewayAdapter +where + H: VirtualTcpSocketFactory + VirtualTcpListenerFactory + VirtualUdpSocketFactory, +{ + operation: Mutex<()>, + started: AtomicBool, + runtime_config: CoreRuntimeConfigStore, + data_plane: Arc>, + host: Arc, + socket_context: SocketContext, + command_runtime: Arc>, + consumer_lease: Mutex>, + tasks: Arc>>, +} + +impl Socks5GatewayAdapter +where + H: VirtualTcpSocketFactory + VirtualTcpListenerFactory + VirtualUdpSocketFactory, +{ + pub(crate) fn new( + runtime_config: CoreRuntimeConfigStore, + data_plane: Arc>, + host: Arc, + dns: Arc, + socket_context: SocketContext, + ) -> Arc { + Arc::new(Self { + operation: Mutex::new(()), + started: AtomicBool::new(false), + runtime_config, + data_plane, + host: host.clone(), + socket_context: socket_context.clone(), + command_runtime: Arc::new(HostSocks5ServerRuntime::new(host, dns, socket_context)), + consumer_lease: Mutex::new(None), + tasks: Arc::new(std::sync::Mutex::new(JoinSet::new())), + }) + } + + async fn start_inner(&self) -> anyhow::Result<()> { + let Some(bind_addr) = self.runtime_config.snapshot().services.gateway.socks5_bind else { + return Ok(()); + }; + let options = TcpListenOptions::socks5(bind_addr); + let bind = options + .bind + .clone() + .with_context(self.socket_context.clone()); + let listener = self.host.bind_tcp(options.with_bind(bind)).await?; + let consumer_lease = self.data_plane.acquire_consumer_lease()?; + + self.tasks.lock().unwrap().spawn(reap_joinset_background( + self.tasks.clone(), + "SOCKS5 gateway adapter", + )); + let data_plane = self.data_plane.clone(); + let command_runtime = self.command_runtime.clone(); + let session_tasks = self.tasks.clone(); + self.tasks.lock().unwrap().spawn(async move { + loop { + match listener.accept().await { + Ok((socket, source_addr)) => { + tracing::info!(?source_addr, "accepted a SOCKS5 connection"); + session_tasks.lock().unwrap().spawn(handle_socks5_stream( + socket, + data_plane.clone(), + source_addr, + command_runtime.clone(), + )); + } + Err(error) => tracing::error!(?error, "SOCKS5 accept failed"), + } + } + }); + self.consumer_lease.lock().await.replace(consumer_lease); + Ok(()) + } + + pub(crate) async fn start(&self) -> anyhow::Result<()> { + let _operation = self.operation.lock().await; + if self.started.load(Ordering::Acquire) { + return Ok(()); + } + if let Err(error) = self.start_inner().await { + self.stop_inner().await; + return Err(error); + } + self.started.store(true, Ordering::Release); + Ok(()) + } + + async fn stop_inner(&self) { + self.started.store(false, Ordering::Release); + self.consumer_lease.lock().await.take(); + let mut tasks = { + let mut tasks = self.tasks.lock().unwrap(); + std::mem::replace(&mut *tasks, JoinSet::new()) + }; + tasks.shutdown().await; + } + + pub(crate) async fn stop(&self) { + let _operation = self.operation.lock().await; + self.stop_inner().await; + } +} diff --git a/easytier/src/gateway/fast_socks5/mod.rs b/easytier-core/src/gateway/socks5/codec.rs similarity index 56% rename from easytier/src/gateway/fast_socks5/mod.rs rename to easytier-core/src/gateway/socks5/codec.rs index 424f0816..8c836cb3 100644 --- a/easytier/src/gateway/fast_socks5/mod.rs +++ b/easytier-core/src/gateway/socks5/codec.rs @@ -1,67 +1,27 @@ -//! Fast SOCKS5 client/server implementation written in Rust async/.await (with tokio). +//! Portable SOCKS5 wire types and codecs: protocol constants, error and +//! address types, and the UDP request header encode/decode helpers. //! -//! This library is maintained by [anyip.io](https://anyip.io/) a residential and mobile socks5 proxy provider. -//! -//! ## Features -//! -//! - An `async`/`.await` [SOCKS5](https://tools.ietf.org/html/rfc1928) implementation. -//! - An `async`/`.await` [SOCKS4 Client](https://www.openssh.com/txt/socks4.protocol) implementation. -//! - An `async`/`.await` [SOCKS4a Client](https://www.openssh.com/txt/socks4a.protocol) implementation. -//! - No **unsafe** code -//! - Built on-top of `tokio` library -//! - Ultra lightweight and scalable -//! - No system dependencies -//! - Cross-platform -//! - Authentication methods: -//! - No-Auth method -//! - Username/Password auth method -//! - Custom auth methods can be implemented via the Authentication Trait -//! - Credentials returned on authentication success -//! - All SOCKS5 RFC errors (replies) should be mapped -//! - `AsyncRead + AsyncWrite` traits are implemented on Socks5Stream & Socks5Socket -//! - `IPv4`, `IPv6`, and `Domains` types are supported -//! - Config helper for Socks5Server -//! - Helpers to run a Socks5Server à la *"std's TcpStream"* via `incoming.next().await` -//! - Examples come with real cases commands scenarios -//! - Can disable `DNS resolving` -//! - Can skip the authentication/handshake process, which will directly handle command's request (useful to save useless round-trips in a current authenticated environment) -//! - Can disable command execution (useful if you just want to forward the request to a different server) -//! -//! -//! ## Install -//! -//! Open in [crates.io](https://crates.io/crates/fast-socks5). -//! -//! -//! ## Examples -//! -//! Please check [`examples`](https://github.com/dizda/fast-socks5/tree/master/examples) directory. +//! DNS resolution, socket creation, listeners, and concrete command execution +//! are deliberately supplied by host adapters outside this module. -#![forbid(unsafe_code)] - -pub mod server; -pub mod util; +mod target_addr; use anyhow::Context; use std::fmt; use std::io; +pub(crate) use target_addr::TargetAddr; +pub(super) use target_addr::{AddrError, ToTargetAddr, read_address}; use thiserror::Error; -use util::target_addr::TargetAddr; -use util::target_addr::ToTargetAddr; -use util::target_addr::read_address; use tokio::io::AsyncReadExt; use tracing::error; -use crate::read_exact; - #[rustfmt::skip] -pub mod consts { +pub(super) mod consts { pub const SOCKS5_VERSION: u8 = 0x05; pub const SOCKS5_AUTH_METHOD_NONE: u8 = 0x00; - pub const SOCKS5_AUTH_METHOD_GSSAPI: u8 = 0x01; pub const SOCKS5_AUTH_METHOD_PASSWORD: u8 = 0x02; pub const SOCKS5_AUTH_METHOD_NOT_ACCEPTABLE: u8 = 0xff; @@ -77,7 +37,6 @@ pub mod consts { pub const SOCKS5_REPLY_GENERAL_FAILURE: u8 = 0x01; pub const SOCKS5_REPLY_CONNECTION_NOT_ALLOWED: u8 = 0x02; pub const SOCKS5_REPLY_NETWORK_UNREACHABLE: u8 = 0x03; - pub const SOCKS5_REPLY_HOST_UNREACHABLE: u8 = 0x04; pub const SOCKS5_REPLY_CONNECTION_REFUSED: u8 = 0x05; pub const SOCKS5_REPLY_TTL_EXPIRED: u8 = 0x06; pub const SOCKS5_REPLY_COMMAND_NOT_SUPPORTED: u8 = 0x07; @@ -85,27 +44,16 @@ pub mod consts { } #[derive(Debug, PartialEq)] -pub enum Socks5Command { +pub(super) enum Socks5Command { TCPConnect, TCPBind, UDPAssociate, } -#[allow(dead_code)] impl Socks5Command { #[inline] #[rustfmt::skip] - fn as_u8(&self) -> u8 { - match self { - Socks5Command::TCPConnect => consts::SOCKS5_CMD_TCP_CONNECT, - Socks5Command::TCPBind => consts::SOCKS5_CMD_TCP_BIND, - Socks5Command::UDPAssociate => consts::SOCKS5_CMD_UDP_ASSOCIATE, - } - } - - #[inline] - #[rustfmt::skip] - fn from_u8(code: u8) -> Option { + pub fn from_u8(code: u8) -> Option { match code { consts::SOCKS5_CMD_TCP_CONNECT => Some(Socks5Command::TCPConnect), consts::SOCKS5_CMD_TCP_BIND => Some(Socks5Command::TCPBind), @@ -116,7 +64,7 @@ impl Socks5Command { } #[derive(Debug, PartialEq)] -pub enum AuthenticationMethod { +pub(super) enum AuthenticationMethod { None, Password { username: String, password: String }, } @@ -124,17 +72,7 @@ pub enum AuthenticationMethod { impl AuthenticationMethod { #[inline] #[rustfmt::skip] - fn as_u8(&self) -> u8 { - match self { - AuthenticationMethod::None => consts::SOCKS5_AUTH_METHOD_NONE, - AuthenticationMethod::Password {..} => - consts::SOCKS5_AUTH_METHOD_PASSWORD - } - } - - #[inline] - #[rustfmt::skip] - fn from_u8(code: u8) -> Option { + pub fn from_u8(code: u8) -> Option { match code { consts::SOCKS5_AUTH_METHOD_NONE => Some(AuthenticationMethod::None), consts::SOCKS5_AUTH_METHOD_PASSWORD => Some(AuthenticationMethod::Password { username: "test".to_string(), password: "test".to_string()}), @@ -165,13 +103,9 @@ impl fmt::Display for AuthenticationMethod { //} #[derive(Error, Debug)] -pub enum SocksError { +pub(crate) enum SocksError { #[error("i/o error: {0}")] Io(#[from] io::Error), - #[error("the data for key `{0}` is not available")] - Redaction(String), - #[error("invalid header (expected {expected:?}, found {found:?})")] - InvalidHeader { expected: String, found: String }, #[error("Auth method unacceptable `{0:?}`.")] AuthMethodUnacceptable(Vec), @@ -187,19 +121,16 @@ pub enum SocksError { #[error("Error with reply: {0}.")] ReplyError(#[from] ReplyError), - #[error("Argument input error: `{0}`.")] - ArgumentInputError(&'static str), - // #[error("Other: `{0}`.")] #[error(transparent)] Other(#[from] anyhow::Error), } -pub type Result = core::result::Result; +pub(crate) type Result = core::result::Result; /// SOCKS5 reply code #[derive(Error, Debug, Copy, Clone)] -pub enum ReplyError { +pub(crate) enum ReplyError { #[error("Succeeded")] Succeeded, #[error("General failure")] @@ -208,14 +139,10 @@ pub enum ReplyError { ConnectionNotAllowed, #[error("Network unreachable")] NetworkUnreachable, - #[error("Host unreachable")] - HostUnreachable, #[error("Connection refused")] ConnectionRefused, #[error("Connection timeout")] ConnectionTimeout, - #[error("TTL expired")] - TtlExpired, #[error("Command not supported")] CommandNotSupported, #[error("Address type not supported")] @@ -232,33 +159,13 @@ impl ReplyError { ReplyError::GeneralFailure => consts::SOCKS5_REPLY_GENERAL_FAILURE, ReplyError::ConnectionNotAllowed => consts::SOCKS5_REPLY_CONNECTION_NOT_ALLOWED, ReplyError::NetworkUnreachable => consts::SOCKS5_REPLY_NETWORK_UNREACHABLE, - ReplyError::HostUnreachable => consts::SOCKS5_REPLY_HOST_UNREACHABLE, ReplyError::ConnectionRefused => consts::SOCKS5_REPLY_CONNECTION_REFUSED, ReplyError::ConnectionTimeout => consts::SOCKS5_REPLY_TTL_EXPIRED, - ReplyError::TtlExpired => consts::SOCKS5_REPLY_TTL_EXPIRED, ReplyError::CommandNotSupported => consts::SOCKS5_REPLY_COMMAND_NOT_SUPPORTED, ReplyError::AddressTypeNotSupported => consts::SOCKS5_REPLY_ADDRESS_TYPE_NOT_SUPPORTED, // ReplyError::OtherReply(c) => c, } } - - #[inline] - #[rustfmt::skip] - pub fn from_u8(code: u8) -> ReplyError { - match code { - consts::SOCKS5_REPLY_SUCCEEDED => ReplyError::Succeeded, - consts::SOCKS5_REPLY_GENERAL_FAILURE => ReplyError::GeneralFailure, - consts::SOCKS5_REPLY_CONNECTION_NOT_ALLOWED => ReplyError::ConnectionNotAllowed, - consts::SOCKS5_REPLY_NETWORK_UNREACHABLE => ReplyError::NetworkUnreachable, - consts::SOCKS5_REPLY_HOST_UNREACHABLE => ReplyError::HostUnreachable, - consts::SOCKS5_REPLY_CONNECTION_REFUSED => ReplyError::ConnectionRefused, - consts::SOCKS5_REPLY_TTL_EXPIRED => ReplyError::TtlExpired, - consts::SOCKS5_REPLY_COMMAND_NOT_SUPPORTED => ReplyError::CommandNotSupported, - consts::SOCKS5_REPLY_ADDRESS_TYPE_NOT_SUPPORTED => ReplyError::AddressTypeNotSupported, -// _ => ReplyError::OtherReply(code), - _ => unreachable!("ReplyError code unsupported."), - } - } } /// Generate UDP header @@ -283,7 +190,7 @@ impl ReplyError { /// o DST.PORT desired destination port /// o DATA user data /// ``` -pub fn new_udp_header(target_addr: T) -> Result> { +pub(super) fn new_udp_header(target_addr: T) -> Result> { let mut header = vec![ 0, 0, // RSV 0, // FRAG @@ -294,14 +201,21 @@ pub fn new_udp_header(target_addr: T) -> Result> { } /// Parse data from UDP client on raw buffer, return (frag, target_addr, payload). -pub async fn parse_udp_request(mut req: &[u8]) -> Result<(u8, TargetAddr, &[u8])> { - let rsv = read_exact!(req, [0u8; 2]).context("Malformed request")?; +pub(super) async fn parse_udp_request(mut req: &[u8]) -> Result<(u8, TargetAddr, &[u8])> { + let mut rsv = [0u8; 2]; + req.read_exact(&mut rsv) + .await + .context("Malformed request")?; if !rsv.eq(&[0u8; 2]) { return Err(ReplyError::GeneralFailure.into()); } - let [frag, atyp] = read_exact!(req, [0u8; 2]).context("Malformed request")?; + let mut frag_and_type = [0u8; 2]; + req.read_exact(&mut frag_and_type) + .await + .context("Malformed request")?; + let [frag, atyp] = frag_and_type; let target_addr = read_address(&mut req, atyp).await.map_err(|e| { // print explicit error @@ -312,3 +226,49 @@ pub async fn parse_udp_request(mut req: &[u8]) -> Result<(u8, TargetAddr, &[u8]) Ok((frag, target_addr, req)) } + +#[cfg(test)] +mod tests { + use std::net::SocketAddr; + + use super::*; + + #[tokio::test] + async fn udp_ipv4_header_round_trips() { + let destination: SocketAddr = "10.42.0.7:5353".parse().unwrap(); + let mut packet = new_udp_header(destination).unwrap(); + packet.extend_from_slice(b"payload"); + + let (fragment, parsed_destination, payload) = parse_udp_request(&packet).await.unwrap(); + + assert_eq!(fragment, 0); + assert_eq!(parsed_destination, TargetAddr::Ip(destination)); + assert_eq!(payload, b"payload"); + } + + #[tokio::test] + async fn udp_domain_header_round_trips() { + let mut packet = new_udp_header(("peer.example", 443)).unwrap(); + packet.extend_from_slice(b"hello"); + + let (_, parsed_destination, payload) = parse_udp_request(&packet).await.unwrap(); + + assert_eq!( + parsed_destination, + TargetAddr::Domain("peer.example".to_string(), 443) + ); + assert_eq!(payload, b"hello"); + } + + #[tokio::test] + async fn udp_header_rejects_nonzero_reserved_field() { + let err = parse_udp_request(&[1, 0, 0, consts::SOCKS5_ADDR_TYPE_IPV4]) + .await + .expect_err("nonzero reserved field must fail"); + + assert!(matches!( + err, + SocksError::ReplyError(ReplyError::GeneralFailure) + )); + } +} diff --git a/easytier/src/gateway/fast_socks5/util/target_addr.rs b/easytier-core/src/gateway/socks5/codec/target_addr.rs similarity index 75% rename from easytier/src/gateway/fast_socks5/util/target_addr.rs rename to easytier-core/src/gateway/socks5/codec/target_addr.rs index 76833a2c..3b6e073b 100644 --- a/easytier/src/gateway/fast_socks5/util/target_addr.rs +++ b/easytier-core/src/gateway/socks5/codec/target_addr.rs @@ -1,7 +1,5 @@ -use crate::gateway::fast_socks5::SocksError; -use crate::gateway::fast_socks5::consts; -use crate::gateway::fast_socks5::consts::SOCKS5_ADDR_TYPE_IPV4; -use crate::read_exact; +use super::{SocksError, consts}; +use consts::SOCKS5_ADDR_TYPE_IPV4; use anyhow::Context; use std::fmt; @@ -10,7 +8,6 @@ use std::net::{Ipv4Addr, Ipv6Addr, SocketAddr, SocketAddrV4, SocketAddrV6}; use std::vec::IntoIter; use thiserror::Error; use tokio::io::{AsyncRead, AsyncReadExt}; -use tokio::net::lookup_host; use tracing::debug; @@ -50,35 +47,6 @@ pub enum TargetAddr { } impl TargetAddr { - pub async fn resolve_dns(self) -> anyhow::Result { - match self { - TargetAddr::Ip(ip) => Ok(TargetAddr::Ip(ip)), - TargetAddr::Domain(domain, port) => { - debug!("Attempt to DNS resolve the domain {}...", &domain); - - let socket_addr = lookup_host((&domain[..], port)) - .await - .context(AddrError::DNSResolutionFailed)? - .next() - .ok_or(AddrError::Custom( - "Can't fetch DNS to the domain.".to_string(), - ))?; - debug!("domain name resolved to {}", socket_addr); - - // has been converted to an ip - Ok(TargetAddr::Ip(socket_addr)) - } - } - } - - pub fn is_ip(&self) -> bool { - matches!(self, TargetAddr::Ip(_)) - } - - pub fn is_domain(&self) -> bool { - !self.is_ip() - } - pub fn to_be_bytes(&self) -> anyhow::Result> { let mut buf = vec![]; match self { @@ -114,16 +82,14 @@ impl TargetAddr { } } -// async-std ToSocketAddrs doesn't supports external trait implementation -// @see https://github.com/async-rs/async-std/issues/539 impl std::net::ToSocketAddrs for TargetAddr { type Iter = IntoIter; - fn to_socket_addrs(&self) -> io::Result> { - match *self { - TargetAddr::Ip(addr) => Ok(vec![addr].into_iter()), + fn to_socket_addrs(&self) -> io::Result { + match self { + TargetAddr::Ip(address) => Ok(vec![*address].into_iter()), TargetAddr::Domain(_, _) => Err(io::Error::other( - "Domain name has to be explicitly resolved, please use TargetAddr::resolve_dns().", + "domain name has to be explicitly resolved", )), } } @@ -205,16 +171,33 @@ pub async fn read_address( let addr = match atyp { consts::SOCKS5_ADDR_TYPE_IPV4 => { debug!("Address type `IPv4`"); - Addr::V4(read_exact!(stream, [0u8; 4]).context(AddrError::IPv4Unreadable)?) + let mut bytes = [0u8; 4]; + stream + .read_exact(&mut bytes) + .await + .context(AddrError::IPv4Unreadable)?; + Addr::V4(bytes) } consts::SOCKS5_ADDR_TYPE_IPV6 => { debug!("Address type `IPv6`"); - Addr::V6(read_exact!(stream, [0u8; 16]).context(AddrError::IPv6Unreadable)?) + let mut bytes = [0u8; 16]; + stream + .read_exact(&mut bytes) + .await + .context(AddrError::IPv6Unreadable)?; + Addr::V6(bytes) } consts::SOCKS5_ADDR_TYPE_DOMAIN_NAME => { debug!("Address type `domain`"); - let len = read_exact!(stream, [0]).context(AddrError::DomainLenUnreadable)?[0]; - let domain = read_exact!(stream, vec![0u8; len as usize]) + let mut len = [0]; + stream + .read_exact(&mut len) + .await + .context(AddrError::DomainLenUnreadable)?; + let mut domain = vec![0u8; len[0] as usize]; + stream + .read_exact(&mut domain) + .await .context(AddrError::DomainContentUnreadable)?; // make sure the bytes are correct utf8 string let domain = String::from_utf8(domain).context(AddrError::Utf8)?; @@ -225,7 +208,11 @@ pub async fn read_address( }; // Find port number - let port = read_exact!(stream, [0u8; 2]).context(AddrError::PortNumberUnreadable)?; + let mut port = [0u8; 2]; + stream + .read_exact(&mut port) + .await + .context(AddrError::PortNumberUnreadable)?; // Convert (u8 * 2) into u16 let port = (port[0] as u16) << 8 | port[1] as u16; @@ -238,3 +225,24 @@ pub async fn read_address( Ok(addr) } + +#[cfg(test)] +mod tests { + use super::*; + use std::net::ToSocketAddrs; + + #[test] + fn target_addr_to_socket_addrs_keeps_ip_only_contract() { + let address: SocketAddr = "192.0.2.8:443".parse().unwrap(); + let resolved = TargetAddr::Ip(address) + .to_socket_addrs() + .unwrap() + .collect::>(); + assert_eq!(resolved, vec![address]); + + let error = TargetAddr::Domain("peer.example".into(), 443) + .to_socket_addrs() + .unwrap_err(); + assert_eq!(error.kind(), io::ErrorKind::Other); + } +} diff --git a/easytier-core/src/gateway/socks5/host.rs b/easytier-core/src/gateway/socks5/host.rs new file mode 100644 index 00000000..f62fe12d --- /dev/null +++ b/easytier-core/src/gateway/socks5/host.rs @@ -0,0 +1,366 @@ +use std::{ + io, + net::{IpAddr, Ipv6Addr, SocketAddr, SocketAddrV6}, + sync::{Arc, Mutex}, +}; + +use crate::{ + host::dns::{DnsQuery, DnsResolver}, + socket::{ + IpVersion, SocketContext, + udp::{UdpBindOptions, VirtualUdpSocket, VirtualUdpSocketFactory}, + }, +}; + +use super::{ + codec::{AddrError, Result, SocksError, TargetAddr, new_udp_header, parse_udp_request}, + server::{Socks5ServerRuntime, Socks5UdpAssociation}, +}; + +/// Portable SOCKS command runtime backed exclusively by host capabilities. +pub(crate) struct HostSocks5ServerRuntime +where + H: VirtualUdpSocketFactory, +{ + host: Arc, + dns: Arc, + socket_context: SocketContext, +} + +impl HostSocks5ServerRuntime +where + H: VirtualUdpSocketFactory, +{ + pub fn new(host: Arc, dns: Arc, socket_context: SocketContext) -> Self { + Self { + host, + dns, + socket_context, + } + } + + async fn resolve_target(&self, target: TargetAddr) -> Result { + match target { + TargetAddr::Ip(address) => Ok(address), + TargetAddr::Domain(domain, port) => { + tracing::debug!(%domain, "attempting SOCKS DNS resolution"); + let address = self + .dns + .resolve(DnsQuery::new(domain.clone(), self.socket_context.clone())) + .await + .map_err(|error| { + SocksError::Other(error.context(AddrError::DNSResolutionFailed)) + })? + .into_iter() + .next() + .ok_or_else(|| { + SocksError::Other(anyhow::Error::new(AddrError::Custom(format!( + "cannot resolve SOCKS domain {domain}" + )))) + })?; + let address = SocketAddr::new(address, port); + tracing::debug!(%address, "SOCKS domain resolved"); + Ok(address) + } + } + } + + fn udp_bind_options(&self) -> UdpBindOptions { + UdpBindOptions::socks5() + .with_context(self.socket_context.clone().with_ip_version(IpVersion::V6)) + .with_local_addr(Some(SocketAddr::V6(SocketAddrV6::new( + Ipv6Addr::UNSPECIFIED, + 0, + 0, + 0, + )))) + } +} + +#[async_trait::async_trait] +impl Socks5ServerRuntime for HostSocks5ServerRuntime +where + H: VirtualUdpSocketFactory, +{ + async fn resolve_dns(&self, target_addr: TargetAddr) -> Result { + self.resolve_target(target_addr).await.map(TargetAddr::Ip) + } + + async fn bind_udp_association(&self) -> Result> { + let inbound = self + .host + .bind_udp(self.udp_bind_options()) + .await + .map_err(SocksError::Other)?; + Ok(Box::new(HostSocks5UdpAssociation { + inbound, + host: self.host.clone(), + socket_context: self.socket_context.clone(), + client: Mutex::new(None), + })) + } +} + +struct HostSocks5UdpAssociation +where + H: VirtualUdpSocketFactory, +{ + inbound: Arc, + host: Arc, + socket_context: SocketContext, + client: Mutex>, +} + +#[async_trait::async_trait] +impl Socks5UdpAssociation for HostSocks5UdpAssociation +where + H: VirtualUdpSocketFactory, +{ + fn local_addr(&self) -> io::Result { + self.inbound.local_addr() + } + + async fn transfer(self: Box) -> Result<()> { + let outbound = self + .host + .bind_udp( + UdpBindOptions::socks5() + .with_context(self.socket_context.clone().with_ip_version(IpVersion::V6)) + .with_local_addr(Some(SocketAddr::V6(SocketAddrV6::new( + Ipv6Addr::UNSPECIFIED, + 0, + 0, + 0, + )))), + ) + .await + .map_err(SocksError::Other)?; + tokio::try_join!( + transfer_requests(self.as_ref(), outbound.as_ref()), + transfer_responses(self.as_ref(), outbound.as_ref()), + )?; + Ok(()) + } +} + +async fn transfer_requests( + association: &HostSocks5UdpAssociation, + outbound: &H::Socket, +) -> Result<()> +where + H: VirtualUdpSocketFactory, +{ + let mut buffer = vec![0u8; 0x10000]; + loop { + let (size, client) = association.inbound.recv_from(&mut buffer).await?; + { + let mut pinned = association.client.lock().unwrap(); + match *pinned { + None => *pinned = Some(client), + Some(current) if current == client => {} + Some(_) => continue, + } + } + + let (fragment, target, data) = parse_udp_request(&buffer[..size]).await?; + if fragment != 0 { + tracing::debug!(fragment, "discarding fragmented SOCKS UDP request"); + return Ok(()); + } + + let target = match target { + TargetAddr::Ip(address) => address, + TargetAddr::Domain(_, _) => { + return Err(io::Error::other( + "SOCKS UDP domain targets must be explicitly resolved by the client", + ) + .into()); + } + }; + outbound.send_to(data, ipv4_mapped_addr(target)).await?; + } +} + +async fn transfer_responses( + association: &HostSocks5UdpAssociation, + outbound: &H::Socket, +) -> Result<()> +where + H: VirtualUdpSocketFactory, +{ + let mut buffer = vec![0u8; 0x10000]; + loop { + let (size, remote) = outbound.recv_from(&mut buffer).await?; + let client = *association.client.lock().unwrap(); + let Some(client) = client else { + continue; + }; + let mut data = new_udp_header(ipv4_mapped_addr(remote))?; + data.extend_from_slice(&buffer[..size]); + association.inbound.send_to(&data, client).await?; + } +} + +fn ipv4_mapped_addr(mut address: SocketAddr) -> SocketAddr { + if let IpAddr::V4(ipv4) = address.ip() { + address.set_ip(IpAddr::V6(ipv4.to_ipv6_mapped())); + } + address +} + +#[cfg(test)] +mod tests { + use super::*; + use std::collections::VecDeque; + use std::sync::atomic::{AtomicUsize, Ordering}; + + struct StaticDns; + + #[async_trait::async_trait] + impl DnsResolver for StaticDns { + async fn resolve(&self, query: DnsQuery) -> anyhow::Result> { + assert_eq!(query.host, "peer.example"); + Ok(vec!["192.0.2.8".parse().unwrap()]) + } + } + + #[derive(Default)] + struct MockSocket { + receives: Mutex, SocketAddr)>>, + sends: Mutex, SocketAddr)>>, + } + + #[async_trait::async_trait] + impl VirtualUdpSocket for MockSocket { + fn local_addr(&self) -> io::Result { + Ok("[::]:42000".parse().unwrap()) + } + + async fn send_to(&self, data: &[u8], addr: SocketAddr) -> io::Result { + self.sends.lock().unwrap().push((data.to_vec(), addr)); + Ok(data.len()) + } + + async fn recv_from(&self, buf: &mut [u8]) -> io::Result<(usize, SocketAddr)> { + let (data, addr) = self.receives.lock().unwrap().pop_front().unwrap(); + buf[..data.len()].copy_from_slice(&data); + Ok((data.len(), addr)) + } + } + + #[derive(Default)] + struct MockHost { + binds: AtomicUsize, + options: Mutex>, + } + + #[async_trait::async_trait] + impl VirtualUdpSocketFactory for MockHost { + type Socket = MockSocket; + + async fn bind_udp(&self, options: UdpBindOptions) -> anyhow::Result> { + self.binds.fetch_add(1, Ordering::Relaxed); + self.options.lock().unwrap().push(options); + Ok(Arc::new(MockSocket::default())) + } + } + + #[tokio::test] + async fn resolves_and_binds_only_through_host_capabilities() { + let host = Arc::new(MockHost::default()); + let context = SocketContext::default().with_socket_mark(Some(7)); + let runtime = + HostSocks5ServerRuntime::new(host.clone(), Arc::new(StaticDns), context.clone()); + + assert_eq!( + runtime + .resolve_dns(TargetAddr::Domain("peer.example".into(), 443)) + .await + .unwrap(), + TargetAddr::Ip("192.0.2.8:443".parse().unwrap()) + ); + let association = runtime.bind_udp_association().await.unwrap(); + assert_eq!(association.local_addr().unwrap().port(), 42000); + assert_eq!(host.binds.load(Ordering::Relaxed), 1); + let options = host.options.lock().unwrap(); + assert_eq!( + options[0].purpose, + crate::socket::udp::UdpSocketPurpose::Socks5 + ); + assert_eq!(options[0].context.socket_mark, context.socket_mark); + assert_eq!(options[0].context.ip_version, IpVersion::V6); + } + + #[tokio::test] + async fn udp_requests_pin_the_first_client_and_map_ipv4_targets() { + let client: SocketAddr = "192.0.2.1:1000".parse().unwrap(); + let other_client: SocketAddr = "192.0.2.2:1000".parse().unwrap(); + let target: SocketAddr = "198.51.100.3:53".parse().unwrap(); + let inbound = Arc::new(MockSocket::default()); + let outbound = MockSocket::default(); + + let mut request = new_udp_header(target).unwrap(); + request.extend_from_slice(b"first"); + let mut ignored = new_udp_header(target).unwrap(); + ignored.extend_from_slice(b"ignored"); + let mut fragmented = new_udp_header(target).unwrap(); + fragmented[2] = 1; + inbound.receives.lock().unwrap().extend([ + (request, client), + (ignored, other_client), + (fragmented, client), + ]); + + let association = HostSocks5UdpAssociation { + inbound, + host: Arc::new(MockHost::default()), + socket_context: SocketContext::default(), + client: Mutex::new(None), + }; + transfer_requests(&association, &outbound).await.unwrap(); + + assert_eq!(*association.client.lock().unwrap(), Some(client)); + let sends = outbound.sends.lock().unwrap(); + assert_eq!(sends.len(), 1); + assert_eq!(sends[0].0, b"first"); + assert_eq!(sends[0].1, "[::ffff:198.51.100.3]:53".parse().unwrap()); + } + + #[tokio::test] + async fn udp_requests_reject_domain_targets_without_dns() { + let client: SocketAddr = "192.0.2.1:1000".parse().unwrap(); + let inbound = Arc::new(MockSocket::default()); + let outbound = MockSocket::default(); + let mut request = new_udp_header(("peer.example", 53)).unwrap(); + request.extend_from_slice(b"query"); + inbound + .receives + .lock() + .unwrap() + .push_back((request, client)); + + let association = HostSocks5UdpAssociation { + inbound, + host: Arc::new(MockHost::default()), + socket_context: SocketContext::default(), + client: Mutex::new(None), + }; + let error = transfer_requests(&association, &outbound) + .await + .unwrap_err(); + + assert!(error.to_string().contains("must be explicitly resolved")); + assert!(outbound.sends.lock().unwrap().is_empty()); + } + + #[test] + fn udp_response_preserves_ipv6_wire_family_for_mapped_ipv4() { + let remote = ipv4_mapped_addr("198.51.100.3:53".parse().unwrap()); + let header = new_udp_header(remote).unwrap(); + + assert_eq!( + header[3], + crate::gateway::socks5::codec::consts::SOCKS5_ADDR_TYPE_IPV6 + ); + } +} diff --git a/easytier-core/src/gateway/socks5/mod.rs b/easytier-core/src/gateway/socks5/mod.rs new file mode 100644 index 00000000..dfa2084e --- /dev/null +++ b/easytier-core/src/gateway/socks5/mod.rs @@ -0,0 +1,13 @@ +//! SOCKS5 protocol and Host-listener Adapter for the gateway data plane. + +#![forbid(unsafe_code)] + +mod adapter; +mod codec; +mod host; +mod server; + +pub(crate) use adapter::Socks5GatewayAdapter; +pub(crate) use codec::{Result, SocksError}; +pub(crate) use host::HostSocks5ServerRuntime; +pub(crate) use server::{AcceptAuthentication, AsyncTcpConnector, Config, Socks5Socket}; diff --git a/easytier/src/gateway/fast_socks5/server.rs b/easytier-core/src/gateway/socks5/server.rs similarity index 69% rename from easytier/src/gateway/fast_socks5/server.rs rename to easytier-core/src/gateway/socks5/server.rs index 5eea5aaf..a35d5b62 100644 --- a/easytier/src/gateway/fast_socks5/server.rs +++ b/easytier-core/src/gateway/socks5/server.rs @@ -1,29 +1,21 @@ -use super::Socks5Command; -use super::new_udp_header; -use super::parse_udp_request; -use super::read_exact; -use super::util::stream::tcp_connect_with_timeout; -use super::util::target_addr::{TargetAddr, read_address}; -use super::{AuthenticationMethod, ReplyError, Result, SocksError, consts}; +use super::codec::{ + AuthenticationMethod, ReplyError, Result, Socks5Command, SocksError, TargetAddr, consts, + read_address, +}; use anyhow::Context; use std::io; use std::net::IpAddr; use std::net::Ipv4Addr; -use std::net::{SocketAddr, ToSocketAddrs as StdToSocketAddrs}; -use std::ops::Deref; +use std::net::SocketAddr; use std::pin::Pin; use std::sync::Arc; use std::task::Poll; use tokio::io::AsyncReadExt; use tokio::io::{AsyncRead, AsyncWrite, AsyncWriteExt}; -use tokio::net::TcpStream; -use tokio::net::UdpSocket; -use tokio::try_join; - use tracing::{debug, error, info, trace}; #[derive(Clone)] -pub struct Config { +pub(crate) struct Config { /// Timeout of the command request request_timeout: u64, /// Avoid useless roundtrips if we don't need the Authentication layer @@ -57,51 +49,15 @@ impl Default for Config { /// Use this trait to handle a custom authentication on your end. #[async_trait::async_trait] -pub trait Authentication: Send + Sync { +pub(crate) trait Authentication: Send + Sync { type Item; async fn authenticate(&self, credentials: Option<(String, String)>) -> Option; } -/// Basic user/pass auth method provided. -pub struct SimpleUserPassword { - pub username: String, - pub password: String, -} - -/// The struct returned when the user has successfully authenticated -pub struct AuthSucceeded { - pub username: String, -} - -/// This is an example to auth via simple credentials. -/// If the auth succeed, we return the username authenticated with, for further uses. -#[async_trait::async_trait] -impl Authentication for SimpleUserPassword { - type Item = AuthSucceeded; - - async fn authenticate(&self, credentials: Option<(String, String)>) -> Option { - if let Some((username, password)) = credentials { - // Client has supplied credentials - if username == self.username && password == self.password { - // Some() will allow the authentication and the credentials - // will be forwarded to the socket - Some(AuthSucceeded { username }) - } else { - // Credentials incorrect, we deny the auth - None - } - } else { - // The client hasn't supplied any credentials, which only happens - // when `Config::allow_no_auth()` is set as `true` - None - } - } -} - /// This will simply return Option::None, which denies the authentication #[derive(Copy, Clone, Default)] -pub struct DenyAuthentication {} +pub(crate) struct DenyAuthentication {} #[async_trait::async_trait] impl Authentication for DenyAuthentication { @@ -114,7 +70,7 @@ impl Authentication for DenyAuthentication { /// While this one will always allow the user in. #[derive(Copy, Clone, Default)] -pub struct AcceptAuthentication {} +pub(crate) struct AcceptAuthentication {} #[async_trait::async_trait] impl Authentication for AcceptAuthentication { @@ -140,112 +96,71 @@ impl Config { self } - /// Enable authentication - /// 'static lifetime for Authentication avoid us to use `dyn Authentication` - /// and set the Arc before calling the function. - pub fn with_authentication(self, authentication: T) -> Config { - Config { - request_timeout: self.request_timeout, - skip_auth: self.skip_auth, - dns_resolve: self.dns_resolve, - execute_command: self.execute_command, - allow_udp: self.allow_udp, - allow_no_auth: self.allow_no_auth, - auth: Some(Arc::new(authentication)), - } - } - /// For some complex scenarios, we may want to either accept Username/Password configuration /// or IP Whitelisting, in case the client send only 2 auth methods rather than 3 (with auth) pub fn set_allow_no_auth(&mut self, value: bool) -> &mut Self { self.allow_no_auth = value; self } - - /// Set whether or not to execute commands - pub fn set_execute_command(&mut self, value: bool) -> &mut Self { - self.execute_command = value; - self - } - - /// Will the server perform dns resolve - pub fn set_dns_resolve(&mut self, value: bool) -> &mut Self { - self.dns_resolve = value; - self - } - - /// Set whether or not to allow udp traffic - pub fn set_udp_support(&mut self, value: bool) -> &mut Self { - self.allow_udp = value; - self - } } #[async_trait::async_trait] -pub trait AsyncTcpConnector { - type S: AsyncRead + AsyncWrite + Unpin + Send + Sync; +pub(crate) trait AsyncTcpConnector { + type S: AsyncRead + AsyncWrite + Unpin + Send; async fn tcp_connect(&self, addr: SocketAddr, timeout_s: u64) -> Result; } -pub struct DefaultTcpConnector {} +#[async_trait::async_trait] +pub(crate) trait Socks5UdpAssociation: Send { + fn local_addr(&self) -> io::Result; + + async fn transfer(self: Box) -> Result<()>; +} #[async_trait::async_trait] -impl AsyncTcpConnector for DefaultTcpConnector { - type S = TcpStream; +pub(crate) trait Socks5ServerRuntime: Send + Sync { + async fn resolve_dns(&self, target_addr: TargetAddr) -> Result; - async fn tcp_connect(&self, addr: SocketAddr, timeout_s: u64) -> Result { - tcp_connect_with_timeout(addr, timeout_s).await - } + async fn bind_udp_association(&self) -> Result>; } /// Wrap TcpStream and contains Socks5 protocol implementation. -pub struct Socks5Socket -{ +pub(crate) struct Socks5Socket< + T: AsyncRead + AsyncWrite + Unpin, + A: Authentication, + C: AsyncTcpConnector, +> { inner: T, config: Arc>, - auth: AuthenticationMethod, target_addr: Option, cmd: Option, /// Socket address which will be used in the reply message. reply_ip: Option, - /// If the client has been authenticated, that's where we store his credentials - /// to be accessed from the socket - credentials: Option, tcp_connector: C, + runtime: Arc, } impl Socks5Socket { - pub fn new(socket: T, config: Arc>, tcp_connector: C) -> Self { + pub fn new( + socket: T, + config: Arc>, + tcp_connector: C, + runtime: Arc, + ) -> Self { Socks5Socket { inner: socket, config, - auth: AuthenticationMethod::None, target_addr: None, cmd: None, reply_ip: None, - credentials: None, tcp_connector, + runtime, } } - /// Set the bind IP address in Socks5Reply. - /// - /// Only the inner socket owner knows the correct reply bind addr, so leave this field to be - /// populated. For those strict clients, users can use this function to set the correct IP - /// address. - /// - /// Most popular SOCKS5 clients [1] [2] ignore BND.ADDR and BND.PORT the reply of command - /// CONNECT, but this field could be useful in some other command, such as UDP ASSOCIATE. - /// - /// [1]: https://github.com/chromium/chromium/blob/bd2c7a8b65ec42d806277dd30f138a673dec233a/net/socket/socks5_client_socket.cc#L481 - /// [2]: https://github.com/curl/curl/blob/d15692ebbad5e9cfb871b0f7f51a73e43762cee2/lib/socks.c#L978 - pub fn set_reply_ip(&mut self, addr: IpAddr) { - self.reply_ip = Some(addr); - } - /// Process clients SOCKS requests /// This is the entry point where a whole request is processed. pub async fn upgrade_to_socks5(mut self) -> Result> { @@ -258,8 +173,7 @@ impl let auth_method = self.can_accept_method(methods).await?; if self.config.auth.is_some() { - let credentials = self.authenticate(auth_method).await?; - self.credentials = Some(credentials); + self.authenticate(auth_method).await?; } } else { debug!("skipping auth"); @@ -279,11 +193,6 @@ impl Ok(self) } - /// Consumes the `Socks5Socket`, returning the wrapped stream. - pub fn into_inner(self) -> T { - self.inner - } - /// Read the authentication method provided by the client. /// A client send a list of methods that he supports, he could send /// @@ -303,8 +212,12 @@ impl async fn get_methods(&mut self) -> Result> { trace!("Socks5Socket: get_methods()"); // read the first 2 bytes which contains the SOCKS version and the methods len() - let [version, methods_len] = - read_exact!(self.inner, [0u8; 2]).context("Can't read methods")?; + let mut header = [0u8; 2]; + self.inner + .read_exact(&mut header) + .await + .context("Can't read methods")?; + let [version, methods_len] = header; debug!( "Handshake headers: [version: {version}, methods len: {len}]", version = version, @@ -318,7 +231,10 @@ impl // {METHODS available from the client} // eg. (non-auth) {0, 1} // eg. (auth) {0, 1, 2} - let methods = read_exact!(self.inner, vec![0u8; methods_len as usize]) + let mut methods = vec![0u8; methods_len as usize]; + self.inner + .read_exact(&mut methods) + .await .context("Can't get methods.")?; debug!("methods supported sent by the client: {:?}", &methods); @@ -389,7 +305,12 @@ impl async fn read_username_password(socket: &mut T) -> Result<(String, String)> { trace!("Socks5Socket: authenticate()"); - let [version, user_len] = read_exact!(socket, [0u8; 2]).context("Can't read user len")?; + let mut header = [0u8; 2]; + socket + .read_exact(&mut header) + .await + .context("Can't read user len")?; + let [version, user_len] = header; debug!( "Auth: [version: {version}, user len: {len}]", version = version, @@ -403,11 +324,19 @@ impl ))); } - let username = - read_exact!(socket, vec![0u8; user_len as usize]).context("Can't get username.")?; + let mut username = vec![0u8; user_len as usize]; + socket + .read_exact(&mut username) + .await + .context("Can't get username.")?; debug!("username bytes: {:?}", &username); - let [pass_len] = read_exact!(socket, [0u8; 1]).context("Can't read pass len")?; + let mut pass_len = [0u8; 1]; + socket + .read_exact(&mut pass_len) + .await + .context("Can't read pass len")?; + let [pass_len] = pass_len; debug!("Auth: [pass len: {len}]", len = pass_len,); if pass_len < 1 { @@ -417,8 +346,11 @@ impl ))); } - let password = - read_exact!(socket, vec![0u8; pass_len as usize]).context("Can't get password.")?; + let mut password = vec![0u8; pass_len as usize]; + socket + .read_exact(&mut password) + .await + .context("Can't get password.")?; debug!("password bytes: {:?}", &password); let username = String::from_utf8(username).context("Failed to convert username")?; @@ -515,8 +447,12 @@ impl /// It the request is correct, it should returns a ['SocketAddr']. /// async fn read_command(&mut self) -> Result<()> { - let [version, cmd, rsv, address_type] = - read_exact!(self.inner, [0u8; 4]).context("Malformed request")?; + let mut header = [0u8; 4]; + self.inner + .read_exact(&mut header) + .await + .context("Malformed request")?; + let [version, cmd, rsv, address_type] = header; debug!( "Request: [version: {version}, command: {cmd}, rev: {rsv}, address_type: {address_type}]", version = version, @@ -569,7 +505,7 @@ impl if let Some(target_addr) = self.target_addr.take() { // decide whether we have to resolve DNS or not self.target_addr = match target_addr { - TargetAddr::Domain(_, _) => Some(target_addr.resolve_dns().await?), + TargetAddr::Domain(_, _) => Some(self.runtime.resolve_dns(target_addr).await?), TargetAddr::Ip(_) => Some(target_addr), }; } @@ -598,15 +534,15 @@ impl /// Connect to the target address that the client wants, /// then forward the data between them (client <=> target address). async fn execute_command_connect(&mut self) -> Result<()> { - // async-std's ToSocketAddrs doesn't supports external trait implementation - // @see https://github.com/async-rs/async-std/issues/539 - let addr = self - .target_addr - .as_ref() - .context("target_addr empty")? - .to_socket_addrs()? - .next() - .context("unreachable")?; + let addr = match self.target_addr.as_ref().context("target_addr empty")? { + TargetAddr::Ip(addr) => *addr, + TargetAddr::Domain(_, _) => { + return Err(io::Error::other( + "domain must be resolved when SOCKS DNS resolution is enabled", + ) + .into()); + } + }; // TCP connect with timeout, to avoid memory leak for connection that takes forever let outbound = self @@ -645,7 +581,7 @@ impl // Listen with UDP6 socket, so the client can connect to it with either // IPv4 or IPv6. - let peer_sock = UdpSocket::bind("[::]:0").await?; + let association = self.runtime.bind_udp_association().await?; // Respect the pre-populated reply IP address. self.inner @@ -653,7 +589,7 @@ impl &ReplyError::Succeeded, SocketAddr::new( self.reply_ip.context("invalid reply ip")?, - peer_sock.local_addr()?.port(), + association.local_addr()?.port(), ), )) .await @@ -661,39 +597,10 @@ impl debug!("Wrote success"); - transfer_udp(peer_sock).await?; + association.transfer().await?; Ok(()) } - - pub fn target_addr(&self) -> Option<&TargetAddr> { - self.target_addr.as_ref() - } - - pub fn auth(&self) -> &AuthenticationMethod { - &self.auth - } - - pub fn cmd(&self) -> &Option { - &self.cmd - } - - /// Borrow the credentials of the user has authenticated with - pub fn get_credentials(&self) -> Option<&<::Item as Deref>::Target> - where - ::Item: Deref, - { - self.credentials.as_deref() - } - - /// Get the credentials of the user has authenticated with - pub fn take_credentials(&mut self) -> Option { - self.credentials.take() - } - - pub fn tcp_connector(&self) -> &C { - &self.tcp_connector - } } /// Copy data between two peers @@ -711,59 +618,6 @@ where Ok(()) } -async fn handle_udp_request(inbound: &UdpSocket, outbound: &UdpSocket) -> Result<()> { - let mut buf = vec![0u8; 0x10000]; - loop { - let (size, client_addr) = inbound.recv_from(&mut buf).await?; - debug!("Server recieve udp from {}", client_addr); - inbound.connect(client_addr).await?; - - let (frag, target_addr, data) = parse_udp_request(&buf[..size]).await?; - - if frag != 0 { - debug!("Discard UDP frag packets sliently."); - return Ok(()); - } - - debug!("Server forward to packet to {}", target_addr); - let mut target_addr = target_addr - .to_socket_addrs()? - .next() - .context("unreachable")?; - - target_addr.set_ip(match target_addr.ip() { - std::net::IpAddr::V4(v4) => std::net::IpAddr::V6(v4.to_ipv6_mapped()), - v6 @ std::net::IpAddr::V6(_) => v6, - }); - outbound.send_to(data, target_addr).await?; - } -} - -async fn handle_udp_response(inbound: &UdpSocket, outbound: &UdpSocket) -> Result<()> { - let mut buf = vec![0u8; 0x10000]; - loop { - let (size, remote_addr) = outbound.recv_from(&mut buf).await?; - debug!("Recieve packet from {}", remote_addr); - - let mut data = new_udp_header(remote_addr)?; - data.extend_from_slice(&buf[..size]); - inbound.send(&data).await?; - } -} - -async fn transfer_udp(inbound: UdpSocket) -> Result<()> { - let outbound = UdpSocket::bind("[::]:0").await?; - - let req_fut = handle_udp_request(&inbound, &outbound); - let res_fut = handle_udp_response(&inbound, &outbound); - match try_join!(req_fut, res_fut) { - Ok(_) => {} - Err(error) => return Err(error), - } - - Ok(()) -} - // Fixes the issue "cannot borrow data in dereference of `Pin<&mut >` as mutable" // // cf. https://users.rust-lang.org/t/take-in-impl-future-cannot-borrow-data-in-a-dereference-of-pin/52042 @@ -840,3 +694,183 @@ fn new_reply(error: &ReplyError, sock_addr: SocketAddr) -> Vec { reply } + +#[cfg(test)] +mod tests { + use std::sync::Mutex; + + use tokio::io::{AsyncReadExt, AsyncWriteExt, DuplexStream}; + + use super::*; + + struct SimpleUserPassword { + username: String, + password: String, + } + + #[async_trait::async_trait] + impl Authentication for SimpleUserPassword { + type Item = (); + + async fn authenticate(&self, credentials: Option<(String, String)>) -> Option { + credentials + .filter(|(username, password)| { + username == &self.username && password == &self.password + }) + .map(|_| ()) + } + } + + impl Config { + fn with_authentication(self, authentication: T) -> Config { + Config { + request_timeout: self.request_timeout, + skip_auth: self.skip_auth, + dns_resolve: self.dns_resolve, + execute_command: self.execute_command, + allow_udp: self.allow_udp, + allow_no_auth: self.allow_no_auth, + auth: Some(Arc::new(authentication)), + } + } + } + + struct TestConnector { + outbound: Mutex>, + } + + #[async_trait::async_trait] + impl AsyncTcpConnector for TestConnector { + type S = DuplexStream; + + async fn tcp_connect(&self, _addr: SocketAddr, _timeout_s: u64) -> Result { + self.outbound + .lock() + .unwrap() + .take() + .context("test outbound already taken") + .map_err(Into::into) + } + } + + struct TestRuntime; + + #[async_trait::async_trait] + impl Socks5ServerRuntime for TestRuntime { + async fn resolve_dns(&self, target_addr: TargetAddr) -> Result { + Ok(target_addr) + } + + async fn bind_udp_association(&self) -> Result> { + Err(SocksError::Io(io::Error::new( + io::ErrorKind::Unsupported, + "UDP is not used by this test runtime", + ))) + } + } + + fn connector(outbound: DuplexStream) -> TestConnector { + TestConnector { + outbound: Mutex::new(Some(outbound)), + } + } + + #[tokio::test] + async fn no_auth_connect_handshake_transfers_bidirectionally() { + let (server_stream, mut client_stream) = tokio::io::duplex(1024); + let (outbound, mut destination_stream) = tokio::io::duplex(1024); + let mut config = Config::::default(); + config.set_allow_no_auth(true); + let socket = Socks5Socket::new( + server_stream, + Arc::new(config), + connector(outbound), + Arc::new(TestRuntime), + ); + let task = tokio::spawn(socket.upgrade_to_socks5()); + + client_stream.write_all(&[5, 1, 0]).await.unwrap(); + let mut method_reply = [0u8; 2]; + client_stream.read_exact(&mut method_reply).await.unwrap(); + assert_eq!(method_reply, [5, 0]); + + client_stream + .write_all(&[5, 1, 0, 1, 10, 42, 0, 7, 0, 80]) + .await + .unwrap(); + let mut connect_reply = [0u8; 10]; + client_stream.read_exact(&mut connect_reply).await.unwrap(); + assert_eq!(connect_reply[0..2], [5, 0]); + + client_stream.write_all(b"request").await.unwrap(); + let mut request = [0u8; 7]; + destination_stream.read_exact(&mut request).await.unwrap(); + assert_eq!(&request, b"request"); + + destination_stream.write_all(b"response").await.unwrap(); + let mut response = [0u8; 8]; + client_stream.read_exact(&mut response).await.unwrap(); + assert_eq!(&response, b"response"); + + drop(client_stream); + drop(destination_stream); + assert!(task.await.unwrap().is_ok()); + } + + #[tokio::test] + async fn unresolved_domain_keeps_dns_disabled_failure_semantics() { + let (server_stream, _client_stream) = tokio::io::duplex(128); + let (outbound, _destination_stream) = tokio::io::duplex(128); + let mut socket = Socks5Socket::new( + server_stream, + Arc::new(Config::::default()), + connector(outbound), + Arc::new(TestRuntime), + ); + socket.target_addr = Some(TargetAddr::Domain("peer.example".into(), 443)); + + let error = socket.execute_command_connect().await.unwrap_err(); + + assert!(matches!(error, SocksError::Io(_))); + } + + #[tokio::test] + async fn password_authentication_rejection_preserves_wire_reply() { + let (server_stream, mut client_stream) = tokio::io::duplex(128); + let (outbound, _destination_stream) = tokio::io::duplex(128); + let config = + Config::::default().with_authentication(SimpleUserPassword { + username: "user".to_string(), + password: "correct".to_string(), + }); + let socket = Socks5Socket::new( + server_stream, + Arc::new(config), + connector(outbound), + Arc::new(TestRuntime), + ); + let task = tokio::spawn(socket.upgrade_to_socks5()); + + client_stream.write_all(&[5, 1, 2]).await.unwrap(); + let mut method_reply = [0u8; 2]; + client_stream.read_exact(&mut method_reply).await.unwrap(); + assert_eq!(method_reply, [5, 2]); + + client_stream + .write_all(&[ + 1, 4, b'u', b's', b'e', b'r', 5, b'w', b'r', b'o', b'n', b'g', + ]) + .await + .unwrap(); + let mut auth_reply = [0u8; 2]; + client_stream.read_exact(&mut auth_reply).await.unwrap(); + assert_eq!(auth_reply, [1, 0xff]); + + let err = task + .await + .unwrap() + .err() + .expect("wrong password must reject authentication"); + assert!(matches!(err, SocksError::AuthenticationRejected(_))); + } +} diff --git a/easytier-core/src/gateway/udp_broadcast.rs b/easytier-core/src/gateway/udp_broadcast.rs new file mode 100644 index 00000000..ee95a92a --- /dev/null +++ b/easytier-core/src/gateway/udp_broadcast.rs @@ -0,0 +1,740 @@ +use std::net::Ipv4Addr; + +use cidr::Ipv4Inet; +use pnet_packet::{ + ip::IpNextHeaderProtocols, + ipv4::{self, Ipv4Flags, Ipv4Packet, MutableIpv4Packet}, + udp::{self, MutableUdpPacket, UdpPacket}, +}; + +#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] +pub struct PhysicalInterface { + addr: Ipv4Addr, + directed_broadcast: Ipv4Addr, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub struct NonContiguousIpv4Netmask(Ipv4Addr); + +impl NonContiguousIpv4Netmask { + pub fn netmask(self) -> Ipv4Addr { + self.0 + } +} + +impl PhysicalInterface { + pub fn from_observation( + addr: Ipv4Addr, + netmask: Option, + is_internal: bool, + virtual_addr: Ipv4Addr, + ) -> Result, NonContiguousIpv4Netmask> { + if is_internal || addr == virtual_addr { + return Ok(None); + } + + let Some(netmask) = netmask else { + return Ok(None); + }; + let prefix = prefix_len_from_netmask(netmask).ok_or(NonContiguousIpv4Netmask(netmask))?; + Ok(Self::from_ip_and_prefix(addr, prefix)) + } + + pub fn from_ip_and_prefix(addr: Ipv4Addr, prefix: u8) -> Option { + if should_ignore_interface_addr(addr) || prefix > 30 { + return None; + } + + Some(Self { + addr, + directed_broadcast: directed_broadcast(addr, prefix)?, + }) + } + + pub fn address(&self) -> Ipv4Addr { + self.addr + } + + pub fn directed_broadcast(&self) -> Ipv4Addr { + self.directed_broadcast + } +} + +#[derive(Debug, Clone)] +pub struct BroadcastRelayConfig { + virtual_ipv4: Ipv4Inet, + physical_interfaces: Vec, +} + +impl BroadcastRelayConfig { + pub fn new(virtual_ipv4: Ipv4Inet, physical_interfaces: Vec) -> Self { + let mut eligible_interfaces = Vec::with_capacity(physical_interfaces.len()); + for interface in physical_interfaces { + if interface.addr == virtual_ipv4.address() || eligible_interfaces.contains(&interface) + { + continue; + } + eligible_interfaces.push(interface); + } + + Self { + virtual_ipv4, + physical_interfaces: eligible_interfaces, + } + } + + pub fn virtual_ipv4(&self) -> &Ipv4Inet { + &self.virtual_ipv4 + } + + pub fn physical_interfaces(&self) -> &[PhysicalInterface] { + &self.physical_interfaces + } + + fn is_physical_source(&self, addr: Ipv4Addr) -> bool { + self.physical_interfaces + .iter() + .any(|iface| iface.addr == addr) + } + + fn normalize_destination(&self, dst: Ipv4Addr) -> Option { + if dst.is_broadcast() || dst.is_multicast() { + return Some(dst); + } + + self.physical_interfaces + .iter() + .any(|iface| iface.directed_broadcast == dst) + .then_some(self.virtual_ipv4.last_address()) + } +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct NormalizedPacket { + pub packet: Vec, + pub destination: Ipv4Addr, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub struct UdpPacketSummary { + pub src: Ipv4Addr, + pub dst: Ipv4Addr, + pub src_port: u16, + pub dst_port: u16, + pub ip_len: usize, + pub udp_len: usize, + pub payload_len: usize, +} + +impl UdpPacketSummary { + pub fn parse(packet: &[u8]) -> Option { + let ipv4_packet = Ipv4Packet::new(packet)?; + if ipv4_packet.get_version() != 4 + || ipv4_packet.get_next_level_protocol() != IpNextHeaderProtocols::Udp + { + return None; + } + + let header_len = usize::from(ipv4_packet.get_header_length()) * 4; + let total_len = usize::from(ipv4_packet.get_total_length()); + if header_len < Ipv4Packet::minimum_packet_size() + || total_len < header_len + UdpPacket::minimum_packet_size() + || total_len > packet.len() + { + return None; + } + + let udp_packet = UdpPacket::new(&packet[header_len..total_len])?; + let udp_len = usize::from(udp_packet.get_length()); + if udp_len < UdpPacket::minimum_packet_size() || header_len + udp_len != total_len { + return None; + } + + Some(Self { + src: ipv4_packet.get_source(), + dst: ipv4_packet.get_destination(), + src_port: udp_packet.get_source(), + dst_port: udp_packet.get_destination(), + ip_len: total_len, + udp_len, + payload_len: udp_len - UdpPacket::minimum_packet_size(), + }) + } +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum UdpBroadcastPacketRejection { + MalformedIpv4, + NotUdpIpv4, + Fragmented, + BadIpv4Length, + IgnoredSource, + VirtualSourceDuplicate, + NonPhysicalSource, + UnsupportedDestination, + LoopbackDestination, + MalformedUdp, + BadUdpLength, +} + +impl UdpBroadcastPacketRejection { + pub fn reason(self) -> &'static str { + match self { + Self::MalformedIpv4 => "malformed_ipv4", + Self::NotUdpIpv4 => "not_udp_ipv4", + Self::Fragmented => "fragmented", + Self::BadIpv4Length => "bad_ipv4_length", + Self::IgnoredSource => "ignored_source", + Self::VirtualSourceDuplicate => "virtual_source_duplicate", + Self::NonPhysicalSource => "non_physical_source", + Self::UnsupportedDestination => "unsupported_destination", + Self::LoopbackDestination => "loopback_destination", + Self::MalformedUdp => "malformed_udp", + Self::BadUdpLength => "bad_udp_length", + } + } +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +struct ParsedUdpBroadcastPacket { + header_len: usize, + udp_len: usize, + normalized_destination: Ipv4Addr, +} + +fn should_ignore_interface_addr(addr: Ipv4Addr) -> bool { + addr.is_unspecified() || addr.is_loopback() || addr.is_multicast() || addr.is_broadcast() +} + +fn prefix_len_from_netmask(mask: Ipv4Addr) -> Option { + let raw = u32::from(mask); + let prefix = raw.count_ones() as u8; + let expected = if prefix == 0 { + 0 + } else { + u32::MAX << (32 - prefix) + }; + (raw == expected).then_some(prefix) +} + +fn directed_broadcast(addr: Ipv4Addr, prefix: u8) -> Option { + if prefix > 32 { + return None; + } + + let mask = if prefix == 0 { + 0 + } else { + u32::MAX << (32 - prefix) + }; + Some(Ipv4Addr::from(u32::from(addr) | !mask)) +} + +fn parse_udp_broadcast( + packet: &[u8], + config: &BroadcastRelayConfig, +) -> Result { + let ipv4_packet = Ipv4Packet::new(packet).ok_or(UdpBroadcastPacketRejection::MalformedIpv4)?; + if ipv4_packet.get_version() != 4 + || ipv4_packet.get_next_level_protocol() != IpNextHeaderProtocols::Udp + { + return Err(UdpBroadcastPacketRejection::NotUdpIpv4); + } + + if ipv4_packet.get_fragment_offset() != 0 + || ipv4_packet.get_flags() & Ipv4Flags::MoreFragments != 0 + { + return Err(UdpBroadcastPacketRejection::Fragmented); + } + + let header_len = usize::from(ipv4_packet.get_header_length()) * 4; + let total_len = usize::from(ipv4_packet.get_total_length()); + if header_len < Ipv4Packet::minimum_packet_size() + || total_len < header_len + UdpPacket::minimum_packet_size() + || total_len > packet.len() + { + return Err(UdpBroadcastPacketRejection::BadIpv4Length); + } + + let src = ipv4_packet.get_source(); + let dst = ipv4_packet.get_destination(); + if should_ignore_interface_addr(src) { + return Err(UdpBroadcastPacketRejection::IgnoredSource); + } + if src == config.virtual_ipv4.address() { + return Err(UdpBroadcastPacketRejection::VirtualSourceDuplicate); + } + if !config.is_physical_source(src) { + return Err(UdpBroadcastPacketRejection::NonPhysicalSource); + } + + let normalized_destination = config + .normalize_destination(dst) + .ok_or(UdpBroadcastPacketRejection::UnsupportedDestination)?; + if normalized_destination.is_loopback() { + return Err(UdpBroadcastPacketRejection::LoopbackDestination); + } + + let udp_packet = UdpPacket::new(&packet[header_len..total_len]) + .ok_or(UdpBroadcastPacketRejection::MalformedUdp)?; + let udp_len = usize::from(udp_packet.get_length()); + if udp_len < UdpPacket::minimum_packet_size() || header_len + udp_len != total_len { + return Err(UdpBroadcastPacketRejection::BadUdpLength); + } + + Ok(ParsedUdpBroadcastPacket { + header_len, + udp_len, + normalized_destination, + }) +} + +pub fn normalize_udp_broadcast_packet( + packet: &[u8], + config: &BroadcastRelayConfig, +) -> Result { + let parsed = parse_udp_broadcast(packet, config)?; + let header_len = parsed.header_len; + let packet_len = header_len + parsed.udp_len; + let destination = parsed.normalized_destination; + let virtual_ipv4 = config.virtual_ipv4.address(); + let mut normalized = packet[..packet_len].to_vec(); + + { + let mut ipv4_packet = MutableIpv4Packet::new(&mut normalized) + .ok_or(UdpBroadcastPacketRejection::MalformedIpv4)?; + ipv4_packet.set_source(virtual_ipv4); + ipv4_packet.set_destination(destination); + ipv4_packet.set_total_length(packet_len as u16); + ipv4_packet.set_checksum(0); + } + + { + let mut udp_packet = MutableUdpPacket::new(&mut normalized[header_len..packet_len]) + .ok_or(UdpBroadcastPacketRejection::MalformedUdp)?; + udp_packet.set_checksum(0); + let checksum = udp::ipv4_checksum(&udp_packet.to_immutable(), &virtual_ipv4, &destination); + udp_packet.set_checksum(checksum); + } + + { + let mut ipv4_packet = MutableIpv4Packet::new(&mut normalized) + .ok_or(UdpBroadcastPacketRejection::MalformedIpv4)?; + let checksum = ipv4::checksum(&ipv4_packet.to_immutable()); + ipv4_packet.set_checksum(checksum); + } + + Ok(NormalizedPacket { + packet: normalized, + destination, + }) +} + +#[derive(Clone)] +pub struct UdpBroadcastRelayStats { + packets_captured: crate::foundation::stats::CounterHandle, + packets_ignored: crate::foundation::stats::CounterHandle, + packets_forwarded: crate::foundation::stats::CounterHandle, + packets_forward_failed: crate::foundation::stats::CounterHandle, +} + +impl UdpBroadcastRelayStats { + pub(crate) fn new( + packets_captured: crate::foundation::stats::CounterHandle, + packets_ignored: crate::foundation::stats::CounterHandle, + packets_forwarded: crate::foundation::stats::CounterHandle, + packets_forward_failed: crate::foundation::stats::CounterHandle, + ) -> Self { + Self { + packets_captured, + packets_ignored, + packets_forwarded, + packets_forward_failed, + } + } + + pub fn record_captured(&self) { + self.packets_captured.inc(); + } + + pub fn record_ignored(&self) { + self.packets_ignored.inc(); + } + + pub fn record_forwarded(&self) { + self.packets_forwarded.inc(); + } + + pub fn record_forward_failed(&self) { + self.packets_forward_failed.inc(); + } +} + +#[cfg(test)] +mod tests { + use super::*; + use pnet_packet::{MutablePacket, Packet}; + + fn config() -> BroadcastRelayConfig { + BroadcastRelayConfig::new( + "10.144.144.1/24".parse().unwrap(), + vec![PhysicalInterface::from_ip_and_prefix(Ipv4Addr::new(192, 168, 1, 7), 24).unwrap()], + ) + } + + fn build_udp_packet(src: Ipv4Addr, dst: Ipv4Addr, payload: &[u8]) -> Vec { + let mut packet = vec![0; 20 + 8 + payload.len()]; + { + let mut ipv4_packet = MutableIpv4Packet::new(&mut packet).unwrap(); + ipv4_packet.set_version(4); + ipv4_packet.set_header_length(5); + ipv4_packet.set_total_length((20 + 8 + payload.len()) as u16); + ipv4_packet.set_ttl(64); + ipv4_packet.set_next_level_protocol(IpNextHeaderProtocols::Udp); + ipv4_packet.set_source(src); + ipv4_packet.set_destination(dst); + } + + { + let mut udp_packet = MutableUdpPacket::new(&mut packet[20..]).unwrap(); + udp_packet.set_source(12345); + udp_packet.set_destination(37020); + udp_packet.set_length((8 + payload.len()) as u16); + udp_packet.payload_mut().copy_from_slice(payload); + let checksum = udp::ipv4_checksum(&udp_packet.to_immutable(), &src, &dst); + udp_packet.set_checksum(checksum); + } + + { + let mut ipv4_packet = MutableIpv4Packet::new(&mut packet).unwrap(); + let checksum = ipv4::checksum(&ipv4_packet.to_immutable()); + ipv4_packet.set_checksum(checksum); + } + + packet + } + + fn assert_valid_checksums(packet: &[u8]) { + let ipv4_packet = Ipv4Packet::new(packet).unwrap(); + assert_eq!(ipv4::checksum(&ipv4_packet), ipv4_packet.get_checksum()); + let udp_packet = UdpPacket::new(ipv4_packet.payload()).unwrap(); + assert_eq!( + udp::ipv4_checksum( + &udp_packet, + &ipv4_packet.get_source(), + &ipv4_packet.get_destination() + ), + udp_packet.get_checksum() + ); + } + + #[test] + fn rewrites_limited_broadcast() { + let packet = build_udp_packet(Ipv4Addr::new(192, 168, 1, 7), Ipv4Addr::BROADCAST, b"hello"); + + let normalized = normalize_udp_broadcast_packet(&packet, &config()).unwrap(); + let ipv4_packet = Ipv4Packet::new(&normalized.packet).unwrap(); + + assert_eq!(normalized.destination, Ipv4Addr::BROADCAST); + assert_eq!(ipv4_packet.get_source(), Ipv4Addr::new(10, 144, 144, 1)); + assert_eq!(ipv4_packet.get_destination(), Ipv4Addr::BROADCAST); + assert_eq!(&ipv4_packet.payload()[8..], b"hello"); + assert_valid_checksums(&normalized.packet); + } + + #[test] + fn rewrites_directed_broadcast() { + let packet = build_udp_packet( + Ipv4Addr::new(192, 168, 1, 7), + Ipv4Addr::new(192, 168, 1, 255), + b"directed", + ); + + let normalized = normalize_udp_broadcast_packet(&packet, &config()).unwrap(); + let ipv4_packet = Ipv4Packet::new(&normalized.packet).unwrap(); + + assert_eq!(normalized.destination, Ipv4Addr::new(10, 144, 144, 255)); + assert_eq!(ipv4_packet.get_source(), Ipv4Addr::new(10, 144, 144, 1)); + assert_eq!( + ipv4_packet.get_destination(), + Ipv4Addr::new(10, 144, 144, 255) + ); + assert_eq!(&ipv4_packet.payload()[8..], b"directed"); + assert_valid_checksums(&normalized.packet); + } + + #[test] + fn preserves_multicast_destination() { + let multicast = Ipv4Addr::new(239, 255, 255, 250); + let packet = build_udp_packet(Ipv4Addr::new(192, 168, 1, 7), multicast, b"multicast"); + + let normalized = normalize_udp_broadcast_packet(&packet, &config()).unwrap(); + let ipv4_packet = Ipv4Packet::new(&normalized.packet).unwrap(); + + assert_eq!(normalized.destination, multicast); + assert_eq!(ipv4_packet.get_source(), Ipv4Addr::new(10, 144, 144, 1)); + assert_eq!(ipv4_packet.get_destination(), multicast); + assert_eq!(&ipv4_packet.payload()[8..], b"multicast"); + assert_valid_checksums(&normalized.packet); + } + + #[test] + fn rejects_malformed_packets() { + assert_eq!( + normalize_udp_broadcast_packet(&[], &config()), + Err(UdpBroadcastPacketRejection::MalformedIpv4) + ); + + let mut packet = + build_udp_packet(Ipv4Addr::new(192, 168, 1, 7), Ipv4Addr::BROADCAST, b"bad"); + packet[2..4].copy_from_slice(&10u16.to_be_bytes()); + assert_eq!( + normalize_udp_broadcast_packet(&packet, &config()), + Err(UdpBroadcastPacketRejection::BadIpv4Length) + ); + } + + #[test] + fn rejects_fragments() { + let mut packet = build_udp_packet( + Ipv4Addr::new(192, 168, 1, 7), + Ipv4Addr::BROADCAST, + b"fragment", + ); + { + let mut ipv4_packet = MutableIpv4Packet::new(&mut packet).unwrap(); + ipv4_packet.set_flags(Ipv4Flags::MoreFragments); + } + + assert_eq!( + normalize_udp_broadcast_packet(&packet, &config()), + Err(UdpBroadcastPacketRejection::Fragmented) + ); + } + + #[test] + fn rejects_non_broadcast_destinations() { + let packet = build_udp_packet( + Ipv4Addr::new(192, 168, 1, 7), + Ipv4Addr::new(192, 168, 1, 10), + b"unicast", + ); + + assert_eq!( + normalize_udp_broadcast_packet(&packet, &config()), + Err(UdpBroadcastPacketRejection::UnsupportedDestination) + ); + } + + #[test] + fn rejects_virtual_source_duplicates() { + let packet = build_udp_packet(Ipv4Addr::new(10, 144, 144, 1), Ipv4Addr::BROADCAST, b"loop"); + + assert_eq!( + normalize_udp_broadcast_packet(&packet, &config()), + Err(UdpBroadcastPacketRejection::VirtualSourceDuplicate) + ); + } + + #[test] + fn rejects_non_udp_ipv4_packets() { + let mut packet = + build_udp_packet(Ipv4Addr::new(192, 168, 1, 7), Ipv4Addr::BROADCAST, b"tcp"); + MutableIpv4Packet::new(&mut packet) + .unwrap() + .set_next_level_protocol(IpNextHeaderProtocols::Tcp); + + assert_eq!( + normalize_udp_broadcast_packet(&packet, &config()), + Err(UdpBroadcastPacketRejection::NotUdpIpv4) + ); + } + + #[test] + fn rejects_ignored_and_non_physical_sources() { + let ignored = build_udp_packet(Ipv4Addr::LOCALHOST, Ipv4Addr::BROADCAST, b"ignored"); + assert_eq!( + normalize_udp_broadcast_packet(&ignored, &config()), + Err(UdpBroadcastPacketRejection::IgnoredSource) + ); + + let non_physical = + build_udp_packet(Ipv4Addr::new(192, 168, 1, 8), Ipv4Addr::BROADCAST, b"other"); + assert_eq!( + normalize_udp_broadcast_packet(&non_physical, &config()), + Err(UdpBroadcastPacketRejection::NonPhysicalSource) + ); + } + + #[test] + fn rejects_loopback_destination_after_mapping() { + let config = BroadcastRelayConfig::new( + "127.0.0.1/24".parse().unwrap(), + vec![PhysicalInterface::from_ip_and_prefix(Ipv4Addr::new(192, 168, 1, 7), 24).unwrap()], + ); + let packet = build_udp_packet( + Ipv4Addr::new(192, 168, 1, 7), + Ipv4Addr::new(192, 168, 1, 255), + b"loopback", + ); + + assert_eq!( + normalize_udp_broadcast_packet(&packet, &config), + Err(UdpBroadcastPacketRejection::LoopbackDestination) + ); + } + + #[test] + fn summarizes_packets_and_rejects_bad_udp_length() { + let mut packet = build_udp_packet( + Ipv4Addr::new(192, 168, 1, 7), + Ipv4Addr::BROADCAST, + b"summary", + ); + assert_eq!( + UdpPacketSummary::parse(&packet), + Some(UdpPacketSummary { + src: Ipv4Addr::new(192, 168, 1, 7), + dst: Ipv4Addr::BROADCAST, + src_port: 12345, + dst_port: 37020, + ip_len: 35, + udp_len: 15, + payload_len: 7, + }) + ); + + packet[24..26].copy_from_slice(&8u16.to_be_bytes()); + assert_eq!( + normalize_udp_broadcast_packet(&packet, &config()), + Err(UdpBroadcastPacketRejection::BadUdpLength) + ); + assert_eq!(UdpPacketSummary::parse(&packet), None); + } + + #[test] + fn rejection_reasons_preserve_log_values() { + let reasons = [ + (UdpBroadcastPacketRejection::MalformedIpv4, "malformed_ipv4"), + (UdpBroadcastPacketRejection::NotUdpIpv4, "not_udp_ipv4"), + (UdpBroadcastPacketRejection::Fragmented, "fragmented"), + ( + UdpBroadcastPacketRejection::BadIpv4Length, + "bad_ipv4_length", + ), + (UdpBroadcastPacketRejection::IgnoredSource, "ignored_source"), + ( + UdpBroadcastPacketRejection::VirtualSourceDuplicate, + "virtual_source_duplicate", + ), + ( + UdpBroadcastPacketRejection::NonPhysicalSource, + "non_physical_source", + ), + ( + UdpBroadcastPacketRejection::UnsupportedDestination, + "unsupported_destination", + ), + ( + UdpBroadcastPacketRejection::LoopbackDestination, + "loopback_destination", + ), + (UdpBroadcastPacketRejection::MalformedUdp, "malformed_udp"), + (UdpBroadcastPacketRejection::BadUdpLength, "bad_udp_length"), + ]; + + for (rejection, reason) in reasons { + assert_eq!(rejection.reason(), reason); + } + } + + #[test] + fn detects_directed_broadcast_from_prefix() { + let physical = + PhysicalInterface::from_ip_and_prefix(Ipv4Addr::new(172, 16, 5, 10), 20).unwrap(); + assert_eq!( + physical.directed_broadcast(), + Ipv4Addr::new(172, 16, 15, 255) + ); + assert_eq!( + prefix_len_from_netmask(Ipv4Addr::new(255, 255, 240, 0)), + Some(20) + ); + assert_eq!(prefix_len_from_netmask(Ipv4Addr::new(255, 0, 255, 0)), None); + } + + #[test] + fn classifies_physical_interface_observations() { + let addr = Ipv4Addr::new(192, 168, 1, 7); + let virtual_addr = Ipv4Addr::new(10, 144, 144, 1); + let netmask = Ipv4Addr::new(255, 255, 255, 0); + let non_contiguous = Ipv4Addr::new(255, 0, 255, 0); + + assert_eq!( + PhysicalInterface::from_observation(addr, Some(netmask), false, virtual_addr), + Ok(PhysicalInterface::from_ip_and_prefix(addr, 24)) + ); + assert_eq!( + PhysicalInterface::from_observation(addr, None, false, virtual_addr), + Ok(None) + ); + assert_eq!( + PhysicalInterface::from_observation(addr, Some(non_contiguous), true, virtual_addr), + Ok(None) + ); + assert_eq!( + PhysicalInterface::from_observation( + virtual_addr, + Some(non_contiguous), + false, + virtual_addr, + ), + Ok(None) + ); + assert_eq!( + PhysicalInterface::from_observation(addr, Some(non_contiguous), false, virtual_addr,), + Err(NonContiguousIpv4Netmask(non_contiguous)) + ); + } + + #[test] + fn keeps_link_local_interfaces() { + let physical = + PhysicalInterface::from_ip_and_prefix(Ipv4Addr::new(169, 254, 13, 10), 16).unwrap(); + assert_eq!( + physical.directed_broadcast(), + Ipv4Addr::new(169, 254, 255, 255) + ); + } + + #[test] + fn rejects_ineligible_interface_addresses_and_prefixes() { + for addr in [ + Ipv4Addr::UNSPECIFIED, + Ipv4Addr::LOCALHOST, + Ipv4Addr::new(239, 1, 2, 3), + Ipv4Addr::BROADCAST, + ] { + assert_eq!(PhysicalInterface::from_ip_and_prefix(addr, 24), None); + } + assert_eq!( + PhysicalInterface::from_ip_and_prefix(Ipv4Addr::new(192, 168, 1, 7), 31), + None + ); + } + + #[test] + fn config_excludes_virtual_and_duplicate_interfaces() { + let virtual_interface = + PhysicalInterface::from_ip_and_prefix(Ipv4Addr::new(10, 144, 144, 1), 24).unwrap(); + let physical_interface = + PhysicalInterface::from_ip_and_prefix(Ipv4Addr::new(192, 168, 1, 7), 24).unwrap(); + + let config = BroadcastRelayConfig::new( + "10.144.144.1/24".parse().unwrap(), + vec![virtual_interface, physical_interface, physical_interface], + ); + + assert_eq!(config.physical_interfaces(), &[physical_interface]); + } +} diff --git a/easytier-core/src/gateway/vpn_portal.rs b/easytier-core/src/gateway/vpn_portal.rs new file mode 100644 index 00000000..353cd2bf --- /dev/null +++ b/easytier-core/src/gateway/vpn_portal.rs @@ -0,0 +1,713 @@ +use std::{ + net::{IpAddr, Ipv4Addr}, + sync::Arc, +}; + +use async_trait::async_trait; +use cidr::{Ipv4Cidr, Ipv4Inet}; +use dashmap::DashMap; +use futures::StreamExt; +use tokio::sync::Mutex; +use tokio::task::JoinSet; +use tokio_util::sync::CancellationToken; + +use crate::{ + config::runtime::CoreRuntimeConfigStore, + events::{CoreEvent, CoreEventSink}, + packet::{PacketType, ZCPacket, ZCPacketType}, + peers::{ + PeerPacketFilter, + peer_manager::{PeerManagerCore, PipelineRegistrationGuard}, + }, + socket::SocketListener, + tunnel::{Tunnel, mpsc::MpscTunnel, mpsc::MpscTunnelSender}, +}; + +const IPV4_HEADER_LEN: usize = 20; + +pub struct VpnPortalClient { + endpoint_addr: Option, + value: V, +} + +impl VpnPortalClient { + pub fn endpoint_addr(&self) -> Option<&url::Url> { + self.endpoint_addr.as_ref() + } + + pub fn value(&self) -> &V { + &self.value + } +} + +pub struct VpnPortalClientTable { + entries: DashMap>>, +} + +impl Default for VpnPortalClientTable { + fn default() -> Self { + Self { + entries: DashMap::new(), + } + } +} + +impl VpnPortalClientTable { + pub fn new() -> Self { + Self::default() + } + + pub fn len(&self) -> usize { + self.entries.len() + } + + pub fn is_empty(&self) -> bool { + self.entries.is_empty() + } + + pub fn endpoint_addrs(&self) -> Vec> { + self.entries + .iter() + .map(|entry| entry.value().endpoint_addr.clone()) + .collect() + } + + pub fn route_peer_packet(&self, packet: &ZCPacket) -> VpnPortalPeerPacketRoute { + let Some(header) = packet.peer_manager_header() else { + return VpnPortalPeerPacketRoute::Pass; + }; + if header.packet_type != PacketType::Data as u8 { + return VpnPortalPeerPacketRoute::Pass; + } + + let payload = packet.payload(); + if payload.len() < IPV4_HEADER_LEN { + return VpnPortalPeerPacketRoute::Drop; + } + if payload[0] >> 4 != 4 { + return VpnPortalPeerPacketRoute::Pass; + } + let destination = ipv4_address(&payload[16..20]); + let Some(client) = self + .entries + .get(&destination) + .map(|entry| entry.value().clone()) + else { + return VpnPortalPeerPacketRoute::Pass; + }; + + VpnPortalPeerPacketRoute::Deliver { + destination, + client, + } + } + + fn insert(&self, address: Ipv4Addr, client: Arc>) { + self.entries.insert(address, client); + } + + fn remove_if_current(&self, address: &Ipv4Addr, client: &Arc>) -> bool { + let removed = self + .entries + .remove_if(address, |_, current| Arc::ptr_eq(current, client)) + .is_some(); + if self.entries.capacity() - self.entries.len() > 16 { + self.entries.shrink_to_fit(); + } + removed + } +} + +pub enum VpnPortalPeerPacketRoute { + Pass, + Drop, + Deliver { + destination: Ipv4Addr, + client: Arc>, + }, +} + +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub struct VpnPortalClientPacket { + pub source: Ipv4Addr, + pub destination: Ipv4Addr, +} + +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub enum VpnPortalClientRemoval { + NotRegistered, + Removed(Ipv4Addr), + EntryChangedOrMissing(Ipv4Addr), +} + +pub struct VpnPortalClientSession { + table: Arc>, + client: Arc>, + registered_ip: Option, +} + +pub type VpnPortalListener = Box>>; + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct VpnPortalClientConfigPlan { + pub client_cidr: Ipv4Cidr, + pub allowed_ips: Vec, + pub listener_url: url::Url, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct VpnPortalInfoSnapshot { + pub vpn_type: String, + pub client_config: String, + pub connected_clients: Vec, +} + +#[async_trait] +pub trait VpnPortalHost: Send + Sync + 'static { + /// Creates already-listening protocol engines. Core owns accepting from the + /// returned listeners and all portable session lifecycle after this seam. + async fn start_listeners(&self) -> anyhow::Result>; + + fn name(&self) -> String; + + fn render_client_config(&self, plan: &VpnPortalClientConfigPlan) -> String; + + fn not_started_client_config(&self) -> String { + "ERROR: VPN Portal Not Started".to_owned() + } +} + +struct VpnPortalRuntime { + cancel: CancellationToken, + tasks: JoinSet<()>, + listener_urls: Vec, + _pipeline: PipelineRegistrationGuard, +} + +struct VpnPortalSessionEventGuard { + events: Arc, + portal: String, + client: String, +} + +impl Drop for VpnPortalSessionEventGuard { + fn drop(&mut self) { + self.events.emit(CoreEvent::VpnPortalClientDisconnected { + portal: self.portal.clone(), + client: self.client.clone(), + }); + } +} + +struct VpnPortalPeerPacketFilter { + clients: Arc>, +} + +#[async_trait] +impl PeerPacketFilter for VpnPortalPeerPacketFilter { + async fn try_process_packet_from_peer(&self, packet: ZCPacket) -> Option { + let client = match self.clients.route_peer_packet(&packet) { + VpnPortalPeerPacketRoute::Pass => return Some(packet), + VpnPortalPeerPacketRoute::Drop => return None, + VpnPortalPeerPacketRoute::Deliver { client, .. } => client, + }; + + let payload_offset = packet.payload_offset(); + let packet = + ZCPacket::new_from_buf(packet.inner().split_off(payload_offset), ZCPacketType::WG); + if let Err(error) = client.value().try_send(packet) { + tracing::debug!(?error, "failed to send packet to VPN portal client"); + } + None + } +} + +pub struct VpnPortalModule { + operation: Mutex<()>, + peer_manager: Arc, + runtime_config: CoreRuntimeConfigStore, + host: Option>, + events: Arc, + clients: Arc>, + runtime: Mutex>, +} + +impl VpnPortalModule { + pub fn new( + peer_manager: Arc, + runtime_config: CoreRuntimeConfigStore, + host: Option>, + events: Arc, + ) -> Arc { + Arc::new(Self { + operation: Mutex::new(()), + peer_manager, + runtime_config, + host, + events, + clients: Arc::new(VpnPortalClientTable::new()), + runtime: Mutex::new(None), + }) + } + + pub async fn start(&self) -> anyhow::Result<()> { + let _operation = self.operation.lock().await; + if self.runtime.lock().await.is_some() { + return Ok(()); + } + if self + .runtime_config + .snapshot() + .peer + .vpn_portal_cidr + .is_none() + { + return Ok(()); + } + let Some(host) = self.host.as_ref() else { + return Ok(()); + }; + + let listeners = host.start_listeners().await?; + if listeners.is_empty() { + anyhow::bail!("VPN portal host returned no active listeners"); + } + + let cancel = CancellationToken::new(); + let mut tasks = JoinSet::new(); + let mut listener_urls = Vec::with_capacity(listeners.len()); + for listener in listeners { + let local_url = listener.local_url(); + listener_urls.push(local_url); + tasks.spawn(Self::run_listener( + listener, + self.peer_manager.clone(), + self.clients.clone(), + self.events.clone(), + cancel.clone(), + )); + } + let pipeline = self + .peer_manager + .add_managed_packet_process_pipeline(Box::new(VpnPortalPeerPacketFilter { + clients: self.clients.clone(), + })) + .await; + *self.runtime.lock().await = Some(VpnPortalRuntime { + cancel, + tasks, + listener_urls: listener_urls.clone(), + _pipeline: pipeline, + }); + for local_url in listener_urls { + self.events + .emit(CoreEvent::VpnPortalStarted(local_url.to_string())); + } + Ok(()) + } + + async fn run_listener( + mut listener: VpnPortalListener, + peer_manager: Arc, + clients: Arc>, + events: Arc, + cancel: CancellationToken, + ) { + let mut sessions = JoinSet::new(); + let mut accepting = true; + loop { + while sessions.try_join_next().is_some() {} + if !accepting && sessions.is_empty() { + break; + } + tokio::select! { + _ = cancel.cancelled() => { + sessions.shutdown().await; + return; + }, + accepted = listener.accept(), if accepting => { + match accepted { + Ok(tunnel) => { + sessions.spawn(Self::run_session( + tunnel, + peer_manager.clone(), + clients.clone(), + events.clone(), + )); + } + Err(error) => { + tracing::warn!(?error, "VPN portal listener stopped accepting"); + accepting = false; + } + } + } + _ = sessions.join_next(), if !sessions.is_empty() => {} + } + } + } + + async fn run_session( + tunnel: Box, + peer_manager: Arc, + clients: Arc>, + events: Arc, + ) { + let info = tunnel.info().unwrap_or_default(); + let portal = info.local_addr.clone().unwrap_or_default().to_string(); + let client = info.remote_addr.clone().unwrap_or_default().to_string(); + let endpoint = info.remote_addr.clone().map(Into::into); + let mut tunnel = MpscTunnel::new(tunnel, None); + let mut stream = tunnel.get_stream(); + + events.emit(CoreEvent::VpnPortalClientConnected { + portal: portal.clone(), + client: client.clone(), + }); + let _event_guard = VpnPortalSessionEventGuard { + events, + portal, + client, + }; + let mut session = VpnPortalClientSession::new(clients, endpoint, tunnel.get_sink()); + loop { + let message = match stream.next().await { + Some(Ok(message)) => message, + Some(Err(error)) => { + tracing::error!(?error, "failed to receive from VPN portal client"); + break; + } + None => break, + }; + + assert_eq!(message.packet_type(), ZCPacketType::WG); + let payload = message.inner(); + let Some(packet) = session.observe_ipv4_payload(&payload) else { + tracing::error!(?payload, "failed to parse VPN portal IPv4 packet"); + continue; + }; + let _ = peer_manager + .send_msg_by_ip( + ZCPacket::new_with_payload(&payload), + IpAddr::V4(packet.destination), + false, + ) + .await; + } + + match session.close() { + VpnPortalClientRemoval::Removed(address) => { + tracing::info!(?address, "removed VPN portal client from table") + } + VpnPortalClientRemoval::EntryChangedOrMissing(address) => tracing::info!( + ?address, + "VPN portal client endpoint changed; retaining replacement" + ), + VpnPortalClientRemoval::NotRegistered => {} + } + } + + pub async fn stop(&self) { + let _operation = self.operation.lock().await; + let Some(mut runtime) = self.runtime.lock().await.take() else { + return; + }; + runtime.cancel.cancel(); + while runtime.tasks.join_next().await.is_some() {} + self.clients.entries.clear(); + } + + pub async fn info_snapshot(&self) -> VpnPortalInfoSnapshot { + let Some(host) = self.host.as_ref() else { + return VpnPortalInfoSnapshot { + vpn_type: "null".to_owned(), + client_config: String::new(), + connected_clients: Vec::new(), + }; + }; + let runtime = self.runtime.lock().await; + let started = runtime.is_some(); + let listener_url = runtime + .as_ref() + .and_then(|runtime| runtime.listener_urls.first().cloned()); + drop(runtime); + let plan = match listener_url { + Some(listener_url) => self.client_config_plan(listener_url).await, + None => None, + }; + VpnPortalInfoSnapshot { + vpn_type: host.name(), + client_config: if started { + plan.as_ref() + .map_or_else(String::new, |plan| host.render_client_config(plan)) + } else { + host.not_started_client_config() + }, + connected_clients: self + .clients + .endpoint_addrs() + .into_iter() + .map(|endpoint| endpoint.map(|url| url.to_string()).unwrap_or_default()) + .collect(), + } + } + + async fn client_config_plan( + &self, + listener_url: url::Url, + ) -> Option { + let config = self.runtime_config.snapshot(); + let client_cidr = config.peer.vpn_portal_cidr?; + let routes = self.peer_manager.list_route_snapshots().await; + let mut allowed_ips = routes + .iter() + .flat_map(|route| route.proxy_cidrs.iter().cloned()) + .collect::>(); + let local_ipv4 = config + .peer + .runtime + .core + .routes + .ipv4 + .as_ref() + .and_then(|prefix| { + let IpAddr::V4(address) = prefix.address else { + return None; + }; + Ipv4Inet::new(address, prefix.prefix_len).ok() + }); + if let Some(ipv4) = routes + .iter() + .filter_map(|route| route.ipv4_addr.map(Into::into)) + .chain(local_ipv4) + .next() + { + allowed_ips.push(ipv4.network().to_string()); + } + allowed_ips.push(client_cidr.to_string()); + Some(VpnPortalClientConfigPlan { + client_cidr, + allowed_ips, + listener_url, + }) + } +} + +impl VpnPortalClientSession { + pub fn new( + table: Arc>, + endpoint_addr: Option, + value: V, + ) -> Self { + Self { + table, + client: Arc::new(VpnPortalClient { + endpoint_addr, + value, + }), + registered_ip: None, + } + } + + pub fn observe_ipv4_payload(&mut self, payload: &[u8]) -> Option { + if payload.len() < IPV4_HEADER_LEN { + return None; + } + let packet = VpnPortalClientPacket { + source: ipv4_address(&payload[12..16]), + destination: ipv4_address(&payload[16..20]), + }; + + if self.registered_ip.is_none() { + self.table.insert(packet.source, self.client.clone()); + self.registered_ip = Some(packet.source); + } + Some(packet) + } + + pub fn registered_ip(&self) -> Option { + self.registered_ip + } + + pub fn close(&mut self) -> VpnPortalClientRemoval { + let Some(address) = self.registered_ip.take() else { + return VpnPortalClientRemoval::NotRegistered; + }; + if self.table.remove_if_current(&address, &self.client) { + VpnPortalClientRemoval::Removed(address) + } else { + VpnPortalClientRemoval::EntryChangedOrMissing(address) + } + } +} + +impl Drop for VpnPortalClientSession { + fn drop(&mut self) { + let _ = self.close(); + } +} + +fn ipv4_address(bytes: &[u8]) -> Ipv4Addr { + Ipv4Addr::new(bytes[0], bytes[1], bytes[2], bytes[3]) +} + +#[cfg(test)] +mod tests { + use super::*; + + fn ipv4_payload(source: [u8; 4], destination: [u8; 4], version: u8) -> Vec { + let mut payload = vec![0u8; IPV4_HEADER_LEN]; + payload[0] = version << 4 | 5; + payload[12..16].copy_from_slice(&source); + payload[16..20].copy_from_slice(&destination); + payload + } + + fn peer_packet(payload: &[u8], packet_type: PacketType) -> ZCPacket { + let mut packet = ZCPacket::new_with_payload(payload); + packet.fill_peer_manager_hdr(1, 2, packet_type as u8); + packet + } + + #[test] + fn session_registers_first_source_and_routes_peer_packet() { + let table = Arc::new(VpnPortalClientTable::new()); + let endpoint = Some("wg://198.51.100.2:51820".parse().unwrap()); + let mut session = VpnPortalClientSession::new(table.clone(), endpoint, "client"); + + let observed = session + .observe_ipv4_payload(&ipv4_payload([10, 10, 0, 2], [10, 10, 0, 3], 4)) + .unwrap(); + assert_eq!(observed.source, Ipv4Addr::new(10, 10, 0, 2)); + assert_eq!(session.registered_ip(), Some(observed.source)); + + let packet = peer_packet( + &ipv4_payload([10, 10, 0, 3], [10, 10, 0, 2], 4), + PacketType::Data, + ); + let VpnPortalPeerPacketRoute::Deliver { + destination, + client, + } = table.route_peer_packet(&packet) + else { + panic!("registered destination must be delivered"); + }; + assert_eq!(destination, observed.source); + assert_eq!(client.value(), &"client"); + } + + #[test] + fn closing_old_endpoint_does_not_remove_replacement() { + let table = Arc::new(VpnPortalClientTable::new()); + let payload = ipv4_payload([10, 10, 0, 2], [10, 10, 0, 3], 4); + let mut old = VpnPortalClientSession::new( + table.clone(), + Some("wg://198.51.100.2:51820".parse().unwrap()), + "old", + ); + let mut replacement = VpnPortalClientSession::new( + table.clone(), + Some("wg://198.51.100.3:51820".parse().unwrap()), + "replacement", + ); + old.observe_ipv4_payload(&payload).unwrap(); + replacement.observe_ipv4_payload(&payload).unwrap(); + + assert_eq!( + old.close(), + VpnPortalClientRemoval::EntryChangedOrMissing(Ipv4Addr::new(10, 10, 0, 2)) + ); + assert_eq!(table.len(), 1); + + let routed = peer_packet( + &ipv4_payload([10, 10, 0, 3], [10, 10, 0, 2], 4), + PacketType::Data, + ); + let VpnPortalPeerPacketRoute::Deliver { client, .. } = table.route_peer_packet(&routed) + else { + panic!("replacement must remain registered"); + }; + assert_eq!(client.value(), &"replacement"); + } + + #[test] + fn dropping_old_session_does_not_remove_same_endpoint_replacement() { + let table = Arc::new(VpnPortalClientTable::new()); + let endpoint = Some("wg://198.51.100.2:51820".parse().unwrap()); + let payload = ipv4_payload([10, 10, 0, 2], [10, 10, 0, 3], 4); + let mut old = VpnPortalClientSession::new(table.clone(), endpoint.clone(), "old"); + let mut replacement = VpnPortalClientSession::new(table.clone(), endpoint, "replacement"); + old.observe_ipv4_payload(&payload).unwrap(); + replacement.observe_ipv4_payload(&payload).unwrap(); + + drop(old); + + let routed = peer_packet( + &ipv4_payload([10, 10, 0, 3], [10, 10, 0, 2], 4), + PacketType::Data, + ); + let VpnPortalPeerPacketRoute::Deliver { client, .. } = table.route_peer_packet(&routed) + else { + panic!("same-endpoint replacement must remain registered"); + }; + assert_eq!(client.value(), &"replacement"); + } + + #[test] + fn close_removes_matching_entry_and_non_data_packets_pass() { + let table = Arc::new(VpnPortalClientTable::new()); + let mut session = VpnPortalClientSession::new(table.clone(), None, ()); + session + .observe_ipv4_payload(&ipv4_payload([10, 10, 0, 2], [10, 10, 0, 3], 4)) + .unwrap(); + + let non_data = peer_packet( + &ipv4_payload([10, 10, 0, 3], [10, 10, 0, 2], 4), + PacketType::Ping, + ); + assert!(matches!( + table.route_peer_packet(&non_data), + VpnPortalPeerPacketRoute::Pass + )); + assert_eq!( + session.close(), + VpnPortalClientRemoval::Removed(Ipv4Addr::new(10, 10, 0, 2)) + ); + assert!(table.is_empty()); + } + + #[test] + fn dropping_session_removes_matching_entry() { + let table = Arc::new(VpnPortalClientTable::new()); + { + let mut session = VpnPortalClientSession::new(table.clone(), None, ()); + session + .observe_ipv4_payload(&ipv4_payload([10, 10, 0, 2], [10, 10, 0, 3], 4)) + .unwrap(); + assert_eq!(table.len(), 1); + } + assert!(table.is_empty()); + } + + #[test] + fn peer_route_rejects_non_ipv4_payload() { + let table = VpnPortalClientTable::<()>::new(); + let packet = peer_packet( + &ipv4_payload([10, 10, 0, 3], [10, 10, 0, 2], 6), + PacketType::Data, + ); + assert!(matches!( + table.route_peer_packet(&packet), + VpnPortalPeerPacketRoute::Pass + )); + } + + #[test] + fn peer_route_drops_short_data_payload() { + let table = VpnPortalClientTable::<()>::new(); + let packet = peer_packet(&[0u8; IPV4_HEADER_LEN - 1], PacketType::Data); + assert!(matches!( + table.route_peer_packet(&packet), + VpnPortalPeerPacketRoute::Drop + )); + } +} diff --git a/easytier-core/src/host/dns.rs b/easytier-core/src/host/dns.rs new file mode 100644 index 00000000..cca66c09 --- /dev/null +++ b/easytier-core/src/host/dns.rs @@ -0,0 +1,425 @@ +use std::{io, net::IpAddr, sync::Arc, task::Poll}; + +use async_trait::async_trait; + +use crate::socket::SocketContext; + +use super::socket::{HostOperationId, HostSocketRuntime}; + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct DnsQuery { + pub host: String, + pub context: SocketContext, +} + +impl DnsQuery { + pub fn new(host: impl Into, context: SocketContext) -> Self { + Self { + host: host.into(), + context, + } + } +} + +#[async_trait] +pub trait DnsResolver: Send + Sync + 'static { + async fn resolve(&self, query: DnsQuery) -> anyhow::Result>; +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct DnsSrvRecord { + pub priority: u16, + pub weight: u16, + pub port: u16, + pub target: String, +} + +/// Resolves non-address DNS records used by EasyTier endpoint discovery. +#[async_trait] +pub trait DnsRecordResolver: Send + Sync + 'static { + async fn resolve_txt(&self, query: DnsQuery) -> anyhow::Result; + + async fn resolve_srv(&self, query: DnsQuery) -> anyhow::Result>; +} + +/// Mechanical asynchronous DNS below core's resolver seam. +/// +/// Submit methods must return without waiting for DNS. Completion methods own +/// their returned data and must not retain guest-memory borrows. Cancellation +/// removes pending or completed-but-unobserved operation state. +pub trait HostDnsIo: Send + Sync + 'static { + fn submit_resolve(&self, operation: HostOperationId, query: &DnsQuery) -> io::Result<()>; + + fn take_resolve(&self, operation: HostOperationId) -> Poll>>; + + fn submit_txt(&self, operation: HostOperationId, query: &DnsQuery) -> io::Result<()>; + + fn take_txt(&self, operation: HostOperationId) -> Poll>; + + fn submit_srv(&self, operation: HostOperationId, query: &DnsQuery) -> io::Result<()>; + + fn take_srv(&self, operation: HostOperationId) -> Poll>>; + + fn cancel_operation(&self, operation: HostOperationId) -> io::Result<()>; +} + +pub struct HostDnsResolver +where + D: HostDnsIo, +{ + runtime: HostSocketRuntime, + io: Arc, +} + +impl Clone for HostDnsResolver +where + D: HostDnsIo, +{ + fn clone(&self) -> Self { + Self { + runtime: self.runtime.clone(), + io: self.io.clone(), + } + } +} + +impl HostDnsResolver +where + D: HostDnsIo, +{ + pub fn new(runtime: HostSocketRuntime, io: Arc) -> Self { + Self { runtime, io } + } + + async fn run_operation( + &self, + submit: impl FnOnce(&D, HostOperationId) -> io::Result<()>, + take: impl Fn(&D, HostOperationId) -> Poll>, + ) -> io::Result { + self.runtime + .run_operation(self.io.clone(), submit, take, |io, operation| { + io.cancel_operation(operation) + }) + .await + } +} + +#[async_trait] +impl DnsResolver for HostDnsResolver +where + D: HostDnsIo, +{ + async fn resolve(&self, query: DnsQuery) -> anyhow::Result> { + Ok(self + .run_operation( + |io, operation| io.submit_resolve(operation, &query), + HostDnsIo::take_resolve, + ) + .await?) + } +} + +#[async_trait] +impl DnsRecordResolver for HostDnsResolver +where + D: HostDnsIo, +{ + async fn resolve_txt(&self, query: DnsQuery) -> anyhow::Result { + Ok(self + .run_operation( + |io, operation| io.submit_txt(operation, &query), + HostDnsIo::take_txt, + ) + .await?) + } + + async fn resolve_srv(&self, query: DnsQuery) -> anyhow::Result> { + Ok(self + .run_operation( + |io, operation| io.submit_srv(operation, &query), + HostDnsIo::take_srv, + ) + .await?) + } +} + +#[cfg(test)] +mod tests { + use std::{collections::HashMap, sync::Mutex}; + + use crate::socket::{IpVersion, SocketContext}; + + use super::*; + + enum TestDnsOperation { + Resolve { + query: DnsQuery, + result: Option>>, + }, + Txt { + query: DnsQuery, + result: Option>, + }, + Srv { + query: DnsQuery, + result: Option>>, + }, + } + + #[derive(Default)] + struct TestDnsIo { + operations: Mutex>, + cancelled: Mutex>, + } + + impl TestDnsIo { + fn operation( + &self, + predicate: impl Fn(&TestDnsOperation) -> bool, + ) -> (HostOperationId, DnsQuery) { + self.operations + .lock() + .unwrap() + .iter() + .find_map(|(operation, value)| { + if !predicate(value) { + return None; + } + let query = match value { + TestDnsOperation::Resolve { query, .. } + | TestDnsOperation::Txt { query, .. } + | TestDnsOperation::Srv { query, .. } => query, + }; + Some((*operation, query.clone())) + }) + .unwrap() + } + + fn complete_resolve(&self, operation: HostOperationId, addresses: Vec) { + let mut operations = self.operations.lock().unwrap(); + let TestDnsOperation::Resolve { result, .. } = operations.get_mut(&operation).unwrap() + else { + panic!("operation is not an address query"); + }; + *result = Some(Ok(addresses)); + } + + fn complete_txt(&self, operation: HostOperationId, text: String) { + let mut operations = self.operations.lock().unwrap(); + let TestDnsOperation::Txt { result, .. } = operations.get_mut(&operation).unwrap() + else { + panic!("operation is not a TXT query"); + }; + *result = Some(Ok(text)); + } + + fn complete_srv(&self, operation: HostOperationId, records: Vec) { + let mut operations = self.operations.lock().unwrap(); + let TestDnsOperation::Srv { result, .. } = operations.get_mut(&operation).unwrap() + else { + panic!("operation is not an SRV query"); + }; + *result = Some(Ok(records)); + } + } + + impl HostDnsIo for TestDnsIo { + fn submit_resolve(&self, operation: HostOperationId, query: &DnsQuery) -> io::Result<()> { + self.operations.lock().unwrap().insert( + operation, + TestDnsOperation::Resolve { + query: query.clone(), + result: None, + }, + ); + Ok(()) + } + + fn take_resolve(&self, operation: HostOperationId) -> Poll>> { + let mut operations = self.operations.lock().unwrap(); + let Some(TestDnsOperation::Resolve { result, .. }) = operations.get_mut(&operation) + else { + return Poll::Ready(Err(io::ErrorKind::NotFound.into())); + }; + let Some(result) = result.take() else { + return Poll::Pending; + }; + operations.remove(&operation); + Poll::Ready(result) + } + + fn submit_txt(&self, operation: HostOperationId, query: &DnsQuery) -> io::Result<()> { + self.operations.lock().unwrap().insert( + operation, + TestDnsOperation::Txt { + query: query.clone(), + result: None, + }, + ); + Ok(()) + } + + fn take_txt(&self, operation: HostOperationId) -> Poll> { + let mut operations = self.operations.lock().unwrap(); + let Some(TestDnsOperation::Txt { result, .. }) = operations.get_mut(&operation) else { + return Poll::Ready(Err(io::ErrorKind::NotFound.into())); + }; + let Some(result) = result.take() else { + return Poll::Pending; + }; + operations.remove(&operation); + Poll::Ready(result) + } + + fn submit_srv(&self, operation: HostOperationId, query: &DnsQuery) -> io::Result<()> { + self.operations.lock().unwrap().insert( + operation, + TestDnsOperation::Srv { + query: query.clone(), + result: None, + }, + ); + Ok(()) + } + + fn take_srv(&self, operation: HostOperationId) -> Poll>> { + let mut operations = self.operations.lock().unwrap(); + let Some(TestDnsOperation::Srv { result, .. }) = operations.get_mut(&operation) else { + return Poll::Ready(Err(io::ErrorKind::NotFound.into())); + }; + let Some(result) = result.take() else { + return Poll::Pending; + }; + operations.remove(&operation); + Poll::Ready(result) + } + + fn cancel_operation(&self, operation: HostOperationId) -> io::Result<()> { + self.operations.lock().unwrap().remove(&operation); + self.cancelled.lock().unwrap().push(operation); + Ok(()) + } + } + + fn query(host: &str, ip_version: IpVersion) -> DnsQuery { + DnsQuery::new( + host, + SocketContext { + ip_version, + socket_mark: Some(7), + netns: None, + }, + ) + } + + #[test] + fn query_keeps_host_and_socket_context() { + let context = SocketContext::default(); + assert_eq!( + DnsQuery::new("example.com", context.clone()), + DnsQuery { + host: "example.com".to_owned(), + context, + } + ); + } + + #[test] + fn srv_record_keeps_dns_selection_fields() { + let record = DnsSrvRecord { + priority: 10, + weight: 20, + port: 11010, + target: "peer.example.com.".to_owned(), + }; + + assert_eq!(record.priority, 10); + assert_eq!(record.weight, 20); + assert_eq!(record.port, 11010); + assert_eq!(record.target, "peer.example.com."); + } + + #[tokio::test] + async fn forwards_all_query_types_and_wakes_completed_futures() { + let io = Arc::new(TestDnsIo::default()); + let runtime = HostSocketRuntime::new(); + let resolver = HostDnsResolver::new(runtime.clone(), io.clone()); + + let address_query = query("peer.example", IpVersion::V4); + let address_task = tokio::spawn({ + let resolver = resolver.clone(); + let query = address_query.clone(); + async move { resolver.resolve(query).await } + }); + tokio::task::yield_now().await; + let (address_operation, submitted) = + io.operation(|value| matches!(value, TestDnsOperation::Resolve { .. })); + assert_eq!(submitted, address_query); + let addresses = vec!["192.0.2.1".parse().unwrap()]; + io.complete_resolve(address_operation, addresses.clone()); + runtime.notify_completions(); + assert_eq!(address_task.await.unwrap().unwrap(), addresses); + + let txt_query = query("_easytier.example", IpVersion::Both); + let mut txt = Box::pin(resolver.resolve_txt(txt_query.clone())); + assert!(futures::poll!(&mut txt).is_pending()); + let (txt_operation, submitted) = + io.operation(|value| matches!(value, TestDnsOperation::Txt { .. })); + assert_eq!(submitted, txt_query); + io.complete_txt(txt_operation, "tcp://peer.example:11010".to_owned()); + runtime.notify_completions(); + assert_eq!(txt.await.unwrap(), "tcp://peer.example:11010"); + + let srv_query = query("_easytier._udp.example", IpVersion::V6); + let mut srv = Box::pin(resolver.resolve_srv(srv_query.clone())); + assert!(futures::poll!(&mut srv).is_pending()); + let (srv_operation, submitted) = + io.operation(|value| matches!(value, TestDnsOperation::Srv { .. })); + assert_eq!(submitted, srv_query); + let records = vec![DnsSrvRecord { + priority: 10, + weight: 20, + port: 11010, + target: "peer.example.".to_owned(), + }]; + io.complete_srv(srv_operation, records.clone()); + runtime.notify_completions(); + assert_eq!(srv.await.unwrap(), records); + } + + #[tokio::test] + async fn dropping_pending_or_unobserved_completion_cancels_host_state() { + let io = Arc::new(TestDnsIo::default()); + let runtime = HostSocketRuntime::new(); + let resolver = HostDnsResolver::new(runtime.clone(), io.clone()); + + let pending_operation = { + let mut resolve = + Box::pin(resolver.resolve_txt(query("pending.example", IpVersion::Both))); + assert!(futures::poll!(&mut resolve).is_pending()); + assert_eq!(runtime.inner.wakers.len(), 1); + let (operation, _) = + io.operation(|value| matches!(value, TestDnsOperation::Txt { .. })); + drop(resolve); + assert_eq!(runtime.inner.wakers.len(), 0); + operation + }; + let completed_operation = { + let mut resolve = Box::pin(resolver.resolve(query("cancel.example", IpVersion::Both))); + assert!(futures::poll!(&mut resolve).is_pending()); + assert_eq!(runtime.inner.wakers.len(), 1); + let (operation, _) = + io.operation(|value| matches!(value, TestDnsOperation::Resolve { .. })); + io.complete_resolve(operation, vec!["192.0.2.2".parse().unwrap()]); + drop(resolve); + assert_eq!(runtime.inner.wakers.len(), 0); + operation + }; + + assert_eq!( + *io.cancelled.lock().unwrap(), + vec![pending_operation, completed_operation] + ); + assert!(io.operations.lock().unwrap().is_empty()); + } +} diff --git a/easytier-core/src/host/environment.rs b/easytier-core/src/host/environment.rs new file mode 100644 index 00000000..6f63cd68 --- /dev/null +++ b/easytier-core/src/host/environment.rs @@ -0,0 +1,231 @@ +//! Host-operation bridge for connector environment queries. + +use std::{io, net::SocketAddr, sync::Arc, task::Poll}; + +use crate::socket::SocketContext; + +use super::socket::{HostOperationId, HostSocketRuntime}; + +/// Mechanical asynchronous environment operations below connector policy. +/// +/// Submit calls must take ownership of their complete input before returning +/// and must leave no operation state when they return an error. A `Pending` +/// take retains the operation; a `Ready` take consumes both successful and +/// failed results. Cancellation must remove pending and completed-but-unread +/// state and is idempotent for an already-absent operation. +pub trait HostConnectorEnvironmentIo: Send + Sync + 'static { + fn submit_local_addr_for_remote( + &self, + operation: HostOperationId, + remote_addr: SocketAddr, + context: &SocketContext, + ) -> io::Result<()>; + + fn take_local_addr_for_remote( + &self, + operation: HostOperationId, + ) -> Poll>; + + fn cancel_operation(&self, operation: HostOperationId) -> io::Result<()>; +} + +pub(crate) async fn local_addr_for_remote( + runtime: &HostSocketRuntime, + io: Arc, + remote_addr: SocketAddr, + context: SocketContext, +) -> anyhow::Result +where + I: HostConnectorEnvironmentIo, +{ + Ok(runtime + .run_operation( + io, + |io, operation| io.submit_local_addr_for_remote(operation, remote_addr, &context), + HostConnectorEnvironmentIo::take_local_addr_for_remote, + |io, operation| io.cancel_operation(operation), + ) + .await?) +} + +#[cfg(test)] +mod tests { + use std::{collections::HashMap, sync::Mutex}; + + use super::*; + + #[derive(Debug, Clone, PartialEq, Eq)] + enum Request { + Local(SocketAddr, SocketContext), + } + + #[derive(Debug, Clone, Copy)] + enum Completion { + Ok(SocketAddr), + Err(io::ErrorKind), + } + + #[derive(Default)] + struct TestIo { + requests: Mutex)>>, + cancelled: Mutex>, + } + + impl TestIo { + fn submit(&self, operation: HostOperationId, request: Request) -> io::Result<()> { + let replaced = self + .requests + .lock() + .unwrap() + .insert(operation, (request, None)); + if replaced.is_some() { + return Err(io::ErrorKind::AlreadyExists.into()); + } + Ok(()) + } + + fn take(&self, operation: HostOperationId) -> Poll> { + let mut requests = self.requests.lock().unwrap(); + let Some((_, result)) = requests.get(&operation) else { + return Poll::Ready(Err(io::ErrorKind::NotFound.into())); + }; + let Some(result) = *result else { + return Poll::Pending; + }; + requests.remove(&operation); + Poll::Ready(match result { + Completion::Ok(address) => Ok(address), + Completion::Err(kind) => Err(kind.into()), + }) + } + + fn operation_for(&self, request: Request) -> HostOperationId { + self.requests + .lock() + .unwrap() + .iter() + .find_map(|(operation, (candidate, _))| { + (*candidate == request).then_some(*operation) + }) + .expect("submitted host environment operation") + } + + fn complete(&self, operation: HostOperationId, result: Completion) { + self.requests.lock().unwrap().get_mut(&operation).unwrap().1 = Some(result); + } + } + + impl HostConnectorEnvironmentIo for TestIo { + fn submit_local_addr_for_remote( + &self, + operation: HostOperationId, + remote_addr: SocketAddr, + context: &SocketContext, + ) -> io::Result<()> { + self.submit(operation, Request::Local(remote_addr, context.clone())) + } + + fn take_local_addr_for_remote( + &self, + operation: HostOperationId, + ) -> Poll> { + self.take(operation) + } + + fn cancel_operation(&self, operation: HostOperationId) -> io::Result<()> { + self.requests.lock().unwrap().remove(&operation); + self.cancelled.lock().unwrap().push(operation); + Ok(()) + } + } + + #[tokio::test] + async fn wakes_completed_operations_and_cancels_dropped_futures() { + let runtime = HostSocketRuntime::new(); + let io = Arc::new(TestIo::default()); + let remote = "203.0.113.1:11010".parse().unwrap(); + let context = SocketContext::default().with_socket_mark(Some(7)); + let task = tokio::spawn({ + let runtime = runtime.clone(); + let io = io.clone(); + let context = context.clone(); + async move { local_addr_for_remote(&runtime, io, remote, context).await } + }); + tokio::task::yield_now().await; + + let operation = io.operation_for(Request::Local(remote, context)); + io.complete( + operation, + Completion::Ok("192.0.2.1:40100".parse().unwrap()), + ); + runtime.notify_completions(); + assert_eq!( + task.await.unwrap().unwrap(), + "192.0.2.1:40100".parse().unwrap() + ); + + let pending_remote = "203.0.113.2:11010".parse().unwrap(); + let pending = tokio::spawn({ + let runtime = runtime.clone(); + let io = io.clone(); + async move { + local_addr_for_remote(&runtime, io, pending_remote, SocketContext::default()).await + } + }); + tokio::task::yield_now().await; + let cancelled = io.operation_for(Request::Local(pending_remote, SocketContext::default())); + pending.abort(); + let _ = pending.await; + assert_eq!(*io.cancelled.lock().unwrap(), vec![cancelled]); + assert!(!io.requests.lock().unwrap().contains_key(&cancelled)); + + let failed_remote = "203.0.113.3:11010".parse().unwrap(); + let failed = tokio::spawn({ + let runtime = runtime.clone(); + let io = io.clone(); + async move { + local_addr_for_remote(&runtime, io, failed_remote, SocketContext::default()).await + } + }); + tokio::task::yield_now().await; + let failed_operation = + io.operation_for(Request::Local(failed_remote, SocketContext::default())); + io.complete( + failed_operation, + Completion::Err(io::ErrorKind::AddrNotAvailable), + ); + runtime.notify_completions(); + assert_eq!( + failed + .await + .unwrap() + .unwrap_err() + .downcast_ref::() + .unwrap() + .kind(), + io::ErrorKind::AddrNotAvailable + ); + assert!(!io.requests.lock().unwrap().contains_key(&failed_operation)); + + let unread_remote = "203.0.113.4:11010".parse().unwrap(); + let unread = tokio::spawn({ + let runtime = runtime.clone(); + let io = io.clone(); + async move { + local_addr_for_remote(&runtime, io, unread_remote, SocketContext::default()).await + } + }); + tokio::task::yield_now().await; + let unread_operation = + io.operation_for(Request::Local(unread_remote, SocketContext::default())); + io.complete( + unread_operation, + Completion::Ok("198.51.100.1:44000".parse().unwrap()), + ); + runtime.notify_completions(); + unread.abort(); + let _ = unread.await; + assert!(io.cancelled.lock().unwrap().contains(&unread_operation)); + assert!(!io.requests.lock().unwrap().contains_key(&unread_operation)); + } +} diff --git a/easytier-core/src/host/mod.rs b/easytier-core/src/host/mod.rs new file mode 100644 index 00000000..8dd46e1a --- /dev/null +++ b/easytier-core/src/host/mod.rs @@ -0,0 +1,15 @@ +//! Host capability seams. +//! +//! This module is the single home of every Host capability seam (see +//! CONTEXT.md, "Host capability"): DNS resolution, connector environment +//! facts, packet egress, and host-backed sockets. The socket-flavoured seams +//! — the host operation runtime and identifiers, the TCP stream and UDP +//! socket bridges, socket factories, and TCP listeners — live in [`socket`]. +//! Concrete WASI adapters for these seams live in [`crate::wasi`]. + +pub mod dns; +pub mod environment; +pub mod packet; +pub mod socket; +#[cfg(test)] +pub(crate) mod testkit; diff --git a/easytier-core/src/host/packet.rs b/easytier-core/src/host/packet.rs new file mode 100644 index 00000000..85273e80 --- /dev/null +++ b/easytier-core/src/host/packet.rs @@ -0,0 +1,281 @@ +use std::{io, sync::Arc, task::Poll}; + +use async_trait::async_trait; +use tokio::sync::mpsc; + +use super::socket::{HostOperationId, HostSocketRuntime}; + +/// Receives raw IP packet bytes leaving the EasyTier peer graph. +/// +/// The host decides whether packets go to a TUN device, a Go callback, or a +/// different packet backend. Core's internal packet headers never cross this +/// boundary, and core never performs platform I/O directly. +#[async_trait] +pub trait PacketSink: Send + Sync + 'static { + async fn write_packet(&self, packet: Vec) -> anyhow::Result<()>; +} + +#[async_trait] +impl PacketSink for mpsc::Sender> { + async fn write_packet(&self, packet: Vec) -> anyhow::Result<()> { + self.send(packet) + .await + .map_err(|_| anyhow::anyhow!("packet sink channel is closed")) + } +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] +pub struct HostPacketSinkHandle(pub u64); + +/// Mechanical packet egress below core's packet scheduling seam. +/// +/// A successful `try_write_packet` owns a complete packet copy before it +/// returns. `WouldBlock` has no side effects. Readiness operations only report +/// that another admission attempt may succeed; they never accept a packet. +pub trait HostPacketIo: Send + Sync + 'static { + fn try_write_packet(&self, handle: HostPacketSinkHandle, packet: &[u8]) -> io::Result<()>; + + fn submit_write_ready( + &self, + handle: HostPacketSinkHandle, + operation: HostOperationId, + ) -> io::Result<()>; + + fn take_write_ready(&self, operation: HostOperationId) -> Poll>; + + fn cancel_operation(&self, operation: HostOperationId) -> io::Result<()>; +} + +pub struct HostPacketSink +where + I: HostPacketIo, +{ + runtime: HostSocketRuntime, + io: Arc, + handle: HostPacketSinkHandle, +} + +impl Clone for HostPacketSink +where + I: HostPacketIo, +{ + fn clone(&self) -> Self { + Self { + runtime: self.runtime.clone(), + io: self.io.clone(), + handle: self.handle, + } + } +} + +impl HostPacketSink +where + I: HostPacketIo, +{ + pub fn new(runtime: HostSocketRuntime, io: Arc, handle: HostPacketSinkHandle) -> Self { + Self { + runtime, + io, + handle, + } + } + + async fn wait_writable(&self) -> io::Result<()> { + self.runtime + .run_operation( + self.io.clone(), + |io, operation| io.submit_write_ready(self.handle, operation), + |io, operation| io.take_write_ready(operation), + |io, operation| io.cancel_operation(operation), + ) + .await + } +} + +#[async_trait] +impl PacketSink for HostPacketSink +where + I: HostPacketIo, +{ + async fn write_packet(&self, packet: Vec) -> anyhow::Result<()> { + loop { + match self.io.try_write_packet(self.handle, &packet) { + Ok(()) => return Ok(()), + Err(error) if error.kind() == io::ErrorKind::WouldBlock => { + self.wait_writable().await?; + } + Err(error) => return Err(error.into()), + } + } + } +} + +#[cfg(test)] +mod tests { + use std::{collections::HashMap, sync::Mutex}; + + use super::*; + + #[derive(Default)] + struct TestPacketState { + writable: bool, + packets: Vec<(HostPacketSinkHandle, Vec)>, + waiters: HashMap, + cancelled: Vec, + } + + #[derive(Default)] + struct TestPacketIo { + state: Mutex, + } + + impl TestPacketIo { + fn set_writable(&self) { + let mut state = self.state.lock().unwrap(); + state.writable = true; + for ready in state.waiters.values_mut() { + *ready = true; + } + } + + fn waiter(&self) -> HostOperationId { + *self.state.lock().unwrap().waiters.keys().next().unwrap() + } + } + + impl HostPacketIo for TestPacketIo { + fn try_write_packet(&self, handle: HostPacketSinkHandle, packet: &[u8]) -> io::Result<()> { + let mut state = self.state.lock().unwrap(); + if !state.writable { + return Err(io::ErrorKind::WouldBlock.into()); + } + state.writable = false; + state.packets.push((handle, packet.to_vec())); + Ok(()) + } + + fn submit_write_ready( + &self, + _handle: HostPacketSinkHandle, + operation: HostOperationId, + ) -> io::Result<()> { + let mut state = self.state.lock().unwrap(); + let ready = state.writable; + state.waiters.insert(operation, ready); + Ok(()) + } + + fn take_write_ready(&self, operation: HostOperationId) -> Poll> { + let mut state = self.state.lock().unwrap(); + match state.waiters.get(&operation) { + Some(true) => { + state.waiters.remove(&operation); + Poll::Ready(Ok(())) + } + Some(false) => Poll::Pending, + None => Poll::Ready(Err(io::ErrorKind::NotFound.into())), + } + } + + fn cancel_operation(&self, operation: HostOperationId) -> io::Result<()> { + let mut state = self.state.lock().unwrap(); + state.waiters.remove(&operation); + state.cancelled.push(operation); + Ok(()) + } + } + + fn test_sink( + writable: bool, + ) -> ( + HostSocketRuntime, + Arc, + HostPacketSink, + ) { + let runtime = HostSocketRuntime::new(); + let io = Arc::new(TestPacketIo::default()); + io.state.lock().unwrap().writable = writable; + let sink = HostPacketSink::new(runtime.clone(), io.clone(), HostPacketSinkHandle(41)); + (runtime, io, sink) + } + + #[tokio::test] + async fn admits_complete_packet_without_readiness_wait() { + let (_runtime, io, sink) = test_sink(true); + sink.write_packet(vec![1, 2, 3, 4]).await.unwrap(); + + let state = io.state.lock().unwrap(); + assert_eq!( + state.packets, + vec![(HostPacketSinkHandle(41), vec![1, 2, 3, 4])] + ); + assert!(state.waiters.is_empty()); + } + + #[tokio::test] + async fn waits_for_capacity_then_admits_packet_once() { + let (runtime, io, sink) = test_sink(false); + let task = tokio::spawn(async move { sink.write_packet(vec![5, 6, 7]).await }); + tokio::task::yield_now().await; + assert!(io.state.lock().unwrap().packets.is_empty()); + assert_eq!(runtime.inner.wakers.len(), 1); + + io.set_writable(); + runtime.notify_completions(); + task.await.unwrap().unwrap(); + + let state = io.state.lock().unwrap(); + assert_eq!( + state.packets, + vec![(HostPacketSinkHandle(41), vec![5, 6, 7])] + ); + assert!(state.waiters.is_empty()); + assert!(state.cancelled.is_empty()); + assert_eq!(runtime.inner.wakers.len(), 0); + } + + #[tokio::test] + async fn dropping_pending_waiter_removes_waker_and_host_state() { + let (runtime, io, sink) = test_sink(false); + let operation = { + let mut write = Box::pin(sink.write_packet(vec![7, 8])); + assert!(futures::poll!(&mut write).is_pending()); + assert_eq!(runtime.inner.wakers.len(), 1); + let operation = io.waiter(); + drop(write); + operation + }; + + let state = io.state.lock().unwrap(); + assert!(state.packets.is_empty()); + assert!(state.waiters.is_empty()); + assert_eq!(state.cancelled, vec![operation]); + assert_eq!(runtime.inner.wakers.len(), 0); + } + + #[tokio::test] + async fn dropping_ready_waiter_does_not_admit_packet() { + let (runtime, io, sink) = test_sink(false); + let operation = { + let mut write = Box::pin(sink.write_packet(vec![8, 9])); + assert!(futures::poll!(&mut write).is_pending()); + let operation = io.waiter(); + io.set_writable(); + runtime.notify_completions(); + drop(write); + operation + }; + + { + let state = io.state.lock().unwrap(); + assert!(state.packets.is_empty()); + assert!(state.waiters.is_empty()); + assert_eq!(state.cancelled, vec![operation]); + } + sink.write_packet(vec![8, 9]).await.unwrap(); + assert_eq!( + io.state.lock().unwrap().packets, + vec![(HostPacketSinkHandle(41), vec![8, 9])] + ); + } +} diff --git a/easytier-core/src/host/socket/factory.rs b/easytier-core/src/host/socket/factory.rs new file mode 100644 index 00000000..b1924731 --- /dev/null +++ b/easytier-core/src/host/socket/factory.rs @@ -0,0 +1,503 @@ +use std::{io, net::SocketAddr, sync::Arc, task::Poll}; + +use crate::socket::{ + tcp::{TcpConnectOptions, VirtualTcpSocketFactory}, + udp::{UdpBindOptions, VirtualUdpSocketFactory}, +}; +use async_trait::async_trait; + +use super::{ + HostOperationId, HostSocketHandle, HostSocketIo, HostSocketRuntime, HostTcpIo, HostTcpStream, + udp::{HostUdpIo, HostUdpSocket}, +}; + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct HostTcpConnectResult { + pub handle: HostSocketHandle, + pub local_addr: SocketAddr, + pub peer_addr: SocketAddr, + pub transport_label: Option, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct HostUdpBindResult { + pub handle: HostSocketHandle, + pub local_addr: SocketAddr, +} + +/// Mechanical creation of host-owned socket resources. +/// +/// Submit methods start work without blocking the guest. Completion results own +/// their handles and address metadata. Canceling an operation must atomically +/// stop pending creation or close and discard a resource that completed before +/// core observed it, so dropping a factory future cannot leak a host socket. +pub trait HostSocketFactoryIo: HostSocketIo { + fn submit_tcp_connect( + &self, + operation: HostOperationId, + options: &TcpConnectOptions, + ) -> io::Result<()>; + + fn take_tcp_connect( + &self, + operation: HostOperationId, + ) -> Poll>; + + fn submit_udp_bind( + &self, + operation: HostOperationId, + options: &UdpBindOptions, + ) -> io::Result<()>; + + fn take_udp_bind(&self, operation: HostOperationId) -> Poll>; +} + +pub trait HostSocketBackend: HostSocketFactoryIo + HostTcpIo + HostUdpIo {} + +impl HostSocketBackend for T where T: HostSocketFactoryIo + HostTcpIo + HostUdpIo {} + +pub struct HostSocketFactory +where + B: HostSocketBackend, +{ + runtime: HostSocketRuntime, + backend: Arc, +} + +impl Clone for HostSocketFactory +where + B: HostSocketBackend, +{ + fn clone(&self) -> Self { + Self { + runtime: self.runtime.clone(), + backend: self.backend.clone(), + } + } +} + +impl HostSocketFactory +where + B: HostSocketBackend, +{ + pub fn new(runtime: HostSocketRuntime, backend: Arc) -> Self { + Self { runtime, backend } + } + + async fn connect_tcp(&self, options: TcpConnectOptions) -> io::Result { + let result = self + .runtime + .run_operation( + self.backend.clone(), + |backend, operation| backend.submit_tcp_connect(operation, &options), + |backend, operation| backend.take_tcp_connect(operation), + |backend, operation| backend.cancel_operation(operation), + ) + .await?; + Ok(self.runtime.tcp_stream( + self.backend.clone(), + result.handle, + result.local_addr, + result.peer_addr, + result.transport_label, + )) + } + + async fn bind_udp(&self, options: UdpBindOptions) -> io::Result> { + let context = options.context.clone(); + let result = self + .runtime + .run_operation( + self.backend.clone(), + |backend, operation| backend.submit_udp_bind(operation, &options), + |backend, operation| backend.take_udp_bind(operation), + |backend, operation| backend.cancel_operation(operation), + ) + .await?; + Ok(Arc::new(self.runtime.udp_socket_with_context( + self.backend.clone(), + result.handle, + result.local_addr, + context, + ))) + } +} + +#[async_trait] +impl VirtualTcpSocketFactory for HostSocketFactory +where + B: HostSocketBackend, +{ + type Socket = HostTcpStream; + + async fn connect_tcp(&self, options: TcpConnectOptions) -> anyhow::Result { + Ok(self.connect_tcp(options).await?) + } +} + +#[async_trait] +impl VirtualUdpSocketFactory for HostSocketFactory +where + B: HostSocketBackend, +{ + type Socket = HostUdpSocket; + + async fn bind_udp(&self, options: UdpBindOptions) -> anyhow::Result> { + Ok(self.bind_udp(options).await?) + } +} + +#[cfg(test)] +mod tests { + use std::{ + collections::{HashMap, HashSet}, + sync::Mutex, + }; + + use crate::socket::{ + tcp::{TcpBindOptions, TcpSocketPurpose, VirtualTcpSocket}, + udp::{UdpSocketPurpose, UdpSocketSendMeta, VirtualUdpSocket}, + }; + + use super::*; + + enum TestCreate { + Tcp { + options: TcpConnectOptions, + result: Option>, + }, + Udp { + options: UdpBindOptions, + result: Option>, + }, + } + + #[derive(Default)] + struct TestHostIo { + creates: Mutex>, + cancelled: Mutex>, + closed: Mutex>, + } + + impl TestHostIo { + fn operation(&self, tcp: bool) -> HostOperationId { + self.creates + .lock() + .unwrap() + .iter() + .find_map(|(operation, create)| match (tcp, create) { + (true, TestCreate::Tcp { .. }) | (false, TestCreate::Udp { .. }) => { + Some(*operation) + } + _ => None, + }) + .unwrap() + } + + fn complete_tcp(&self, operation: HostOperationId, result: HostTcpConnectResult) { + let mut creates = self.creates.lock().unwrap(); + let TestCreate::Tcp { + result: completion, .. + } = creates.get_mut(&operation).unwrap() + else { + panic!("operation is not TCP connect"); + }; + *completion = Some(Ok(result)); + } + + fn complete_udp(&self, operation: HostOperationId, result: HostUdpBindResult) { + let mut creates = self.creates.lock().unwrap(); + let TestCreate::Udp { + result: completion, .. + } = creates.get_mut(&operation).unwrap() + else { + panic!("operation is not UDP bind"); + }; + *completion = Some(Ok(result)); + } + } + + impl HostSocketIo for TestHostIo { + fn cancel_operation(&self, operation: HostOperationId) -> io::Result<()> { + if let Some(create) = self.creates.lock().unwrap().remove(&operation) { + let handle = match create { + TestCreate::Tcp { + result: Some(Ok(result)), + .. + } => Some(result.handle), + TestCreate::Udp { + result: Some(Ok(result)), + .. + } => Some(result.handle), + _ => None, + }; + if let Some(handle) = handle { + self.closed.lock().unwrap().insert(handle); + } + } + self.cancelled.lock().unwrap().push(operation); + Ok(()) + } + + fn close(&self, handle: HostSocketHandle) -> io::Result<()> { + self.closed.lock().unwrap().insert(handle); + Ok(()) + } + } + + impl HostSocketFactoryIo for TestHostIo { + fn submit_tcp_connect( + &self, + operation: HostOperationId, + options: &TcpConnectOptions, + ) -> io::Result<()> { + self.creates.lock().unwrap().insert( + operation, + TestCreate::Tcp { + options: options.clone(), + result: None, + }, + ); + Ok(()) + } + + fn take_tcp_connect( + &self, + operation: HostOperationId, + ) -> Poll> { + let mut creates = self.creates.lock().unwrap(); + let Some(TestCreate::Tcp { result, .. }) = creates.get_mut(&operation) else { + return Poll::Ready(Err(io::ErrorKind::NotFound.into())); + }; + let Some(result) = result.take() else { + return Poll::Pending; + }; + creates.remove(&operation); + Poll::Ready(result) + } + + fn submit_udp_bind( + &self, + operation: HostOperationId, + options: &UdpBindOptions, + ) -> io::Result<()> { + self.creates.lock().unwrap().insert( + operation, + TestCreate::Udp { + options: options.clone(), + result: None, + }, + ); + Ok(()) + } + + fn take_udp_bind(&self, operation: HostOperationId) -> Poll> { + let mut creates = self.creates.lock().unwrap(); + let Some(TestCreate::Udp { result, .. }) = creates.get_mut(&operation) else { + return Poll::Ready(Err(io::ErrorKind::NotFound.into())); + }; + let Some(result) = result.take() else { + return Poll::Pending; + }; + creates.remove(&operation); + Poll::Ready(result) + } + } + + impl HostTcpIo for TestHostIo { + fn submit_read( + &self, + _handle: HostSocketHandle, + _operation: HostOperationId, + _capacity: usize, + ) -> io::Result<()> { + Err(io::ErrorKind::Unsupported.into()) + } + + fn take_read(&self, _operation: HostOperationId) -> Poll>> { + Poll::Ready(Err(io::ErrorKind::Unsupported.into())) + } + + fn submit_write( + &self, + _handle: HostSocketHandle, + _operation: HostOperationId, + _source: &[u8], + ) -> io::Result<()> { + Err(io::ErrorKind::Unsupported.into()) + } + + fn take_write(&self, _operation: HostOperationId) -> Poll> { + Poll::Ready(Err(io::ErrorKind::Unsupported.into())) + } + } + + impl HostUdpIo for TestHostIo { + fn submit_recv( + &self, + _handle: HostSocketHandle, + _operation: HostOperationId, + _capacity: usize, + ) -> io::Result<()> { + Err(io::ErrorKind::Unsupported.into()) + } + + fn take_recv( + &self, + _operation: HostOperationId, + ) -> Poll> { + Poll::Ready(Err(io::ErrorKind::Unsupported.into())) + } + + fn try_send( + &self, + _handle: HostSocketHandle, + _source: &[u8], + _peer_addr: SocketAddr, + _meta: UdpSocketSendMeta, + ) -> io::Result<()> { + Err(io::ErrorKind::Unsupported.into()) + } + + fn submit_send_ready( + &self, + _handle: HostSocketHandle, + _operation: HostOperationId, + ) -> io::Result<()> { + Err(io::ErrorKind::Unsupported.into()) + } + + fn take_send_ready(&self, _operation: HostOperationId) -> Poll> { + Poll::Ready(Err(io::ErrorKind::Unsupported.into())) + } + } + + fn test_factory(io: Arc) -> (HostSocketRuntime, HostSocketFactory) { + let runtime = HostSocketRuntime::new(); + let factory = HostSocketFactory::new(runtime.clone(), io); + (runtime, factory) + } + + #[tokio::test] + async fn forwards_tcp_connect_options_and_wraps_completed_handle() { + let io = Arc::new(TestHostIo::default()); + let (runtime, factory) = test_factory(io.clone()); + let options = TcpConnectOptions { + remote_addr: "192.0.2.2:11013".parse().unwrap(), + bind: TcpBindOptions::default() + .with_local_addr(Some("192.0.2.1:0".parse().unwrap())) + .with_socket_mark(Some(7)) + .with_bind_device(Some("host-device".to_owned())) + .with_reuse_port(true), + purpose: TcpSocketPurpose::ManualConnect, + }; + let task = tokio::spawn({ + let factory = factory.clone(); + let options = options.clone(); + async move { VirtualTcpSocketFactory::connect_tcp(&factory, options).await } + }); + tokio::task::yield_now().await; + let operation = io.operation(true); + { + let creates = io.creates.lock().unwrap(); + let TestCreate::Tcp { + options: submitted, .. + } = creates.get(&operation).unwrap() + else { + panic!("operation is not TCP connect"); + }; + assert_eq!(submitted, &options); + } + + io.complete_tcp( + operation, + HostTcpConnectResult { + handle: HostSocketHandle(41), + local_addr: "192.0.2.1:40100".parse().unwrap(), + peer_addr: options.remote_addr, + transport_label: Some("host-tcp".to_owned()), + }, + ); + runtime.notify_completions(); + let stream = task.await.unwrap().unwrap(); + assert_eq!( + stream.local_addr().unwrap(), + "192.0.2.1:40100".parse().unwrap() + ); + assert_eq!(stream.peer_addr().unwrap(), options.remote_addr); + assert_eq!(stream.transport_label(), Some("host-tcp")); + drop(stream); + assert!(io.closed.lock().unwrap().contains(&HostSocketHandle(41))); + } + + #[tokio::test] + async fn forwards_udp_bind_options_and_wraps_completed_handle() { + let io = Arc::new(TestHostIo::default()); + let (runtime, factory) = test_factory(io.clone()); + let options = UdpBindOptions { + context: crate::socket::SocketContext::default().with_socket_mark(Some(9)), + local_addr: Some("[::]:11013".parse().unwrap()), + bind_device: Some("host-device".to_owned()), + reuse_addr: true, + reuse_port: true, + only_v6: true, + purpose: UdpSocketPurpose::PortBoundListener, + }; + let task = tokio::spawn({ + let factory = factory.clone(); + let options = options.clone(); + async move { VirtualUdpSocketFactory::bind_udp(&factory, options).await } + }); + tokio::task::yield_now().await; + let operation = io.operation(false); + { + let creates = io.creates.lock().unwrap(); + let TestCreate::Udp { + options: submitted, .. + } = creates.get(&operation).unwrap() + else { + panic!("operation is not UDP bind"); + }; + assert_eq!(submitted, &options); + } + + io.complete_udp( + operation, + HostUdpBindResult { + handle: HostSocketHandle(42), + local_addr: "[::]:11013".parse().unwrap(), + }, + ); + runtime.notify_completions(); + let socket = task.await.unwrap().unwrap(); + assert_eq!(socket.local_addr().unwrap(), "[::]:11013".parse().unwrap()); + drop(socket); + assert!(io.closed.lock().unwrap().contains(&HostSocketHandle(42))); + } + + #[tokio::test] + async fn cancelling_completed_create_closes_unobserved_handle() { + let io = Arc::new(TestHostIo::default()); + let (_runtime, factory) = test_factory(io.clone()); + let options = TcpConnectOptions::direct_connect("192.0.2.2:11013".parse().unwrap()); + let mut connect = Box::pin(VirtualTcpSocketFactory::connect_tcp( + &factory, + options.clone(), + )); + assert!(futures::poll!(&mut connect).is_pending()); + let operation = io.operation(true); + io.complete_tcp( + operation, + HostTcpConnectResult { + handle: HostSocketHandle(43), + local_addr: "192.0.2.1:40101".parse().unwrap(), + peer_addr: options.remote_addr, + transport_label: None, + }, + ); + drop(connect); + + assert_eq!(*io.cancelled.lock().unwrap(), vec![operation]); + assert!(io.closed.lock().unwrap().contains(&HostSocketHandle(43))); + } +} diff --git a/easytier-core/src/host/socket/listener.rs b/easytier-core/src/host/socket/listener.rs new file mode 100644 index 00000000..62053ad8 --- /dev/null +++ b/easytier-core/src/host/socket/listener.rs @@ -0,0 +1,462 @@ +use std::{io, net::SocketAddr, sync::Arc, task::Poll}; + +use crate::socket::tcp::{TcpListenOptions, VirtualTcpListener, VirtualTcpListenerFactory}; +use async_trait::async_trait; + +use super::{ + HostOperationId, HostSocketHandle, HostSocketIo, HostSocketRuntime, HostTcpIo, HostTcpStream, + factory::HostTcpConnectResult, +}; + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct HostTcpBindResult { + pub handle: HostSocketHandle, + pub local_addr: SocketAddr, +} + +/// Mechanical TCP listener creation and accept readiness. +/// +/// Bind cancellation must close a completed but unobserved listener. Accepted +/// connections stay in a host-owned listener queue until `take_tcp_accept` is +/// called from a guest poll; canceling an accept waiter must not remove one. +pub trait HostTcpListenerIo: HostSocketIo { + fn submit_tcp_bind( + &self, + operation: HostOperationId, + options: &TcpListenOptions, + ) -> io::Result<()>; + + fn take_tcp_bind(&self, operation: HostOperationId) -> Poll>; + + fn submit_tcp_accept( + &self, + handle: HostSocketHandle, + operation: HostOperationId, + ) -> io::Result<()>; + + fn take_tcp_accept(&self, operation: HostOperationId) + -> Poll>; +} + +pub trait HostTcpListenerBackend: HostTcpListenerIo + HostTcpIo {} + +impl HostTcpListenerBackend for T where T: HostTcpListenerIo + HostTcpIo {} + +pub struct HostTcpListenerFactory +where + B: HostTcpListenerBackend, +{ + runtime: HostSocketRuntime, + backend: Arc, +} + +impl Clone for HostTcpListenerFactory +where + B: HostTcpListenerBackend, +{ + fn clone(&self) -> Self { + Self { + runtime: self.runtime.clone(), + backend: self.backend.clone(), + } + } +} + +impl HostTcpListenerFactory +where + B: HostTcpListenerBackend, +{ + pub fn new(runtime: HostSocketRuntime, backend: Arc) -> Self { + Self { runtime, backend } + } + + async fn bind(&self, options: TcpListenOptions) -> io::Result>> { + let result = self + .runtime + .run_operation( + self.backend.clone(), + |backend, operation| backend.submit_tcp_bind(operation, &options), + |backend, operation| backend.take_tcp_bind(operation), + |backend, operation| backend.cancel_operation(operation), + ) + .await?; + Ok(Arc::new(HostTcpListener { + runtime: self.runtime.clone(), + backend: self.backend.clone(), + handle: result.handle, + local_addr: result.local_addr, + })) + } +} + +#[async_trait] +impl VirtualTcpListenerFactory for HostTcpListenerFactory +where + B: HostTcpListenerBackend, +{ + type Listener = HostTcpListener; + + async fn bind_tcp(&self, options: TcpListenOptions) -> anyhow::Result> { + Ok(self.bind(options).await?) + } +} + +pub struct HostTcpListener +where + B: HostTcpListenerBackend, +{ + runtime: HostSocketRuntime, + backend: Arc, + handle: HostSocketHandle, + local_addr: SocketAddr, +} + +impl std::fmt::Debug for HostTcpListener +where + B: HostTcpListenerBackend, +{ + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("HostTcpListener") + .field("handle", &self.handle) + .field("local_addr", &self.local_addr) + .finish_non_exhaustive() + } +} + +#[async_trait] +impl VirtualTcpListener for HostTcpListener +where + B: HostTcpListenerBackend, +{ + type Socket = HostTcpStream; + + fn local_addr(&self) -> io::Result { + Ok(self.local_addr) + } + + async fn accept(&self) -> io::Result<(Self::Socket, SocketAddr)> { + let result = self + .runtime + .run_operation( + self.backend.clone(), + |backend, operation| backend.submit_tcp_accept(self.handle, operation), + |backend, operation| backend.take_tcp_accept(operation), + |backend, operation| backend.cancel_operation(operation), + ) + .await?; + let peer_addr = result.peer_addr; + Ok(( + self.runtime.tcp_stream( + self.backend.clone(), + result.handle, + result.local_addr, + peer_addr, + result.transport_label, + ), + peer_addr, + )) + } +} + +impl Drop for HostTcpListener +where + B: HostTcpListenerBackend, +{ + fn drop(&mut self) { + let _ = self.backend.close(self.handle); + } +} + +#[cfg(test)] +mod tests { + use std::{ + collections::{HashMap, HashSet, VecDeque}, + sync::Mutex, + }; + + use crate::socket::tcp::{TcpBindOptions, TcpListenPurpose, VirtualTcpSocket}; + + use super::*; + + enum TestOperation { + Bind { + options: TcpListenOptions, + result: Option>, + }, + Accept { + handle: HostSocketHandle, + }, + } + + #[derive(Default)] + struct TestHostIo { + operations: Mutex>, + accepted: Mutex>>, + cancelled: Mutex>, + closed: Mutex>, + } + + impl TestHostIo { + fn bind_operation(&self) -> HostOperationId { + self.operation(|operation| matches!(operation, TestOperation::Bind { .. })) + } + + fn accept_operation(&self) -> HostOperationId { + self.operation(|operation| matches!(operation, TestOperation::Accept { .. })) + } + + fn operation(&self, predicate: impl Fn(&TestOperation) -> bool) -> HostOperationId { + self.operations + .lock() + .unwrap() + .iter() + .find_map(|(operation, value)| predicate(value).then_some(*operation)) + .unwrap() + } + + fn complete_bind(&self, operation: HostOperationId, result: HostTcpBindResult) { + let mut operations = self.operations.lock().unwrap(); + let TestOperation::Bind { + result: completion, .. + } = operations.get_mut(&operation).unwrap() + else { + panic!("operation is not bind"); + }; + *completion = Some(Ok(result)); + } + + fn queue_accept(&self, operation: HostOperationId, result: HostTcpConnectResult) { + let operations = self.operations.lock().unwrap(); + let TestOperation::Accept { handle } = operations.get(&operation).unwrap() else { + panic!("operation is not accept"); + }; + self.accepted + .lock() + .unwrap() + .entry(*handle) + .or_default() + .push_back(result); + } + } + + impl HostSocketIo for TestHostIo { + fn cancel_operation(&self, operation: HostOperationId) -> io::Result<()> { + if let Some(TestOperation::Bind { + result: Some(Ok(result)), + .. + }) = self.operations.lock().unwrap().remove(&operation) + { + self.closed.lock().unwrap().insert(result.handle); + } + self.cancelled.lock().unwrap().push(operation); + Ok(()) + } + + fn close(&self, handle: HostSocketHandle) -> io::Result<()> { + self.closed.lock().unwrap().insert(handle); + Ok(()) + } + } + + impl HostTcpListenerIo for TestHostIo { + fn submit_tcp_bind( + &self, + operation: HostOperationId, + options: &TcpListenOptions, + ) -> io::Result<()> { + self.operations.lock().unwrap().insert( + operation, + TestOperation::Bind { + options: options.clone(), + result: None, + }, + ); + Ok(()) + } + + fn take_tcp_bind(&self, operation: HostOperationId) -> Poll> { + let mut operations = self.operations.lock().unwrap(); + let Some(TestOperation::Bind { result, .. }) = operations.get_mut(&operation) else { + return Poll::Ready(Err(io::ErrorKind::NotFound.into())); + }; + let Some(result) = result.take() else { + return Poll::Pending; + }; + operations.remove(&operation); + Poll::Ready(result) + } + + fn submit_tcp_accept( + &self, + handle: HostSocketHandle, + operation: HostOperationId, + ) -> io::Result<()> { + self.operations + .lock() + .unwrap() + .insert(operation, TestOperation::Accept { handle }); + Ok(()) + } + + fn take_tcp_accept( + &self, + operation: HostOperationId, + ) -> Poll> { + let mut operations = self.operations.lock().unwrap(); + let Some(TestOperation::Accept { handle }) = operations.get(&operation) else { + return Poll::Ready(Err(io::ErrorKind::NotFound.into())); + }; + let handle = *handle; + let Some(result) = self + .accepted + .lock() + .unwrap() + .entry(handle) + .or_default() + .pop_front() + else { + return Poll::Pending; + }; + operations.remove(&operation); + Poll::Ready(Ok(result)) + } + } + + impl HostTcpIo for TestHostIo { + fn submit_read( + &self, + _handle: HostSocketHandle, + _operation: HostOperationId, + _capacity: usize, + ) -> io::Result<()> { + Err(io::ErrorKind::Unsupported.into()) + } + + fn take_read(&self, _operation: HostOperationId) -> Poll>> { + Poll::Ready(Err(io::ErrorKind::Unsupported.into())) + } + + fn submit_write( + &self, + _handle: HostSocketHandle, + _operation: HostOperationId, + _source: &[u8], + ) -> io::Result<()> { + Err(io::ErrorKind::Unsupported.into()) + } + + fn take_write(&self, _operation: HostOperationId) -> Poll> { + Poll::Ready(Err(io::ErrorKind::Unsupported.into())) + } + } + + fn listener( + runtime: &HostSocketRuntime, + io: Arc, + handle: HostSocketHandle, + ) -> HostTcpListener { + HostTcpListener { + runtime: runtime.clone(), + backend: io, + handle, + local_addr: "192.0.2.1:11013".parse().unwrap(), + } + } + + #[tokio::test] + async fn forwards_bind_options_and_wraps_accepted_stream() { + let io = Arc::new(TestHostIo::default()); + let runtime = HostSocketRuntime::new(); + let factory = HostTcpListenerFactory::new(runtime.clone(), io.clone()); + let options = TcpListenOptions { + bind: TcpBindOptions::default() + .with_local_addr(Some("192.0.2.1:0".parse().unwrap())) + .with_socket_mark(Some(7)) + .with_bind_device(Some("host-device".to_owned())) + .with_reuse_port(true), + purpose: TcpListenPurpose::ManualConnect, + }; + let bind_task = tokio::spawn({ + let factory = factory.clone(); + let options = options.clone(); + async move { factory.bind_tcp(options).await } + }); + tokio::task::yield_now().await; + let bind_operation = io.bind_operation(); + { + let operations = io.operations.lock().unwrap(); + let TestOperation::Bind { + options: submitted, .. + } = operations.get(&bind_operation).unwrap() + else { + panic!("operation is not bind"); + }; + assert_eq!(submitted, &options); + } + io.complete_bind( + bind_operation, + HostTcpBindResult { + handle: HostSocketHandle(51), + local_addr: "192.0.2.1:11013".parse().unwrap(), + }, + ); + runtime.notify_completions(); + let listener = bind_task.await.unwrap().unwrap(); + + let accept_task = tokio::spawn({ + let listener = listener.clone(); + async move { listener.accept().await } + }); + tokio::task::yield_now().await; + let accept_operation = io.accept_operation(); + io.queue_accept( + accept_operation, + HostTcpConnectResult { + handle: HostSocketHandle(52), + local_addr: "192.0.2.1:11013".parse().unwrap(), + peer_addr: "192.0.2.2:40100".parse().unwrap(), + transport_label: Some("host-accepted".to_owned()), + }, + ); + runtime.notify_completions(); + let (stream, peer_addr) = accept_task.await.unwrap().unwrap(); + assert_eq!(peer_addr, "192.0.2.2:40100".parse().unwrap()); + assert_eq!(stream.transport_label(), Some("host-accepted")); + drop(stream); + drop(listener); + assert!(io.closed.lock().unwrap().contains(&HostSocketHandle(51))); + assert!(io.closed.lock().unwrap().contains(&HostSocketHandle(52))); + } + + #[tokio::test] + async fn cancelling_ready_accept_preserves_connection_for_next_poll() { + let io = Arc::new(TestHostIo::default()); + let runtime = HostSocketRuntime::new(); + let listener = listener(&runtime, io.clone(), HostSocketHandle(53)); + let peer_addr = "192.0.2.2:40101".parse().unwrap(); + let operation = { + let mut accept = Box::pin(listener.accept()); + assert!(futures::poll!(&mut accept).is_pending()); + let operation = io.accept_operation(); + io.queue_accept( + operation, + HostTcpConnectResult { + handle: HostSocketHandle(54), + local_addr: "192.0.2.1:11013".parse().unwrap(), + peer_addr, + transport_label: None, + }, + ); + runtime.notify_completions(); + drop(accept); + operation + }; + assert_eq!(*io.cancelled.lock().unwrap(), vec![operation]); + + let (stream, accepted_peer) = listener.accept().await.unwrap(); + assert_eq!(accepted_peer, peer_addr); + drop(stream); + assert!(io.closed.lock().unwrap().contains(&HostSocketHandle(54))); + } +} diff --git a/easytier-core/src/host/socket/mod.rs b/easytier-core/src/host/socket/mod.rs new file mode 100644 index 00000000..af7f9b34 --- /dev/null +++ b/easytier-core/src/host/socket/mod.rs @@ -0,0 +1,816 @@ +//! Host-backed socket seams. +//! +//! This module holds the socket-flavoured Host capability seams. Core owns +//! socket I/O scheduling, backpressure, and protocol state; the host Adapter +//! behind these traits owns the mechanical endpoint operations (see +//! CONTEXT.md, "Socket"). [`HostSocketIo`] and [`HostTcpIo`] are the base +//! operation traits keyed by [`HostSocketHandle`] and [`HostOperationId`]; +//! [`HostSocketRuntime`] schedules host completions; [`HostTcpStream`] +//! bridges host TCP I/O into core's socket traits. Socket creation lives in +//! [`factory`], TCP listener bind/accept in [`listener`], and the UDP +//! datagram bridge in [`udp`]. Concrete WASI adapters behind these seams live +//! in [`crate::wasi`]. + +use std::{ + collections::HashMap, + fmt, io, + net::SocketAddr, + pin::Pin, + sync::{Arc, LazyLock, Mutex, atomic::Ordering}, + task::{Context, Poll, Waker}, +}; + +use atomic_shim::AtomicU64; +use tokio::io::{AsyncRead, AsyncWrite, ReadBuf}; + +use crate::socket::tcp::VirtualTcpSocket; + +pub mod factory; +pub mod listener; +pub mod udp; + +static NEXT_HOST_OPERATION: LazyLock = LazyLock::new(|| AtomicU64::new(1)); + +#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] +pub struct HostSocketHandle(pub u64); + +#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] +pub struct HostOperationId(pub u64); + +/// Mechanical host I/O below core's socket scheduling seam. +pub trait HostSocketIo: Send + Sync + 'static { + fn cancel_operation(&self, operation: HostOperationId) -> io::Result<()>; + + /// Close must be idempotent. + fn close(&self, handle: HostSocketHandle) -> io::Result<()>; +} + +/// Mechanical host TCP I/O below core's socket scheduling seam. +/// +/// Submit methods must return without waiting for I/O. `submit_write` must take +/// ownership of the complete source before returning and complete only after +/// all accepted bytes are written or an error occurs. Completion methods return +/// host-owned results; they never retain guest-memory borrows. +pub trait HostTcpIo: HostSocketIo { + fn submit_read( + &self, + handle: HostSocketHandle, + operation: HostOperationId, + capacity: usize, + ) -> io::Result<()>; + + fn take_read(&self, operation: HostOperationId) -> Poll>>; + + fn submit_write( + &self, + handle: HostSocketHandle, + operation: HostOperationId, + source: &[u8], + ) -> io::Result<()>; + + fn take_write(&self, operation: HostOperationId) -> Poll>; +} + +#[derive(Default)] +pub(in crate::host) struct WakerRegistry { + wakers: Mutex>, +} + +impl WakerRegistry { + fn register(&self, operation: HostOperationId, waker: &Waker) { + let mut wakers = self.wakers.lock().expect("host waker registry poisoned"); + match wakers.get_mut(&operation) { + Some(registered) if registered.will_wake(waker) => {} + Some(registered) => *registered = waker.clone(), + None => { + wakers.insert(operation, waker.clone()); + } + } + } + + pub(in crate::host) fn remove(&self, operation: HostOperationId) { + self.wakers + .lock() + .expect("host waker registry poisoned") + .remove(&operation); + } + + fn wake_all(&self) { + let wakers = { + let mut registered = self.wakers.lock().expect("host waker registry poisoned"); + std::mem::take(&mut *registered) + }; + for waker in wakers.into_values() { + waker.wake(); + } + } +} + +pub(in crate::host) struct HostSocketRuntimeInner { + pub(in crate::host) completion_epoch: AtomicU64, + pub(in crate::host) wakers: WakerRegistry, +} + +#[derive(Clone)] +pub struct HostSocketRuntime { + pub(in crate::host) inner: Arc, +} + +impl Default for HostSocketRuntime { + fn default() -> Self { + Self::new() + } +} + +impl HostSocketRuntime { + pub fn new() -> Self { + Self { + inner: Arc::new(HostSocketRuntimeInner { + completion_epoch: AtomicU64::new(0), + wakers: WakerRegistry::default(), + }), + } + } + + pub fn tcp_stream( + &self, + io: Arc, + handle: HostSocketHandle, + local_addr: SocketAddr, + peer_addr: SocketAddr, + transport_label: Option, + ) -> HostTcpStream { + HostTcpStream { + runtime: self.clone(), + io, + handle, + local_addr, + peer_addr, + transport_label, + read_operation: None, + read_buffer: None, + read_eof: false, + write_operation: None, + closed: false, + } + } + + /// Wake socket tasks after the host reports one or more completions. + pub fn notify_completions(&self) { + self.inner.completion_epoch.fetch_add(1, Ordering::SeqCst); + self.inner.wakers.wake_all(); + } + + pub(in crate::host) fn next_operation(&self) -> HostOperationId { + loop { + let operation = NEXT_HOST_OPERATION.fetch_add(1, Ordering::Relaxed); + if operation != 0 { + return HostOperationId(operation); + } + } + } + + pub(in crate::host) fn register_pending( + &self, + operation: HostOperationId, + observed_epoch: u64, + context: &Context<'_>, + ) { + self.inner.wakers.register(operation, context.waker()); + if self.inner.completion_epoch.load(Ordering::SeqCst) != observed_epoch { + self.inner.wakers.remove(operation); + context.waker().wake_by_ref(); + } + } + + pub(in crate::host) async fn run_operation( + &self, + io: Arc, + submit: impl FnOnce(&I, HostOperationId) -> io::Result<()>, + take: impl Fn(&I, HostOperationId) -> Poll>, + cancel: fn(&I, HostOperationId) -> io::Result<()>, + ) -> io::Result + where + I: ?Sized + Send + Sync + 'static, + { + let operation = self.next_operation(); + submit(io.as_ref(), operation)?; + let mut pending = PendingHostOperation::new(self.clone(), io, operation, cancel); + futures::future::poll_fn(|context| { + pending.poll(context, |io, operation| take(io, operation)) + }) + .await + } +} + +pub(in crate::host) struct PendingHostOperation +where + I: ?Sized, +{ + runtime: HostSocketRuntime, + io: Arc, + operation: HostOperationId, + cancel: fn(&I, HostOperationId) -> io::Result<()>, + completed: bool, +} + +impl PendingHostOperation +where + I: ?Sized, +{ + pub(in crate::host) fn new( + runtime: HostSocketRuntime, + io: Arc, + operation: HostOperationId, + cancel: fn(&I, HostOperationId) -> io::Result<()>, + ) -> Self { + Self { + runtime, + io, + operation, + cancel, + completed: false, + } + } + + fn poll( + &mut self, + context: &Context<'_>, + take: impl FnOnce(&I, HostOperationId) -> Poll, + ) -> Poll { + let epoch = self.runtime.inner.completion_epoch.load(Ordering::SeqCst); + match take(self.io.as_ref(), self.operation) { + Poll::Pending => { + self.runtime + .register_pending(self.operation, epoch, context); + Poll::Pending + } + Poll::Ready(result) => { + self.complete(); + Poll::Ready(result) + } + } + } + + fn complete(&mut self) { + self.runtime.inner.wakers.remove(self.operation); + self.completed = true; + } + + fn cancel(mut self) -> io::Result<()> { + self.runtime.inner.wakers.remove(self.operation); + self.completed = true; + (self.cancel)(self.io.as_ref(), self.operation) + } +} + +impl Drop for PendingHostOperation +where + I: ?Sized, +{ + fn drop(&mut self) { + if self.completed { + return; + } + self.runtime.inner.wakers.remove(self.operation); + let _ = (self.cancel)(self.io.as_ref(), self.operation); + } +} + +struct ReadBuffer { + data: Vec, + offset: usize, +} + +pub struct HostTcpStream { + runtime: HostSocketRuntime, + io: Arc, + handle: HostSocketHandle, + local_addr: SocketAddr, + peer_addr: SocketAddr, + transport_label: Option, + read_operation: Option>, + read_buffer: Option, + read_eof: bool, + write_operation: Option>, + closed: bool, +} + +impl HostTcpStream { + fn close(&mut self) -> io::Result<()> { + if self.closed { + return Ok(()); + } + + let mut first_error = None; + for pending in [self.read_operation.take(), self.write_operation.take()] + .into_iter() + .flatten() + { + if let Err(error) = pending.cancel() + && first_error.is_none() + { + first_error = Some(error); + } + } + + match self.io.close(self.handle) { + Ok(()) => self.closed = true, + Err(error) if first_error.is_none() => first_error = Some(error), + Err(_) => {} + } + + match first_error { + Some(error) => Err(error), + None => Ok(()), + } + } + + fn copy_buffered_read(&mut self, buffer: &mut ReadBuf<'_>) -> bool { + let Some(pending) = &mut self.read_buffer else { + return false; + }; + let remaining = &pending.data[pending.offset..]; + let copy_len = remaining.len().min(buffer.remaining()); + buffer.put_slice(&remaining[..copy_len]); + pending.offset += copy_len; + if pending.offset == pending.data.len() { + self.read_buffer = None; + } + true + } + + fn poll_write_completion(&mut self, context: &Context<'_>) -> Poll> { + let Some(pending) = &mut self.write_operation else { + return Poll::Ready(Ok(())); + }; + match pending.poll(context, |io, operation| io.take_write(operation)) { + Poll::Pending => Poll::Pending, + Poll::Ready(result) => { + self.write_operation = None; + Poll::Ready(result) + } + } + } +} + +impl fmt::Debug for HostTcpStream { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter + .debug_struct("HostTcpStream") + .field("handle", &self.handle) + .field("local_addr", &self.local_addr) + .field("peer_addr", &self.peer_addr) + .field("transport_label", &self.transport_label) + .field("closed", &self.closed) + .finish_non_exhaustive() + } +} + +impl AsyncRead for HostTcpStream { + fn poll_read( + mut self: Pin<&mut Self>, + context: &mut Context<'_>, + buffer: &mut ReadBuf<'_>, + ) -> Poll> { + if buffer.remaining() == 0 || self.closed || self.read_eof { + return Poll::Ready(Ok(())); + } + if self.copy_buffered_read(buffer) { + return Poll::Ready(Ok(())); + } + + loop { + if self.read_operation.is_none() { + let operation = self.runtime.next_operation(); + if let Err(error) = self + .io + .submit_read(self.handle, operation, buffer.remaining()) + { + return Poll::Ready(Err(error)); + } + self.read_operation = Some(PendingHostOperation::new( + self.runtime.clone(), + self.io.clone(), + operation, + |io, operation| io.cancel_operation(operation), + )); + } + + let completion = self + .read_operation + .as_mut() + .expect("read operation was just installed") + .poll(context, |io, operation| io.take_read(operation)); + match completion { + Poll::Pending => return Poll::Pending, + Poll::Ready(result) => { + self.read_operation = None; + match result { + Ok(data) if data.is_empty() => { + self.read_eof = true; + return Poll::Ready(Ok(())); + } + Ok(data) => { + self.read_buffer = Some(ReadBuffer { data, offset: 0 }); + if self.copy_buffered_read(buffer) { + return Poll::Ready(Ok(())); + } + } + Err(error) => return Poll::Ready(Err(error)), + } + } + } + } + } +} + +impl AsyncWrite for HostTcpStream { + fn poll_write( + mut self: Pin<&mut Self>, + context: &mut Context<'_>, + buffer: &[u8], + ) -> Poll> { + if self.closed { + return Poll::Ready(Err(io::Error::new( + io::ErrorKind::BrokenPipe, + "host TCP stream is closed", + ))); + } + if buffer.is_empty() { + return Poll::Ready(Ok(0)); + } + + match self.poll_write_completion(context) { + Poll::Pending => return Poll::Pending, + Poll::Ready(Err(error)) => return Poll::Ready(Err(error)), + Poll::Ready(Ok(())) => {} + } + + let operation = self.runtime.next_operation(); + if let Err(error) = self.io.submit_write(self.handle, operation, buffer) { + return Poll::Ready(Err(error)); + } + self.write_operation = Some(PendingHostOperation::new( + self.runtime.clone(), + self.io.clone(), + operation, + |io, operation| io.cancel_operation(operation), + )); + Poll::Ready(Ok(buffer.len())) + } + + fn poll_flush(mut self: Pin<&mut Self>, context: &mut Context<'_>) -> Poll> { + self.poll_write_completion(context) + } + + fn poll_shutdown(mut self: Pin<&mut Self>, context: &mut Context<'_>) -> Poll> { + let flush_error = match self.poll_write_completion(context) { + Poll::Pending => return Poll::Pending, + Poll::Ready(Ok(())) => None, + Poll::Ready(Err(error)) => Some(error), + }; + let close_result = self.close(); + match (flush_error, close_result) { + (Some(error), _) => Poll::Ready(Err(error)), + (None, result) => Poll::Ready(result), + } + } +} + +impl VirtualTcpSocket for HostTcpStream { + fn local_addr(&self) -> io::Result { + Ok(self.local_addr) + } + + fn peer_addr(&self) -> io::Result { + Ok(self.peer_addr) + } + + fn transport_label(&self) -> Option<&str> { + self.transport_label.as_deref() + } +} + +impl Drop for HostTcpStream { + fn drop(&mut self) { + let _ = self.close(); + } +} + +#[cfg(test)] +mod tests { + use std::{collections::HashSet, sync::atomic::AtomicBool}; + + use tokio::io::{AsyncReadExt as _, AsyncWriteExt as _}; + + use super::*; + + impl WakerRegistry { + pub(crate) fn len(&self) -> usize { + self.wakers.lock().unwrap().len() + } + } + + enum TestOperation { + Read(Option>>), + Write { + source: Vec, + result: Option>, + }, + } + + #[derive(Default)] + struct TestHostIo { + operations: Mutex>, + cancelled: Mutex>, + closed: Mutex>, + notify_during_take: AtomicBool, + runtime: Mutex>, + } + + impl TestHostIo { + fn operation(&self, read: bool) -> HostOperationId { + self.operations + .lock() + .unwrap() + .iter() + .find_map(|(id, operation)| match (read, operation) { + (true, TestOperation::Read(_)) | (false, TestOperation::Write { .. }) => { + Some(*id) + } + _ => None, + }) + .expect("operation was not submitted") + } + + fn write_source(&self, operation: HostOperationId) -> Vec { + let operations = self.operations.lock().unwrap(); + let TestOperation::Write { source, .. } = operations.get(&operation).unwrap() else { + panic!("operation is not a write"); + }; + source.clone() + } + + fn complete_read(&self, operation: HostOperationId, data: Vec) { + let mut operations = self.operations.lock().unwrap(); + let TestOperation::Read(result) = operations.get_mut(&operation).unwrap() else { + panic!("operation is not a read"); + }; + *result = Some(Ok(data)); + } + + fn complete_write(&self, operation: HostOperationId) { + let mut operations = self.operations.lock().unwrap(); + let TestOperation::Write { result, .. } = operations.get_mut(&operation).unwrap() + else { + panic!("operation is not a write"); + }; + *result = Some(Ok(())); + } + + fn fail_write(&self, operation: HostOperationId) { + let mut operations = self.operations.lock().unwrap(); + let TestOperation::Write { result, .. } = operations.get_mut(&operation).unwrap() + else { + panic!("operation is not a write"); + }; + *result = Some(Err(io::Error::other("write failed"))); + } + } + + impl HostSocketIo for TestHostIo { + fn cancel_operation(&self, operation: HostOperationId) -> io::Result<()> { + self.operations.lock().unwrap().remove(&operation); + self.cancelled.lock().unwrap().push(operation); + Ok(()) + } + + fn close(&self, handle: HostSocketHandle) -> io::Result<()> { + self.closed.lock().unwrap().insert(handle); + Ok(()) + } + } + + impl HostTcpIo for TestHostIo { + fn submit_read( + &self, + _handle: HostSocketHandle, + operation: HostOperationId, + _capacity: usize, + ) -> io::Result<()> { + self.operations + .lock() + .unwrap() + .insert(operation, TestOperation::Read(None)); + Ok(()) + } + + fn take_read(&self, operation: HostOperationId) -> Poll>> { + if self.notify_during_take.swap(false, Ordering::SeqCst) { + self.runtime + .lock() + .unwrap() + .as_ref() + .unwrap() + .notify_completions(); + return Poll::Pending; + } + let mut operations = self.operations.lock().unwrap(); + let Some(TestOperation::Read(result)) = operations.get_mut(&operation) else { + return Poll::Ready(Err(io::Error::new( + io::ErrorKind::NotFound, + "read operation is missing", + ))); + }; + let Some(result) = result.take() else { + return Poll::Pending; + }; + operations.remove(&operation); + Poll::Ready(result) + } + + fn submit_write( + &self, + _handle: HostSocketHandle, + operation: HostOperationId, + source: &[u8], + ) -> io::Result<()> { + self.operations.lock().unwrap().insert( + operation, + TestOperation::Write { + source: source.to_vec(), + result: None, + }, + ); + Ok(()) + } + + fn take_write(&self, operation: HostOperationId) -> Poll> { + let mut operations = self.operations.lock().unwrap(); + let Some(TestOperation::Write { result, .. }) = operations.get_mut(&operation) else { + return Poll::Ready(Err(io::Error::new( + io::ErrorKind::NotFound, + "write operation is missing", + ))); + }; + let Some(result) = result.take() else { + return Poll::Pending; + }; + operations.remove(&operation); + Poll::Ready(result) + } + } + + fn test_stream(io: Arc) -> (HostSocketRuntime, HostTcpStream) { + let runtime = HostSocketRuntime::new(); + *io.runtime.lock().unwrap() = Some(runtime.clone()); + let stream = runtime.tcp_stream( + io, + HostSocketHandle(7), + "192.0.2.1:10000".parse().unwrap(), + "192.0.2.2:11013".parse().unwrap(), + Some("host-test".to_owned()), + ); + (runtime, stream) + } + + #[tokio::test] + async fn pending_completions_wake_reads_and_apply_write_backpressure() { + let io = Arc::new(TestHostIo::default()); + let (runtime, stream) = test_stream(io.clone()); + assert_eq!(stream.transport_label(), Some("host-test")); + let (mut reader, mut writer) = tokio::io::split(stream); + + let read_task = tokio::spawn(async move { + let mut data = [0_u8; 3]; + reader.read_exact(&mut data).await.unwrap(); + data + }); + tokio::task::yield_now().await; + let read_operation = io.operation(true); + io.complete_read(read_operation, b"abc".to_vec()); + runtime.notify_completions(); + assert_eq!(read_task.await.unwrap(), *b"abc"); + + let write_task = tokio::spawn(async move { + writer.write_all(b"one").await.unwrap(); + writer.write_all(b"two").await.unwrap(); + writer.shutdown().await.unwrap(); + }); + tokio::task::yield_now().await; + let first_write = io.operation(false); + assert_eq!(io.write_source(first_write), b"one"); + io.complete_write(first_write); + runtime.notify_completions(); + tokio::task::yield_now().await; + let second_write = io.operation(false); + assert_ne!(second_write, first_write); + assert_eq!(io.write_source(second_write), b"two"); + io.complete_write(second_write); + runtime.notify_completions(); + write_task.await.unwrap(); + + assert!(io.closed.lock().unwrap().contains(&HostSocketHandle(7))); + } + + #[tokio::test] + async fn cancelled_read_keeps_owned_completion_remainder() { + let io = Arc::new(TestHostIo::default()); + let (runtime, mut stream) = test_stream(io.clone()); + let mut large = [0_u8; 4]; + let mut first_read = Box::pin(stream.read(&mut large)); + assert!(futures::poll!(&mut first_read).is_pending()); + drop(first_read); + + let operation = io.operation(true); + io.complete_read(operation, b"abcd".to_vec()); + runtime.notify_completions(); + let mut first = [0_u8; 1]; + stream.read_exact(&mut first).await.unwrap(); + assert_eq!(&first, b"a"); + let mut remainder = [0_u8; 3]; + stream.read_exact(&mut remainder).await.unwrap(); + assert_eq!(&remainder, b"bcd"); + } + + #[tokio::test] + async fn empty_read_completion_reports_eof() { + let io = Arc::new(TestHostIo::default()); + let (runtime, mut stream) = test_stream(io.clone()); + let read_task = tokio::spawn(async move { + let mut byte = [0_u8; 1]; + stream.read(&mut byte).await.unwrap() + }); + tokio::task::yield_now().await; + let operation = io.operation(true); + io.complete_read(operation, Vec::new()); + runtime.notify_completions(); + + assert_eq!(read_task.await.unwrap(), 0); + } + + #[tokio::test] + async fn shutdown_closes_handle_after_buffered_write_error() { + let io = Arc::new(TestHostIo::default()); + let (runtime, mut stream) = test_stream(io.clone()); + stream.write_all(b"data").await.unwrap(); + let operation = io.operation(false); + io.fail_write(operation); + runtime.notify_completions(); + + let error = stream.shutdown().await.unwrap_err(); + assert_eq!(error.to_string(), "write failed"); + assert!(io.closed.lock().unwrap().contains(&HostSocketHandle(7))); + } + + #[tokio::test] + async fn completion_between_poll_and_registration_is_not_lost() { + let io = Arc::new(TestHostIo::default()); + let (_runtime, mut stream) = test_stream(io.clone()); + let operation = { + let mut byte = [0_u8; 1]; + let mut read = Box::pin(stream.read(&mut byte)); + assert!(futures::poll!(&mut read).is_pending()); + io.operation(true) + }; + io.complete_read(operation, b"x".to_vec()); + io.notify_during_take.store(true, Ordering::SeqCst); + + let mut byte = [0_u8; 1]; + tokio::time::timeout( + std::time::Duration::from_secs(1), + stream.read_exact(&mut byte), + ) + .await + .unwrap() + .unwrap(); + assert_eq!(&byte, b"x"); + } + + #[tokio::test] + async fn dropping_pending_stream_removes_waker_cancels_and_closes() { + let io = Arc::new(TestHostIo::default()); + let (runtime, mut stream) = test_stream(io.clone()); + let read_task = tokio::spawn(async move { + let mut byte = [0_u8; 1]; + let _ = stream.read(&mut byte).await; + }); + tokio::task::yield_now().await; + let operation = io.operation(true); + assert_eq!(runtime.inner.wakers.len(), 1); + read_task.abort(); + let _ = read_task.await; + + assert_eq!(runtime.inner.wakers.len(), 0); + assert_eq!(*io.cancelled.lock().unwrap(), vec![operation]); + assert!(io.closed.lock().unwrap().contains(&HostSocketHandle(7))); + } + + #[test] + fn shared_host_io_receives_unique_operation_ids() { + let runtime_a = HostSocketRuntime::new(); + let runtime_b = HostSocketRuntime::new(); + assert_ne!(runtime_a.next_operation(), runtime_b.next_operation()); + } +} diff --git a/easytier-core/src/host/socket/udp.rs b/easytier-core/src/host/socket/udp.rs new file mode 100644 index 00000000..f8aece2e --- /dev/null +++ b/easytier-core/src/host/socket/udp.rs @@ -0,0 +1,522 @@ +use std::{fmt, io, net::SocketAddr, sync::Arc, task::Poll}; + +use async_trait::async_trait; + +use crate::socket::{ + SocketContext, + udp::{UdpSocketRecvMeta, UdpSocketSendMeta, VirtualUdpSocket}, +}; + +use super::{HostOperationId, HostSocketHandle, HostSocketIo, HostSocketRuntime}; + +/// One host-owned UDP receive completion. +#[derive(Debug)] +pub struct HostUdpDatagram { + pub data: Vec, + pub peer_addr: SocketAddr, + pub meta: UdpSocketRecvMeta, +} + +/// Mechanical host UDP I/O below core's datagram scheduling seam. +/// +/// Submit methods register readiness without performing the datagram I/O. +/// Receive datagrams stay in a host-owned socket queue until `take_recv` is +/// called from a guest poll, so canceling a waiter cannot consume a datagram. +/// `try_send` must synchronously copy one complete datagram into a bounded host +/// queue or return `WouldBlock`; it must never retain guest-memory borrows. +pub trait HostUdpIo: HostSocketIo { + fn submit_recv( + &self, + handle: HostSocketHandle, + operation: HostOperationId, + capacity: usize, + ) -> io::Result<()>; + + fn take_recv(&self, operation: HostOperationId) -> Poll>; + + fn try_send( + &self, + handle: HostSocketHandle, + source: &[u8], + peer_addr: SocketAddr, + meta: UdpSocketSendMeta, + ) -> io::Result<()>; + + fn submit_send_ready( + &self, + handle: HostSocketHandle, + operation: HostOperationId, + ) -> io::Result<()>; + + fn take_send_ready(&self, operation: HostOperationId) -> Poll>; +} + +pub struct HostUdpSocket { + runtime: HostSocketRuntime, + io: Arc, + handle: HostSocketHandle, + local_addr: SocketAddr, + context: SocketContext, +} + +impl HostSocketRuntime { + pub fn udp_socket( + &self, + io: Arc, + handle: HostSocketHandle, + local_addr: SocketAddr, + ) -> HostUdpSocket { + self.udp_socket_with_context(io, handle, local_addr, SocketContext::default()) + } + + pub fn udp_socket_with_context( + &self, + io: Arc, + handle: HostSocketHandle, + local_addr: SocketAddr, + context: SocketContext, + ) -> HostUdpSocket { + HostUdpSocket { + runtime: self.clone(), + io, + handle, + local_addr, + context, + } + } +} + +impl HostUdpSocket { + async fn send( + &self, + data: &[u8], + peer_addr: SocketAddr, + meta: UdpSocketSendMeta, + ) -> io::Result { + loop { + match self.io.try_send(self.handle, data, peer_addr, meta) { + Ok(()) => return Ok(data.len()), + Err(error) if error.kind() == io::ErrorKind::WouldBlock => {} + Err(error) => return Err(error), + } + + self.runtime + .run_operation( + self.io.clone(), + |io, operation| io.submit_send_ready(self.handle, operation), + |io, operation| io.take_send_ready(operation), + |io, operation| io.cancel_operation(operation), + ) + .await?; + } + } + + async fn receive(&self, buffer: &mut [u8]) -> io::Result { + let datagram = self + .runtime + .run_operation( + self.io.clone(), + |io, operation| io.submit_recv(self.handle, operation, buffer.len()), + |io, operation| io.take_recv(operation), + |io, operation| io.cancel_operation(operation), + ) + .await?; + + let copy_len = datagram.data.len().min(buffer.len()); + buffer[..copy_len].copy_from_slice(&datagram.data[..copy_len]); + let mut datagram = datagram; + datagram.data.truncate(copy_len); + Ok(datagram) + } +} + +impl fmt::Debug for HostUdpSocket { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter + .debug_struct("HostUdpSocket") + .field("handle", &self.handle) + .field("local_addr", &self.local_addr) + .finish_non_exhaustive() + } +} + +#[async_trait] +impl VirtualUdpSocket for HostUdpSocket { + fn local_addr(&self) -> io::Result { + Ok(self.local_addr) + } + + fn socket_context(&self) -> SocketContext { + self.context.clone() + } + + async fn send_to(&self, data: &[u8], addr: SocketAddr) -> io::Result { + self.send(data, addr, UdpSocketSendMeta::default()).await + } + + async fn recv_from(&self, buffer: &mut [u8]) -> io::Result<(usize, SocketAddr)> { + let datagram = self.receive(buffer).await?; + Ok((datagram.data.len(), datagram.peer_addr)) + } + + async fn send_to_with_meta( + &self, + data: &[u8], + addr: SocketAddr, + meta: UdpSocketSendMeta, + ) -> io::Result { + self.send(data, addr, meta).await + } + + async fn recv_from_with_meta( + &self, + buffer: &mut [u8], + ) -> io::Result<(usize, SocketAddr, UdpSocketRecvMeta)> { + let datagram = self.receive(buffer).await?; + Ok((datagram.data.len(), datagram.peer_addr, datagram.meta)) + } +} + +impl Drop for HostUdpSocket { + fn drop(&mut self) { + let _ = self.io.close(self.handle); + } +} + +#[cfg(test)] +mod tests { + use std::{ + collections::{HashMap, HashSet, VecDeque}, + net::{IpAddr, Ipv4Addr}, + sync::{ + Mutex, + atomic::{AtomicBool, Ordering}, + }, + }; + + use super::*; + + enum TestOperation { + Receive, + SendReady(Option>), + } + + #[derive(Default)] + struct TestHostUdpIo { + operations: Mutex>, + received: Mutex>, + sent: Mutex, SocketAddr, UdpSocketSendMeta)>>, + writable: AtomicBool, + cancelled: Mutex>, + closed: Mutex>, + } + + impl TestHostUdpIo { + fn operation(&self, receive: bool) -> HostOperationId { + self.operations + .lock() + .unwrap() + .iter() + .find_map(|(id, operation)| match (receive, operation) { + (true, TestOperation::Receive) | (false, TestOperation::SendReady(_)) => { + Some(*id) + } + _ => None, + }) + .expect("operation was not submitted") + } + + fn complete_recv(&self, operation: HostOperationId, datagram: HostUdpDatagram) { + let operations = self.operations.lock().unwrap(); + let TestOperation::Receive = operations.get(&operation).unwrap() else { + panic!("operation is not a receive"); + }; + self.received.lock().unwrap().push_back(datagram); + } + + fn complete_send(&self, operation: HostOperationId) { + let mut operations = self.operations.lock().unwrap(); + let TestOperation::SendReady(result) = operations.get_mut(&operation).unwrap() else { + panic!("operation is not a send"); + }; + *result = Some(Ok(())); + self.writable.store(true, Ordering::SeqCst); + } + } + + impl HostSocketIo for TestHostUdpIo { + fn cancel_operation(&self, operation: HostOperationId) -> io::Result<()> { + self.operations.lock().unwrap().remove(&operation); + self.cancelled.lock().unwrap().push(operation); + Ok(()) + } + + fn close(&self, handle: HostSocketHandle) -> io::Result<()> { + self.closed.lock().unwrap().insert(handle); + Ok(()) + } + } + + impl HostUdpIo for TestHostUdpIo { + fn submit_recv( + &self, + _handle: HostSocketHandle, + operation: HostOperationId, + _capacity: usize, + ) -> io::Result<()> { + self.operations + .lock() + .unwrap() + .insert(operation, TestOperation::Receive); + Ok(()) + } + + fn take_recv(&self, operation: HostOperationId) -> Poll> { + let mut operations = self.operations.lock().unwrap(); + let Some(TestOperation::Receive) = operations.get_mut(&operation) else { + return Poll::Ready(Err(io::Error::new( + io::ErrorKind::NotFound, + "receive operation is missing", + ))); + }; + let Some(datagram) = self.received.lock().unwrap().pop_front() else { + return Poll::Pending; + }; + operations.remove(&operation); + Poll::Ready(Ok(datagram)) + } + + fn try_send( + &self, + _handle: HostSocketHandle, + source: &[u8], + peer_addr: SocketAddr, + meta: UdpSocketSendMeta, + ) -> io::Result<()> { + if !self.writable.swap(false, Ordering::SeqCst) { + return Err(io::ErrorKind::WouldBlock.into()); + } + self.sent + .lock() + .unwrap() + .push((source.to_vec(), peer_addr, meta)); + Ok(()) + } + + fn submit_send_ready( + &self, + _handle: HostSocketHandle, + operation: HostOperationId, + ) -> io::Result<()> { + self.operations + .lock() + .unwrap() + .insert(operation, TestOperation::SendReady(None)); + Ok(()) + } + + fn take_send_ready(&self, operation: HostOperationId) -> Poll> { + let mut operations = self.operations.lock().unwrap(); + let Some(TestOperation::SendReady(result)) = operations.get_mut(&operation) else { + return Poll::Ready(Err(io::Error::new( + io::ErrorKind::NotFound, + "send readiness operation is missing", + ))); + }; + let Some(result) = result.take() else { + return Poll::Pending; + }; + operations.remove(&operation); + Poll::Ready(result) + } + } + + fn test_socket(io: Arc) -> (HostSocketRuntime, HostUdpSocket) { + let runtime = HostSocketRuntime::new(); + let socket = + runtime.udp_socket(io, HostSocketHandle(9), "192.0.2.1:11013".parse().unwrap()); + (runtime, socket) + } + + #[tokio::test] + async fn send_waits_for_readiness_then_enqueues_atomically() { + let io = Arc::new(TestHostUdpIo::default()); + let (runtime, socket) = test_socket(io.clone()); + let socket = Arc::new(socket); + let peer_addr = "192.0.2.2:22026".parse().unwrap(); + let source_ip = IpAddr::V4(Ipv4Addr::new(192, 0, 2, 10)); + let task = tokio::spawn({ + let socket = socket.clone(); + async move { + socket + .send_to_with_meta( + b"udp", + peer_addr, + UdpSocketSendMeta { + src_ip: Some(source_ip), + src_ifindex: None, + }, + ) + .await + } + }); + tokio::task::yield_now().await; + + let operation = io.operation(false); + assert!(io.sent.lock().unwrap().is_empty()); + + io.complete_send(operation); + runtime.notify_completions(); + assert_eq!(task.await.unwrap().unwrap(), 3); + assert_eq!( + *io.sent.lock().unwrap(), + vec![( + b"udp".to_vec(), + peer_addr, + UdpSocketSendMeta { + src_ip: Some(source_ip), + src_ifindex: None, + } + )] + ); + drop(socket); + assert!(io.closed.lock().unwrap().contains(&HostSocketHandle(9))); + } + + #[tokio::test] + async fn receive_returns_owned_payload_peer_and_destination_metadata() { + let io = Arc::new(TestHostUdpIo::default()); + let (runtime, socket) = test_socket(io.clone()); + let socket = Arc::new(socket); + let peer_addr = "192.0.2.2:22026".parse().unwrap(); + let destination_ip = IpAddr::V4(Ipv4Addr::new(192, 0, 2, 10)); + let task = tokio::spawn({ + let socket = socket.clone(); + async move { + let mut buffer = [0_u8; 4]; + let result = socket.recv_from_with_meta(&mut buffer).await; + (result, buffer) + } + }); + tokio::task::yield_now().await; + + let operation = io.operation(true); + io.complete_recv( + operation, + HostUdpDatagram { + data: b"data".to_vec(), + peer_addr, + meta: UdpSocketRecvMeta { + dst_ip: Some(destination_ip), + }, + }, + ); + runtime.notify_completions(); + + let (result, buffer) = task.await.unwrap(); + let (length, received_peer, meta) = result.unwrap(); + assert_eq!(length, 4); + assert_eq!(&buffer, b"data"); + assert_eq!(received_peer, peer_addr); + assert_eq!(meta.dst_ip, Some(destination_ip)); + } + + #[tokio::test] + async fn cancelling_completed_receive_preserves_datagram_for_next_poll() { + let io = Arc::new(TestHostUdpIo::default()); + let (runtime, socket) = test_socket(io.clone()); + let peer_addr = "192.0.2.2:22026".parse().unwrap(); + let operation = { + let mut buffer = [0_u8; 1]; + let mut receive = Box::pin(socket.recv_from(&mut buffer)); + assert!(futures::poll!(&mut receive).is_pending()); + let operation = io.operation(true); + assert_eq!(runtime.inner.wakers.len(), 1); + io.complete_recv( + operation, + HostUdpDatagram { + data: b"kept".to_vec(), + peer_addr, + meta: UdpSocketRecvMeta::default(), + }, + ); + runtime.notify_completions(); + drop(receive); + operation + }; + + assert_eq!(runtime.inner.wakers.len(), 0); + assert_eq!(*io.cancelled.lock().unwrap(), vec![operation]); + let mut recovered = [0_u8; 4]; + assert_eq!( + socket.recv_from(&mut recovered).await.unwrap(), + (4, peer_addr) + ); + assert_eq!(&recovered, b"kept"); + drop(socket); + assert!(io.closed.lock().unwrap().contains(&HostSocketHandle(9))); + } + + #[tokio::test] + async fn cancelling_send_readiness_never_enqueues_datagram() { + let io = Arc::new(TestHostUdpIo::default()); + let (runtime, socket) = test_socket(io.clone()); + let peer_addr = "192.0.2.2:22026".parse().unwrap(); + let operation = { + let mut send = Box::pin(socket.send_to(b"cancelled", peer_addr)); + assert!(futures::poll!(&mut send).is_pending()); + let operation = io.operation(false); + io.complete_send(operation); + runtime.notify_completions(); + drop(send); + operation + }; + + assert_eq!(*io.cancelled.lock().unwrap(), vec![operation]); + assert!(io.sent.lock().unwrap().is_empty()); + } + + async fn receive_payload( + runtime: &HostSocketRuntime, + io: &Arc, + socket: Arc, + payload: Vec, + ) -> (usize, [u8; N]) { + let peer_addr = "192.0.2.2:22026".parse().unwrap(); + let task = tokio::spawn(async move { + let mut buffer = [0_u8; N]; + let (length, _) = socket.recv_from(&mut buffer).await.unwrap(); + (length, buffer) + }); + tokio::task::yield_now().await; + let operation = io.operation(true); + io.complete_recv( + operation, + HostUdpDatagram { + data: payload, + peer_addr, + meta: UdpSocketRecvMeta::default(), + }, + ); + runtime.notify_completions(); + task.await.unwrap() + } + + #[tokio::test] + async fn receive_truncates_and_distinguishes_empty_buffers_from_datagrams() { + let io = Arc::new(TestHostUdpIo::default()); + let (runtime, socket) = test_socket(io.clone()); + let socket = Arc::new(socket); + + let (length, buffer) = + receive_payload(&runtime, &io, socket.clone(), b"long".to_vec()).await; + assert_eq!((length, buffer), (2, *b"lo")); + + let (length, buffer) = + receive_payload::<0>(&runtime, &io, socket.clone(), b"consumed".to_vec()).await; + assert_eq!((length, buffer), (0, [])); + + let (length, buffer) = receive_payload(&runtime, &io, socket, Vec::new()).await; + assert_eq!((length, buffer), (0, [0])); + } +} diff --git a/easytier-core/src/host/testkit.rs b/easytier-core/src/host/testkit.rs new file mode 100644 index 00000000..6ab47745 --- /dev/null +++ b/easytier-core/src/host/testkit.rs @@ -0,0 +1,206 @@ +//! Shared virtual-socket and DNS fakes for core unit tests. +//! +//! `gateway` and `instance` tests drive the same portable host seams; both use +//! this kit instead of keeping parallel copies. The fakes are inert by +//! default: TCP connects only succeed for proxy-NAT purposes, listeners never +//! accept, UDP sockets never receive, and DNS answers only literal IP hosts. + +use std::{ + io, + net::{IpAddr, SocketAddr}, + pin::Pin, + sync::{ + Arc, + atomic::{AtomicUsize, Ordering}, + }, + task::{Context, Poll}, +}; + +use tokio::io::{AsyncRead, AsyncWrite, ReadBuf}; + +use super::dns::{DnsQuery, DnsRecordResolver, DnsResolver, DnsSrvRecord}; +use crate::socket::{ + tcp::{ + TcpConnectOptions, TcpListenOptions, TcpListenPurpose, TcpSocketPurpose, + VirtualTcpListener, VirtualTcpListenerFactory, VirtualTcpSocket, VirtualTcpSocketFactory, + }, + udp::{UdpBindOptions, VirtualUdpSocket, VirtualUdpSocketFactory}, +}; + +pub struct TestTcpSocket(pub tokio::io::DuplexStream); + +impl AsyncRead for TestTcpSocket { + fn poll_read( + self: Pin<&mut Self>, + _cx: &mut Context<'_>, + buf: &mut ReadBuf<'_>, + ) -> Poll> { + Pin::new(&mut self.get_mut().0).poll_read(_cx, buf) + } +} + +impl AsyncWrite for TestTcpSocket { + fn poll_write( + self: Pin<&mut Self>, + _cx: &mut Context<'_>, + buf: &[u8], + ) -> Poll> { + Pin::new(&mut self.get_mut().0).poll_write(_cx, buf) + } + + fn poll_flush(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + Pin::new(&mut self.get_mut().0).poll_flush(cx) + } + + fn poll_shutdown(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + Pin::new(&mut self.get_mut().0).poll_shutdown(cx) + } +} + +impl VirtualTcpSocket for TestTcpSocket { + fn local_addr(&self) -> io::Result { + Ok("127.0.0.1:20000".parse().unwrap()) + } + + fn peer_addr(&self) -> io::Result { + Ok("127.0.0.1:20001".parse().unwrap()) + } +} + +pub struct TestTcpListener { + address: SocketAddr, + active_listeners: Arc, +} + +impl Drop for TestTcpListener { + fn drop(&mut self) { + self.active_listeners.fetch_sub(1, Ordering::Relaxed); + } +} + +#[async_trait::async_trait] +impl VirtualTcpListener for TestTcpListener { + type Socket = TestTcpSocket; + + fn local_addr(&self) -> io::Result { + Ok(self.address) + } + + async fn accept(&self) -> io::Result<(Self::Socket, SocketAddr)> { + std::future::pending().await + } +} + +pub struct TestUdpSocket(pub SocketAddr); + +#[async_trait::async_trait] +impl VirtualUdpSocket for TestUdpSocket { + fn local_addr(&self) -> io::Result { + Ok(self.0) + } + + async fn send_to(&self, data: &[u8], _addr: SocketAddr) -> io::Result { + Ok(data.len()) + } + + async fn recv_from(&self, _buf: &mut [u8]) -> io::Result<(usize, SocketAddr)> { + std::future::pending().await + } +} + +#[derive(Default)] +pub struct TestHost { + pub tcp_binds: AtomicUsize, + pub active_tcp_listeners: Arc, + pub udp_binds: AtomicUsize, + pub proxy_nat_connections: + Option>, + pub reject_socks5_listener: bool, +} + +#[async_trait::async_trait] +impl VirtualTcpSocketFactory for TestHost { + type Socket = TestTcpSocket; + + async fn connect_tcp(&self, options: TcpConnectOptions) -> anyhow::Result { + if options.purpose != TcpSocketPurpose::ProxyNat { + anyhow::bail!("test host does not connect non-proxy TCP sockets"); + } + let connections = self + .proxy_nat_connections + .as_ref() + .ok_or_else(|| anyhow::anyhow!("test host proxy NAT is disabled"))?; + let (socket, peer) = tokio::io::duplex(1024); + connections + .send((options.remote_addr, peer)) + .map_err(|_| anyhow::anyhow!("test host proxy NAT receiver is closed"))?; + Ok(TestTcpSocket(socket)) + } +} + +#[async_trait::async_trait] +impl VirtualTcpListenerFactory for TestHost { + type Listener = TestTcpListener; + + async fn bind_tcp(&self, options: TcpListenOptions) -> anyhow::Result> { + self.tcp_binds.fetch_add(1, Ordering::Relaxed); + if self.reject_socks5_listener && options.purpose == TcpListenPurpose::Socks5 { + anyhow::bail!("test host rejected SOCKS5 listener"); + } + let address = options + .bind + .local_addr + .unwrap_or_else(|| "127.0.0.1:20000".parse().unwrap()); + // Ephemeral binds still report a fixed nonzero port: the gateway + // smoltcp connector feeds `local_addr().port()` into the virtual + // stack, which rejects source port 0 as unaddressable. + let address = if address.port() == 0 { + SocketAddr::new(address.ip(), 20000) + } else { + address + }; + self.active_tcp_listeners.fetch_add(1, Ordering::Relaxed); + Ok(Arc::new(TestTcpListener { + address, + active_listeners: self.active_tcp_listeners.clone(), + })) + } +} + +#[async_trait::async_trait] +impl VirtualUdpSocketFactory for TestHost { + type Socket = TestUdpSocket; + + async fn bind_udp(&self, options: UdpBindOptions) -> anyhow::Result> { + self.udp_binds.fetch_add(1, Ordering::Relaxed); + let address = options + .local_addr + .unwrap_or_else(|| "127.0.0.1:20002".parse().unwrap()); + let address = if address.port() == 0 { + SocketAddr::new(address.ip(), 20002) + } else { + address + }; + Ok(Arc::new(TestUdpSocket(address))) + } +} + +pub struct TestDns; + +#[async_trait::async_trait] +impl DnsResolver for TestDns { + async fn resolve(&self, query: DnsQuery) -> anyhow::Result> { + Ok(query.host.parse().into_iter().collect()) + } +} + +#[async_trait::async_trait] +impl DnsRecordResolver for TestDns { + async fn resolve_txt(&self, _query: DnsQuery) -> anyhow::Result { + anyhow::bail!("test DNS has no TXT records") + } + + async fn resolve_srv(&self, _query: DnsQuery) -> anyhow::Result> { + Ok(Vec::new()) + } +} diff --git a/easytier-core/src/instance/build_capabilities.rs b/easytier-core/src/instance/build_capabilities.rs new file mode 100644 index 00000000..ea331d77 --- /dev/null +++ b/easytier-core/src/instance/build_capabilities.rs @@ -0,0 +1,79 @@ +//! Compile-time capability selection for portable Instance validation. +//! +//! Cargo features are localized here so configuration validation remains one +//! unconditional path in `CoreInstance`. + +use crate::config::{peers::PeerRuntimeSnapshot, runtime::CoreRuntimeConfig}; + +use super::CoreInstanceConfig; + +const DHCP_IPV4_AVAILABLE: bool = cfg!(feature = "dhcp-ipv4"); +const SMOLTCP_GATEWAY_AVAILABLE: bool = cfg!(feature = "proxy-smoltcp-stack"); +const PACKET_PROXY_AVAILABLE: bool = cfg!(feature = "proxy-packet"); +const PROXY_CIDR_MONITOR_AVAILABLE: bool = cfg!(feature = "proxy-cidr-monitor"); +const WRAPPED_TRANSPORT_AVAILABLE: bool = cfg!(feature = "wrapped-transport"); +const PUBLIC_IPV6_AVAILABLE: bool = cfg!(feature = "public-ipv6-provider"); +const VPN_PORTAL_AVAILABLE: bool = cfg!(feature = "vpn-portal"); + +fn require(available: bool, requested: bool, capability: &str) -> anyhow::Result<()> { + if requested && !available { + anyhow::bail!("this build does not include {capability}"); + } + Ok(()) +} + +fn validate_snapshot( + runtime: &CoreRuntimeConfig, + peer: &PeerRuntimeSnapshot, +) -> anyhow::Result<()> { + let core = &peer.runtime.core; + require(DHCP_IPV4_AVAILABLE, runtime.dhcp_ipv4, "DHCP IPv4")?; + require( + SMOLTCP_GATEWAY_AVAILABLE, + runtime.gateway.socks5_bind.is_some() || !runtime.gateway.port_forwards.is_empty(), + "the smoltcp gateway", + )?; + require( + PROXY_CIDR_MONITOR_AVAILABLE, + runtime + .manual_routes + .as_ref() + .is_some_and(|routes| !routes.is_empty()), + "the proxy CIDR monitor", + )?; + require( + WRAPPED_TRANSPORT_AVAILABLE, + !core.routes.proxy_networks.is_empty(), + "proxy routing services", + )?; + require( + PACKET_PROXY_AVAILABLE, + runtime + .proxy + .should_start(!core.routes.proxy_networks.is_empty()), + "packet proxy services", + )?; + require( + PUBLIC_IPV6_AVAILABLE, + runtime.public_ipv6_auto + || runtime.public_ipv6_provider.provider_enabled + || runtime.public_ipv6_provider.configured_prefix.is_some(), + "public IPv6 services", + )?; + require( + VPN_PORTAL_AVAILABLE, + peer.vpn_portal_cidr.is_some(), + "the VPN portal", + )?; + Ok(()) +} + +pub(super) fn validate(config: &CoreInstanceConfig) -> anyhow::Result<()> { + validate_snapshot(&config.connectivity.runtime, &config.peer.snapshot) +} + +pub(super) fn validate_runtime( + config: &crate::config::runtime::CoreInstanceRuntimeConfig, +) -> anyhow::Result<()> { + validate_snapshot(&config.services, &config.peer) +} diff --git a/easytier-core/src/instance/config.rs b/easytier-core/src/instance/config.rs new file mode 100644 index 00000000..e42c293c --- /dev/null +++ b/easytier-core/src/instance/config.rs @@ -0,0 +1,386 @@ +//! Portable normalization from the shared TOML model into one core instance. + +use std::collections::BTreeSet; + +use crate::{ + config::{ + IpPrefix, NodeConfig, ProxyNetworkConfig, RouteConfig, + gateway::{GatewayRuntimeConfig, ProxyRuntimeConfig}, + peers::{AclRuleConfig, HostRoutingPolicy, PublicIpv6ProviderConfig}, + runtime::CoreRuntimeConfig, + toml::{ConfigLoader as _, TomlConfig}, + }, + connectivity::{ + direct::DirectConnectorOptions, + manual::{ManualConnectorOptions, discovery::ManualEndpointDiscoveryConfig}, + stun::StunServerConfig, + }, + listener::plan::ListenerRuntimeConfig, + peers::{ + context::PeerRuntimeSnapshotInput, + peer_manager::{PortablePeerManagerConfig, RouteAlgoType}, + }, + socket::{NetNamespace, SocketContext, tcp::TcpBindOptions, udp::UdpBindOptions}, +}; + +use super::{CoreConnectivityConfig, CoreInstanceConfig}; + +const OSPF_UPDATE_MY_FOREIGN_NETWORK_INTERVAL_SEC: u64 = 10; +const MAX_DIRECT_CONNS_PER_PEER_IN_FOREIGN_NETWORK: usize = 3; + +/// Host facts and policy that cannot be derived from the shared TOML model. +/// +/// This input deliberately contains no routes, ACL, peer, gateway, listener, +/// or other portable configuration. Core combines it with TOML through one +/// normalization path for both initial construction and runtime patching. +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct CoreInstanceHostConfig { + pub hostname_fallback: Option, + pub host_routing: HostRoutingPolicy, + pub force_exit_node: bool, + pub allow_interface_bind: bool, + pub smoltcp_available: bool, + pub requires_smoltcp: bool, + pub icmp_failure_is_fatal: bool, + pub public_ipv6_provider_supported: bool, + pub gateway_enabled: bool, + pub easytier_version: String, + pub endpoint_protocols: Vec, +} + +impl Default for CoreInstanceHostConfig { + fn default() -> Self { + Self { + hostname_fallback: None, + host_routing: HostRoutingPolicy::default(), + force_exit_node: false, + allow_interface_bind: true, + smoltcp_available: false, + requires_smoltcp: false, + icmp_failure_is_fatal: false, + public_ipv6_provider_supported: false, + gateway_enabled: true, + easytier_version: env!("CARGO_PKG_VERSION").to_owned(), + endpoint_protocols: ManualEndpointDiscoveryConfig::default().srv_protocols, + } + } +} + +impl CoreInstanceConfig { + /// Normalizes the complete shared TOML model using OS-independent defaults. + /// + /// A Host may still project runtime facts such as a fallback hostname or + /// platform capability after parsing, but it does not need another network + /// configuration schema. + pub fn from_toml(config: &TomlConfig) -> anyhow::Result { + Self::from_toml_with_host(config, &CoreInstanceHostConfig::default()) + } + + /// Normalizes TOML with explicit Host facts and policy. + pub fn from_toml_with_host( + config: &TomlConfig, + host: &CoreInstanceHostConfig, + ) -> anyhow::Result { + let flags = config.get_flags(); + let instance_id = config.get_id(); + let identity: crate::config::NetworkIdentity = config.get_network_identity().into(); + let network_name = identity.network_name.clone(); + let socket_context = SocketContext::default() + .with_socket_mark(flags.socket_mark) + .with_netns(config.get_netns().map(NetNamespace::new)); + let hostname = match config.get_hostname() { + hostname if !hostname.is_empty() => hostname, + _ => host.hostname_fallback.clone().unwrap_or_default(), + }; + let acl = config.get_acl(); + + let peer_snapshot = + crate::config::peers::PeerRuntimeSnapshot::from_host_input(PeerRuntimeSnapshotInput { + node: NodeConfig { + peer_id: None, + instance_id: Some(*instance_id.as_bytes()), + hostname: (!hostname.is_empty()).then_some(hostname), + network_name: network_name.clone(), + }, + routes: RouteConfig { + ipv4: config.get_ipv4().map(|value| IpPrefix { + address: value.address().into(), + prefix_len: value.network_length(), + }), + ipv6: config.get_ipv6().map(|value| IpPrefix { + address: value.address().into(), + prefix_len: value.network_length(), + }), + proxy_networks: config + .get_proxy_cidrs() + .into_iter() + .map(|proxy| ProxyNetworkConfig { + real: IpPrefix { + address: proxy.cidr.first_address().into(), + prefix_len: proxy.cidr.network_length(), + }, + mapped: proxy.mapped_cidr.map(|mapped| IpPrefix { + address: mapped.first_address().into(), + prefix_len: mapped.network_length(), + }), + }) + .collect(), + ..Default::default() + }, + network_identity: identity, + stun_info: Default::default(), + flags: flags.clone(), + secure_mode: config.get_secure_mode(), + host_routing: host.host_routing, + acl: acl.clone(), + easytier_version: host.easytier_version.clone(), + vpn_portal_cidr: config + .get_vpn_portal_config() + .map(|portal| portal.client_cidr), + pinned_peers: config + .get_peers() + .into_iter() + .map(|peer| (peer.uri, peer.peer_public_key)) + .collect(), + ospf_update_my_foreign_network_interval_sec: + OSPF_UPDATE_MY_FOREIGN_NETWORK_INTERVAL_SEC, + max_direct_conns_per_peer_in_foreign_network: + MAX_DIRECT_CONNS_PER_PEER_IN_FOREIGN_NETWORK, + hmac_secret_digest: false, + }); + let peer = PortablePeerManagerConfig { + snapshot: peer_snapshot, + route_algo: RouteAlgoType::Ospf, + exit_nodes: config.get_exit_nodes(), + foreign_context_default_flags: TomlConfig::default().get_flags(), + }; + + let tcp_bind = TcpBindOptions::default().with_context(socket_context.clone()); + let udp_bind = UdpBindOptions::direct_connect().with_context(socket_context.clone()); + let listeners = Some(ListenerRuntimeConfig::new( + config.get_listener_uris(), + flags.enable_ipv6, + socket_context.clone(), + )); + let socks5_bind = config + .get_socks5_portal() + .map(|url| { + let host = url + .host_str() + .ok_or_else(|| anyhow::anyhow!("SOCKS5 portal host is missing"))?; + let port = url + .port() + .ok_or_else(|| anyhow::anyhow!("SOCKS5 portal port is missing"))?; + format!("{host}:{port}") + .parse() + .map_err(|error| anyhow::anyhow!("invalid SOCKS5 portal address: {error}")) + }) + .transpose()?; + let runtime = CoreRuntimeConfig { + acl: AclRuleConfig { + acl, + tcp_whitelist: config.get_tcp_whitelist(), + udp_whitelist: config.get_udp_whitelist(), + whitelist_priority: None, + }, + dhcp_ipv4: config.get_dhcp(), + gateway: GatewayRuntimeConfig { + socks5_bind, + port_forwards: config.get_port_forwards(), + }, + manual_routes: config + .get_routes() + .map(|routes| routes.into_iter().collect::>()), + proxy: ProxyRuntimeConfig { + enable_exit_node: flags.enable_exit_node || host.force_exit_node, + no_tun: flags.no_tun, + forward_by_system: flags.proxy_forward_by_system, + force_smoltcp: host.smoltcp_available + && (flags.use_smoltcp || flags.no_tun || host.requires_smoltcp), + icmp_failure_is_fatal: host.icmp_failure_is_fatal, + udp_response_ipv4_mtu: 1280, + }, + public_ipv6_auto: config.get_ipv6_public_addr_auto(), + public_ipv6_provider: PublicIpv6ProviderConfig { + provider_enabled: config.get_ipv6_public_addr_provider(), + configured_prefix: config.get_ipv6_public_addr_prefix(), + provider_supported: host.public_ipv6_provider_supported, + }, + }; + + Ok(Self { + instance_name: config.get_inst_name(), + peer, + connectivity: CoreConnectivityConfig { + initial_peers: config + .get_peers() + .into_iter() + .map(|peer| peer.uri) + .collect(), + listeners, + runtime, + startup_plan: super::CoreInstanceStartupPlan { + gateway: host.gateway_enabled, + }, + stun: StunServerConfig { + udp_servers: config + .get_stun_servers() + .unwrap_or_else(|| StunServerConfig::default().udp_servers), + udp_v6_servers: config + .get_stun_servers_v6() + .unwrap_or_else(|| StunServerConfig::default().udp_v6_servers), + ..StunServerConfig::default() + }, + endpoint_discovery: ManualEndpointDiscoveryConfig { + user_agent: format!("easytier/{}", host.easytier_version), + network_name: network_name.clone(), + http_tcp_bind: tcp_bind.clone(), + dns_record_context: socket_context, + srv_protocols: host.endpoint_protocols.clone(), + ..Default::default() + }, + manual: ManualConnectorOptions { + bind_device: flags.bind_device, + allow_interface_bind: host.allow_interface_bind, + tcp_bind: tcp_bind.clone(), + udp_bind: udp_bind.clone(), + ..Default::default() + }, + direct: DirectConnectorOptions { + default_protocol: flags.default_protocol, + enable_ipv6: flags.enable_ipv6, + allow_public_server: true, + bind_device: flags.bind_device, + allow_interface_bind: host.allow_interface_bind, + tcp_bind, + udp_bind, + testing: false, + }, + }, + }) + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn shared_toml_normalizes_instance_identity_and_connectivity() { + let config = TomlConfig::new_from_str( + r#" +instance_id = "018f4fb1-7a2c-7d1f-9d89-935b0ad7e135" +instance_name = "wasi-test" +hostname = "portable-host" +ipv4 = "10.144.0.2/24" +listeners = ["tcp://0.0.0.0:11010"] + +[[peer]] +uri = "tcp://127.0.0.1:11010" + +[network_identity] +network_name = "portable" +network_secret = "secret" + +[flags] +disable_p2p = true +"#, + ) + .unwrap(); + + let normalized = CoreInstanceConfig::from_toml(&config).unwrap(); + assert_eq!(normalized.instance_name, "wasi-test"); + assert_eq!( + normalized.peer.snapshot.runtime.core.node.instance_id, + Some(*config.get_id().as_bytes()) + ); + assert_eq!(normalized.peer.snapshot.runtime.core.node.peer_id, None); + assert_eq!( + normalized + .peer + .snapshot + .runtime + .core + .node + .hostname + .as_deref(), + Some("portable-host") + ); + assert_eq!( + normalized.peer.snapshot.runtime.core.node.network_name, + "portable" + ); + assert_eq!(normalized.connectivity.initial_peers.len(), 1); + assert_eq!( + normalized + .connectivity + .listeners + .as_ref() + .unwrap() + .urls + .len(), + 1 + ); + assert!(normalized.peer.snapshot.flags.disable_p2p); + } + + #[test] + fn host_config_supplies_only_platform_policy() { + let config = TomlConfig::default(); + let host = CoreInstanceHostConfig { + hostname_fallback: Some("host-fallback".to_owned()), + host_routing: HostRoutingPolicy { + local_exit_node_fallback: true, + }, + force_exit_node: true, + allow_interface_bind: false, + smoltcp_available: true, + requires_smoltcp: true, + icmp_failure_is_fatal: true, + public_ipv6_provider_supported: true, + gateway_enabled: false, + easytier_version: "host-version".to_owned(), + endpoint_protocols: vec!["host-protocol".to_owned()], + }; + + let normalized = CoreInstanceConfig::from_toml_with_host(&config, &host).unwrap(); + + assert_eq!( + normalized + .peer + .snapshot + .runtime + .core + .node + .hostname + .as_deref(), + Some("host-fallback") + ); + assert!( + normalized + .peer + .snapshot + .runtime + .host_routing + .local_exit_node_fallback + ); + assert_eq!(normalized.peer.snapshot.easytier_version, "host-version"); + assert!(normalized.connectivity.runtime.proxy.enable_exit_node); + assert!(normalized.connectivity.runtime.proxy.force_smoltcp); + assert!(normalized.connectivity.runtime.proxy.icmp_failure_is_fatal); + assert!( + normalized + .connectivity + .runtime + .public_ipv6_provider + .provider_supported + ); + assert!(!normalized.connectivity.startup_plan.gateway); + assert!(!normalized.connectivity.manual.allow_interface_bind); + assert!(!normalized.connectivity.direct.allow_interface_bind); + assert_eq!( + normalized.connectivity.endpoint_discovery.srv_protocols, + ["host-protocol"] + ); + } +} diff --git a/easytier-core/src/instance/data_plane_extension.rs b/easytier-core/src/instance/data_plane_extension.rs new file mode 100644 index 00000000..c956ef18 --- /dev/null +++ b/easytier-core/src/instance/data_plane_extension.rs @@ -0,0 +1,46 @@ +use std::{net::SocketAddr, sync::Arc, time::Duration}; + +use crate::gateway::{ + DataPlaneError, DataPlaneSession, DataPlaneTcpListener, DataPlaneTcpStream, DataPlaneUdpSocket, +}; + +use super::{CoreInstance, CoreInstanceHost}; + +impl CoreInstance +where + H: CoreInstanceHost, +{ + pub fn data_plane_session(&self) -> Arc> { + self.data_plane_session.clone() + } + + pub async fn data_plane_tcp_connect( + &self, + dst_addr: SocketAddr, + timeout: Duration, + ) -> Result { + self.data_plane_runtime + .data_plane_tcp_connect(dst_addr, timeout) + .await + } + + pub async fn data_plane_tcp_bind( + &self, + local_port: u16, + timeout: Duration, + ) -> Result { + self.data_plane_runtime + .data_plane_tcp_bind(local_port, timeout) + .await + } + + pub async fn data_plane_udp_bind( + &self, + local_port: u16, + timeout: Duration, + ) -> Result { + self.data_plane_runtime + .data_plane_udp_bind(local_port, timeout) + .await + } +} diff --git a/easytier-core/src/instance/lifecycle.rs b/easytier-core/src/instance/lifecycle.rs new file mode 100644 index 00000000..65b3c06f --- /dev/null +++ b/easytier-core/src/instance/lifecycle.rs @@ -0,0 +1,249 @@ +//! Serial start/stop orchestration for `CoreInstance`. +//! +//! Each runtime Module owns its internal resources. `CoreInstance` only +//! chooses their composition order and provides one outer cancellation and +//! cleanup path. + +use std::sync::{Arc, Weak}; + +#[cfg(feature = "dhcp-ipv4")] +use crate::gateway::dhcp::{DhcpIpv4Host, DhcpIpv4RouteSource}; + +use super::{CoreInstance, CoreInstanceHost, CoreInstanceState}; + +struct RecoveryGuard +where + F: FnOnce(), +{ + recovery: Option, +} + +impl RecoveryGuard +where + F: FnOnce(), +{ + fn new(recovery: F) -> Self { + Self { + recovery: Some(recovery), + } + } + + fn disarm(&mut self) { + self.recovery.take(); + } +} + +impl Drop for RecoveryGuard +where + F: FnOnce(), +{ + fn drop(&mut self) { + if let Some(recovery) = self.recovery.take() { + recovery(); + } + } +} + +impl CoreInstance +where + H: CoreInstanceHost, +{ + fn recovery_guard(self: &Arc) -> RecoveryGuard> { + let weak: Weak = Arc::downgrade(self); + RecoveryGuard::new(move || { + if let Some(instance) = weak.upgrade() { + tokio::spawn(async move { + instance.stop().await; + }); + } + }) + } + + async fn start_listener(&self) -> anyhow::Result<()> { + match &self.listener { + Some(listener) => listener.start().await, + None => Ok(()), + } + } + + #[cfg(feature = "dhcp-ipv4")] + async fn start_dhcp_ipv4(&self, host: Option>) -> anyhow::Result<()> { + if !self.runtime_config.snapshot().services.dhcp_ipv4 { + return Ok(()); + } + let host = host.ok_or_else(|| { + anyhow::anyhow!("DHCP IPv4 is enabled but no host adapter was provided") + })?; + let route_source: Arc = self.peer_manager.clone(); + self.dhcp_ipv4 + .start(route_source, self.runtime_config.clone(), host) + .await; + Ok(()) + } + + #[cfg(feature = "proxy-packet")] + async fn start_packet_proxy(&self) -> anyhow::Result<()> { + let config = self.runtime_config.snapshot(); + let has_proxy_networks = !config.peer.runtime.core.routes.proxy_networks.is_empty(); + if config.services.proxy.should_start(has_proxy_networks) { + self.packet_proxy + .start() + .await + .map_err(anyhow::Error::new)?; + } + Ok(()) + } + + async fn start_components(&self) -> anyhow::Result<()> { + #[cfg(feature = "public-ipv6-provider")] + self.public_ipv6_provider.validate_before_start().await?; + self.peer_manager + .get_route() + .set_route_cost_fn(self.peer_center.get_cost_calculator()) + .await; + + self.start_listener().await?; + if let Some(packet_egress) = &self.packet_egress { + packet_egress.start()?; + } + self.peer_manager.run().await.map_err(anyhow::Error::from)?; + self.direct.run(); + #[cfg(feature = "tcp-hole-punch")] + self.tcp_hole_punch.run(); + self.manual.start(); + #[cfg(feature = "public-ipv6-provider")] + self.public_ipv6_provider.start().await; + + let dhcp_host = self.instance_runtime.prepare(self.packet_plane()).await?; + #[cfg(feature = "dhcp-ipv4")] + self.start_dhcp_ipv4(dhcp_host).await?; + #[cfg(not(feature = "dhcp-ipv4"))] + let _ = dhcp_host; + #[cfg(feature = "wrapped-transport")] + if let Some(wrapped_transport) = &self.wrapped_transport { + wrapped_transport.start().await?; + } + #[cfg(feature = "proxy-packet")] + self.start_packet_proxy().await?; + self.udp_hole_punch.start().await?; + self.peer_center.init().await; + #[cfg(feature = "proxy-cidr-monitor")] + self.proxy_cidr_monitor + .start(&self.peer_manager, self.runtime_config.clone()) + .await; + #[cfg(feature = "vpn-portal")] + self.vpn_portal.start().await?; + #[cfg(feature = "proxy-smoltcp-stack")] + self.data_plane_runtime.start_runtime().await?; + #[cfg(feature = "proxy-smoltcp-stack")] + self.data_plane_session + .start() + .map_err(anyhow::Error::new)?; + #[cfg(feature = "proxy-smoltcp-stack")] + if self.startup_plan.gateway { + self.socks5_adapter.start().await?; + self.port_forward_adapter.start().await?; + } + Ok(()) + } + + async fn stop_components(&self) { + #[cfg(feature = "vpn-portal")] + self.vpn_portal.stop().await; + #[cfg(feature = "public-ipv6-provider")] + self.public_ipv6_provider.stop().await; + #[cfg(feature = "dhcp-ipv4")] + self.dhcp_ipv4.stop().await; + #[cfg(feature = "proxy-cidr-monitor")] + self.proxy_cidr_monitor.stop().await; + if let Some(listener) = &self.listener { + listener.stop().await; + } + self.udp_hole_punch.stop().await; + #[cfg(feature = "proxy-smoltcp-stack")] + self.port_forward_adapter.stop().await; + #[cfg(feature = "proxy-smoltcp-stack")] + self.socks5_adapter.stop().await; + #[cfg(feature = "proxy-smoltcp-stack")] + self.data_plane_session.stop(); + #[cfg(feature = "proxy-smoltcp-stack")] + self.data_plane_runtime.stop_runtime().await; + #[cfg(feature = "wrapped-transport")] + if let Some(wrapped_transport) = &self.wrapped_transport { + wrapped_transport.stop().await; + } + #[cfg(feature = "proxy-packet")] + self.packet_proxy.stop().await; + self.manual.stop().await; + #[cfg(feature = "tcp-hole-punch")] + self.tcp_hole_punch.stop().await; + self.direct.stop().await; + self.peer_center.stop().await; + + // Host packet tasks can still call the packet plane, so stop them + // before clearing PeerManager resources. + self.instance_runtime.shutdown().await; + self.peer_manager.clear_resources().await; + if let Some(packet_egress) = &self.packet_egress { + packet_egress.stop().await; + } + } + + /// Starts the complete instance through one serial composition path. + pub async fn start(self: &Arc) -> anyhow::Result<()> { + let _operation = self.operation.lock().await; + let state = self.state(); + if state != CoreInstanceState::Created { + anyhow::bail!("core instance cannot start from state {state:?}"); + } + + self.latest_error.write().take(); + self.set_state(CoreInstanceState::Starting); + let mut recovery = self.recovery_guard(); + let result = tokio::select! { + _ = self.cancel.cancelled() => { + Err(anyhow::anyhow!("core instance start cancelled")) + } + result = self.start_components() => result, + }; + let result = match result { + Ok(()) => { + self.set_state(CoreInstanceState::Running); + if self.cancel.is_cancelled() { + Err(anyhow::anyhow!("core instance start cancelled")) + } else { + Ok(()) + } + } + Err(error) => Err(error), + }; + + if let Err(error) = result { + self.latest_error.write().replace(format!("{error:#}")); + self.cancel.cancel(); + self.set_state(CoreInstanceState::Stopping); + self.stop_components().await; + self.set_state(CoreInstanceState::Stopped); + recovery.disarm(); + return Err(error); + } + + recovery.disarm(); + Ok(()) + } + + pub async fn stop(self: &Arc) { + self.cancel.cancel(); + let mut recovery = self.recovery_guard(); + let _operation = self.operation.lock().await; + if self.state() == CoreInstanceState::Stopped { + recovery.disarm(); + return; + } + + self.set_state(CoreInstanceState::Stopping); + self.stop_components().await; + self.set_state(CoreInstanceState::Stopped); + recovery.disarm(); + } +} diff --git a/easytier-core/src/instance/management.rs b/easytier-core/src/instance/management.rs new file mode 100644 index 00000000..7aa9a55e --- /dev/null +++ b/easytier-core/src/instance/management.rs @@ -0,0 +1,211 @@ +use std::{any::Any, net::IpAddr, sync::Arc}; + +use url::Url; + +use crate::{ + config::peers::AclWhitelistSnapshot, + connectivity::manual::ManualConnectorSnapshot, + foundation::stats::MetricSnapshot, + peers::{ + conn::peer_conn::PeerConnId, + credential_manager::{CredentialCreateOptions, CredentialInfo, GeneratedCredential}, + peer_manager::PeerSnapshot, + }, +}; + +use super::{CoreInstance, CoreInstanceHost, CorePacketPlane}; + +impl CoreInstance +where + H: CoreInstanceHost, +{ + pub fn instance_id(&self) -> uuid::Uuid { + self.peer_manager.instance_id() + } + + pub fn instance_name(&self) -> &str { + &self.instance_name + } + + pub fn add_connector(&self, url: Url) -> anyhow::Result<()> { + self.manual.add_connector(url) + } + + pub fn remove_connector(&self, url: &Url) -> bool { + self.manual.remove_connector(url) + } + + pub fn clear_connectors(&self) { + self.manual.clear_connectors(); + } + + pub fn list_connectors(&self) -> Vec { + self.manual.list_connectors() + } + + pub fn running_listeners(&self) -> Vec { + self.running_listeners.running_listeners() + } + + pub fn peer_id(&self) -> crate::config::PeerId { + self.peer_manager.my_peer_id() + } + + pub fn packet_plane(&self) -> Arc { + self.packet_plane.clone() + } + + /// Recovers the concrete Host runtime Adapter for host-specific integration. + pub fn runtime_host(&self) -> Option<&T> { + let runtime_host: &dyn Any = self.instance_runtime.as_ref(); + runtime_host.downcast_ref() + } + + pub fn attach_tun_fd(&self, fd: i32) -> anyhow::Result<()> { + self.instance_runtime.attach_tun_fd(fd) + } + + pub fn latest_error(&self) -> Option { + self.latest_error.read().clone() + } + + pub fn is_ready(&self) -> bool { + self.state() == super::CoreInstanceState::Running + } + + pub fn management_events(&self) -> Vec { + self.instance_runtime.management_events() + } + + pub fn global_peer_map_snapshot(&self) -> crate::proto::peer_rpc::GetGlobalPeerMapResponse { + self.peer_center.global_peer_map_snapshot() + } + + pub async fn peer_snapshots(&self) -> Vec { + self.peer_manager.list_peer_snapshots().await + } + + pub async fn node_snapshot(&self) -> crate::peers::peer_manager::NodeSnapshot { + let mut snapshot = self + .peer_manager + .node_snapshot(self.running_listeners()) + .await; + snapshot.ip_list = self + .direct + .local_address_observations_with_stun(&snapshot.stun_info) + .await; + snapshot + } + + pub async fn route_snapshots(&self) -> Vec { + self.peer_manager.list_route_snapshots().await + } + + pub async fn dump_route(&self) -> String { + self.peer_manager.dump_route().await + } + + pub async fn local_public_ipv6_info( + &self, + ) -> crate::proto::core_peer::peer::ListPublicIpv6InfoResponse { + self.peer_manager.local_public_ipv6_info().await + } + + pub async fn foreign_network_route_infos( + &self, + ) -> crate::proto::peer_rpc::RouteForeignNetworkInfos { + self.peer_manager.foreign_network_route_infos().await + } + + pub async fn foreign_network_snapshots( + &self, + include_trusted_keys: bool, + ) -> std::collections::HashMap + { + self.peer_manager + .list_foreign_network_infos(include_trusted_keys) + .await + } + + pub async fn foreign_network_route_summary( + &self, + ) -> crate::proto::peer_rpc::RouteForeignNetworkSummary { + self.peer_manager.foreign_network_route_summary().await + } + + pub fn acl_stats(&self) -> crate::proto::acl::AclStats { + self.peer_manager.acl_stats() + } + + pub fn acl_whitelist_snapshot(&self) -> AclWhitelistSnapshot { + let config = self.runtime_config.snapshot(); + AclWhitelistSnapshot::from(&config.services.acl) + } + + pub fn generate_credential( + &self, + options: CredentialCreateOptions, + ) -> anyhow::Result { + if !self.peer_manager.can_manage_credentials() { + anyhow::bail!("only admin nodes (with network_secret) can generate credentials"); + } + if options.ttl.is_zero() { + anyhow::bail!("ttl_seconds must be positive"); + } + let generated = self + .peer_manager + .credential_manager() + .generate_credential_with_options( + options.groups, + options.allow_relay, + options.allowed_proxy_cidrs, + options.ttl, + options.credential_id, + options.reusable, + ); + self.peer_manager.notify_credential_changed(); + Ok(generated) + } + + pub fn revoke_credential(&self, credential_id: &str) -> anyhow::Result { + if !self.peer_manager.can_manage_credentials() { + anyhow::bail!("only admin nodes (with network_secret) can revoke credentials"); + } + let revoked = self + .peer_manager + .credential_manager() + .revoke_credential(credential_id); + if revoked { + self.peer_manager.notify_credential_changed(); + } + Ok(revoked) + } + + pub fn credential_snapshots(&self) -> Vec { + self.peer_manager.credential_manager().list_credentials() + } + + pub fn metric_snapshots(&self) -> Vec { + self.peer_manager.stats_manager().get_all_metrics() + } + + pub fn prometheus_metrics(&self) -> String { + self.peer_manager.stats_manager().export_prometheus() + } + + pub async fn close_peer_conn( + &self, + peer_id: crate::config::PeerId, + conn_id: &PeerConnId, + ) -> Result<(), crate::peers::error::Error> { + self.peer_manager.close_peer_conn(peer_id, conn_id).await + } + + pub async fn update_exit_nodes(&self, exit_nodes: Vec) { + self.peer_manager.update_exit_nodes(exit_nodes).await; + } + + pub async fn refresh_acl_groups(&self) { + self.peer_manager.get_route().refresh_acl_groups().await; + } +} diff --git a/easytier-core/src/instance/management_extension.rs b/easytier-core/src/instance/management_extension.rs new file mode 100644 index 00000000..93c0f46b --- /dev/null +++ b/easytier-core/src/instance/management_extension.rs @@ -0,0 +1,22 @@ +use std::sync::Arc; + +use crate::config::runtime::CoreInstanceRuntimeConfig; + +use super::{CoreInstance, CoreInstanceHost, CoreInstanceHostConfig}; + +impl CoreInstance +where + H: CoreInstanceHost, +{ + pub fn toml_config(&self) -> Option { + self.management.toml_config() + } + + pub(crate) fn runtime_config_snapshot(&self) -> Arc { + self.runtime_config.snapshot() + } + + pub(crate) fn host_config(&self) -> &CoreInstanceHostConfig { + self.management.host_config() + } +} diff --git a/easytier-core/src/instance/management_state.rs b/easytier-core/src/instance/management_state.rs new file mode 100644 index 00000000..6f67af3c --- /dev/null +++ b/easytier-core/src/instance/management_state.rs @@ -0,0 +1,34 @@ +use crate::{config::toml::TomlConfig, instance::CoreInstanceHostConfig}; + +pub(super) struct ManagementState { + #[cfg(feature = "management")] + toml_config: Option, + #[cfg(feature = "management")] + host_config: CoreInstanceHostConfig, +} + +impl ManagementState { + pub(super) fn new( + toml_config: Option, + host_config: CoreInstanceHostConfig, + ) -> Self { + #[cfg(not(feature = "management"))] + let _ = (toml_config, host_config); + Self { + #[cfg(feature = "management")] + toml_config, + #[cfg(feature = "management")] + host_config, + } + } + + #[cfg(feature = "management")] + pub(super) fn toml_config(&self) -> Option { + self.toml_config.clone() + } + + #[cfg(feature = "management")] + pub(super) fn host_config(&self) -> &CoreInstanceHostConfig { + &self.host_config + } +} diff --git a/easytier-core/src/instance/manager.rs b/easytier-core/src/instance/manager.rs new file mode 100644 index 00000000..b57891b0 --- /dev/null +++ b/easytier-core/src/instance/manager.rs @@ -0,0 +1,664 @@ +//! Canonical process-level ownership and lifecycle for EasyTier instances. + +use std::{ + collections::{HashMap, hash_map::Entry}, + error::Error, + fmt, + path::PathBuf, + sync::{ + Arc, Mutex, + atomic::{AtomicUsize, Ordering}, + }, +}; + +use dashmap::DashMap; +use uuid::Uuid; + +use crate::config::toml::TomlConfig; +use crate::instance::{CoreInstance, CoreInstanceHost}; +use crate::process_runtime::CoreProcessRuntime; +#[cfg(feature = "management")] +use crate::{ + config::toml::{ConfigLoader as _, ConfigSource}, + management::network_instance_running_info, +}; +#[cfg(feature = "management")] +use easytier_proto::api::manage::NetworkInstanceRunningInfo; + +/// Stable identity required by the instance collection. +pub trait ManagedInstance: Send + Sync + 'static { + fn instance_id(&self) -> Uuid; +} + +impl ManagedInstance for CoreInstance +where + H: CoreInstanceHost, +{ + fn instance_id(&self) -> Uuid { + self.instance_id() + } +} + +/// Host-specific construction seam for one complete instance record. +pub trait InstanceFactory: Send + Sync + 'static { + type Instance: ManagedInstance; + type CreateContext; + type Error; + + fn create( + &self, + config: TomlConfig, + context: Self::CreateContext, + ) -> Result, Self::Error>; +} + +/// Error returned while constructing or registering an instance. +#[derive(Debug)] +pub enum InstanceCreateError { + Factory(E), + AlreadyExists { instance_id: Uuid }, +} + +impl fmt::Display for InstanceCreateError +where + E: fmt::Display, +{ + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + match self { + Self::Factory(error) => write!(formatter, "failed to create instance: {error:#}"), + Self::AlreadyExists { instance_id } => { + write!(formatter, "instance {instance_id} already exists") + } + } + } +} + +impl Error for InstanceCreateError where E: fmt::Debug + fmt::Display {} + +/// Supplies the process runtime used by every instance created by this +/// factory. +pub trait ProcessRuntimeProvider: InstanceFactory { + fn process_runtime(&self) -> Arc; +} + +#[derive(Clone, Copy, Default)] +pub struct ConfigFilePermission(u8); + +impl ConfigFilePermission { + pub const READ_ONLY: u8 = 1 << 0; + pub const NO_DELETE: u8 = 1 << 1; + + pub fn with_flag(self, flag: u8) -> Self { + Self(self.0 | flag) + } + + pub fn remove_flag(self, flag: u8) -> Self { + Self(self.0 & !flag) + } + + pub fn has_flag(&self, flag: u8) -> bool { + self.0 & flag != 0 + } +} + +impl From for ConfigFilePermission { + fn from(value: u8) -> Self { + Self(value) + } +} + +impl From for ConfigFilePermission { + fn from(value: u32) -> Self { + Self(value as u8) + } +} + +impl From for u8 { + fn from(value: ConfigFilePermission) -> Self { + value.0 + } +} + +impl From for u32 { + fn from(value: ConfigFilePermission) -> Self { + value.0 as u32 + } +} + +impl fmt::Debug for ConfigFilePermission { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + let access = if self.has_flag(Self::READ_ONLY) { + "READ_ONLY" + } else { + "EDITABLE" + }; + let deletion = if self.has_flag(Self::NO_DELETE) { + "NO_DELETE" + } else { + "DELETABLE" + }; + write!(formatter, "{access}|{deletion}") + } +} + +#[derive(Debug, Clone)] +pub struct ConfigFileControl { + pub path: Option, + pub permission: ConfigFilePermission, +} + +impl ConfigFileControl { + pub const STATIC_CONFIG: Self = Self { + path: None, + permission: ConfigFilePermission( + ConfigFilePermission::READ_ONLY | ConfigFilePermission::NO_DELETE, + ), + }; + + pub fn new(path: Option, permission: ConfigFilePermission) -> Self { + Self { path, permission } + } + + pub fn is_read_only(&self) -> bool { + self.permission.has_flag(ConfigFilePermission::READ_ONLY) + } + + pub fn set_read_only(&mut self, read_only: bool) { + self.permission = if read_only { + self.permission.with_flag(ConfigFilePermission::READ_ONLY) + } else { + self.permission.remove_flag(ConfigFilePermission::READ_ONLY) + }; + } + + pub fn is_no_delete(&self) -> bool { + self.permission.has_flag(ConfigFilePermission::NO_DELETE) + } + + pub fn set_no_delete(&mut self, no_delete: bool) { + self.permission = if no_delete { + self.permission.with_flag(ConfigFilePermission::NO_DELETE) + } else { + self.permission.remove_flag(ConfigFilePermission::NO_DELETE) + }; + } + + pub fn is_deletable(&self) -> bool { + !self.is_no_delete() + } +} + +pub struct DaemonGuard { + guard: Option>, + notifier: Arc, +} + +impl Drop for DaemonGuard { + fn drop(&mut self) { + drop(self.guard.take()); + self.notifier.notify_one(); + } +} + +struct ActiveStopGuard { + active_stops: Arc, + notifier: Arc, +} + +impl Drop for ActiveStopGuard { + fn drop(&mut self) { + let previous = self.active_stops.fetch_sub(1, Ordering::AcqRel); + debug_assert!(previous > 0); + self.notifier.notify_one(); + } +} + +/// Owns all process-level Instance state and operations. +pub struct InstanceManager { + factory: F, + instances: Mutex>>, + config_controls: DashMap, + notifier: Arc, + config_dir: Option, + daemon_guard: Arc<()>, + mutation_lock: Arc>, + runtime_handle: Option, + active_stops: Arc, +} + +impl InstanceManager { + pub fn new(factory: F, runtime_handle: Option) -> Self { + Self { + factory, + instances: Mutex::new(HashMap::new()), + config_controls: DashMap::new(), + notifier: Arc::new(tokio::sync::Notify::new()), + config_dir: None, + daemon_guard: Arc::new(()), + mutation_lock: Arc::new(tokio::sync::Mutex::new(())), + runtime_handle, + active_stops: Arc::new(AtomicUsize::new(0)), + } + } + + pub fn with_config_path(mut self, config_dir: Option) -> Self { + self.config_dir = config_dir; + self + } + + pub fn create( + &self, + config: TomlConfig, + context: F::CreateContext, + ) -> Result, InstanceCreateError> { + let instance = self + .factory + .create(config, context) + .map_err(InstanceCreateError::Factory)?; + let instance_id = instance.instance_id(); + let mut instances = self.instances.lock().expect("instance map lock poisoned"); + + match instances.entry(instance_id) { + Entry::Vacant(entry) => { + entry.insert(instance.clone()); + Ok(instance) + } + Entry::Occupied(_) => Err(InstanceCreateError::AlreadyExists { instance_id }), + } + } + + pub(crate) fn get(&self, instance_id: Uuid) -> Option> { + self.instances + .lock() + .expect("instance map lock poisoned") + .get(&instance_id) + .cloned() + } + + pub(crate) fn list(&self) -> Vec> { + self.instances + .lock() + .expect("instance map lock poisoned") + .values() + .cloned() + .collect() + } + + pub(crate) fn remove(&self, instance_id: Uuid) -> Option> { + self.instances + .lock() + .expect("instance map lock poisoned") + .remove(&instance_id) + } + + pub fn mutation_lock(&self) -> Arc> { + self.mutation_lock.clone() + } + + pub fn config_dir(&self) -> Option<&PathBuf> { + self.config_dir.as_ref() + } + + pub fn register_daemon(&self) -> DaemonGuard { + DaemonGuard { + guard: Some(self.daemon_guard.clone()), + notifier: self.notifier.clone(), + } + } +} + +impl InstanceManager { + pub fn process_runtime(&self) -> Arc { + self.factory.process_runtime() + } +} + +impl InstanceManager +where + F: InstanceFactory, CreateContext = ()>, + F::Error: fmt::Debug + fmt::Display + Send + Sync + 'static, + H: CoreInstanceHost, +{ + pub fn run_network_instance( + &self, + config: TomlConfig, + control: ConfigFileControl, + ) -> anyhow::Result { + let runtime = self + .runtime_handle + .clone() + .or_else(|| tokio::runtime::Handle::try_current().ok()) + .ok_or_else(|| anyhow::anyhow!("tokio runtime not found, cannot start instance"))?; + let instance = self.create(config, ()).map_err(anyhow::Error::new)?; + let instance_id = instance.instance_id(); + self.config_controls.insert(instance_id, control); + let notifier = self.notifier.clone(); + runtime.spawn(async move { + if let Err(error) = instance.start().await { + tracing::error!(%error, %instance_id, "instance failed to start"); + } + notifier.notify_one(); + }); + Ok(instance_id) + } + + pub async fn delete_network_instances( + &self, + instance_ids: impl IntoIterator, + ) -> anyhow::Result> { + let runtime = self + .runtime_handle + .clone() + .or_else(|| tokio::runtime::Handle::try_current().ok()) + .ok_or_else(|| anyhow::anyhow!("tokio runtime not found, cannot stop instance"))?; + self.active_stops.fetch_add(1, Ordering::AcqRel); + let active_stop = ActiveStopGuard { + active_stops: self.active_stops.clone(), + notifier: self.notifier.clone(), + }; + let mut removed = Vec::new(); + for instance_id in instance_ids { + self.config_controls.remove(&instance_id); + if let Some(instance) = self.remove(instance_id) { + removed.push(instance); + } + } + if removed.is_empty() { + drop(active_stop); + return Ok(self.instance_ids()); + } + + runtime + .spawn(async move { + let _active_stop = active_stop; + for instance in removed { + instance.stop().await; + } + }) + .await + .map_err(|error| anyhow::anyhow!("instance stop task failed: {error}"))?; + Ok(self.instance_ids()) + } + + pub async fn retain_network_instances(&self, retained: &[Uuid]) -> anyhow::Result> { + let removed = self + .list() + .into_iter() + .map(|instance| instance.instance_id()) + .filter(|instance_id| !retained.contains(instance_id)) + .collect::>(); + self.delete_network_instances(removed).await + } + + pub fn instance_ids(&self) -> Vec { + self.list() + .into_iter() + .map(|instance| instance.instance_id()) + .collect() + } + + pub fn instance(&self, instance_id: Uuid) -> Option>> { + self.get(instance_id) + } + + pub fn instances(&self) -> Vec>> { + self.list() + } + + pub fn config_control(&self, instance_id: Uuid) -> Option { + self.config_controls + .get(&instance_id) + .map(|control| control.clone()) + } + + pub fn attach_tun_fd(&self, instance_id: Uuid, fd: i32) -> anyhow::Result<()> { + self.get(instance_id) + .ok_or_else(|| anyhow::anyhow!("instance {instance_id} not found"))? + .attach_tun_fd(fd) + } + + pub fn data_plane_runtime_handle(&self, instance_id: &Uuid) -> Option { + self.instance(*instance_id)?; + self.runtime_handle + .clone() + .or_else(|| tokio::runtime::Handle::try_current().ok()) + } + + pub async fn wait(&self) { + loop { + let instance_running = self + .list() + .iter() + .any(|instance| instance.state() != crate::instance::CoreInstanceState::Stopped); + let daemon_running = Arc::strong_count(&self.daemon_guard) > 1; + let instance_stopping = self.active_stops.load(Ordering::Acquire) != 0; + if !instance_running && !instance_stopping && !daemon_running { + return; + } + self.notifier.notified().await; + } + } + + #[cfg(feature = "management")] + pub fn config(&self, instance_id: Uuid) -> Option { + self.get(instance_id) + .and_then(|instance| instance.toml_config()) + } + + #[cfg(feature = "management")] + pub fn config_source(&self, instance_id: Uuid) -> Option { + self.config(instance_id) + .map(|config| config.get_network_config_source()) + } + + #[cfg(feature = "management")] + pub async fn network_info(&self, instance_id: Uuid) -> Option { + let instance = self.get(instance_id)?; + network_instance_running_info(instance.as_ref()).await.ok() + } + + #[cfg(feature = "management")] + pub async fn collect_network_infos( + &self, + ) -> anyhow::Result> { + let mut result = std::collections::BTreeMap::new(); + for instance in self.list() { + result.insert( + instance.instance_id(), + network_instance_running_info(instance.as_ref()).await?, + ); + } + Ok(result) + } + + #[cfg(feature = "management")] + pub fn collect_network_infos_sync( + &self, + ) -> anyhow::Result> { + self.runtime_handle + .as_ref() + .ok_or_else(|| anyhow::anyhow!("InstanceManager runtime handle is unavailable"))? + .block_on(self.collect_network_infos()) + } + + #[cfg(feature = "proxy-smoltcp-stack")] + pub fn data_plane_session( + &self, + instance_id: &Uuid, + ) -> Option>> { + self.instance(*instance_id) + .map(|instance| instance.data_plane_session()) + } + + #[cfg(feature = "proxy-smoltcp-stack")] + pub async fn data_plane_tcp_connect( + &self, + instance_id: &Uuid, + dst_addr: std::net::SocketAddr, + timeout: std::time::Duration, + ) -> anyhow::Result { + Ok(self + .instance(*instance_id) + .ok_or_else(|| anyhow::anyhow!("instance {instance_id} not found"))? + .data_plane_tcp_connect(dst_addr, timeout) + .await?) + } + + #[cfg(feature = "proxy-smoltcp-stack")] + pub async fn data_plane_tcp_bind( + &self, + instance_id: &Uuid, + local_port: u16, + timeout: std::time::Duration, + ) -> anyhow::Result { + Ok(self + .instance(*instance_id) + .ok_or_else(|| anyhow::anyhow!("instance {instance_id} not found"))? + .data_plane_tcp_bind(local_port, timeout) + .await?) + } + + #[cfg(feature = "proxy-smoltcp-stack")] + pub async fn data_plane_udp_bind( + &self, + instance_id: &Uuid, + local_port: u16, + timeout: std::time::Duration, + ) -> anyhow::Result { + Ok(self + .instance(*instance_id) + .ok_or_else(|| anyhow::anyhow!("instance {instance_id} not found"))? + .data_plane_udp_bind(local_port, timeout) + .await?) + } +} + +#[cfg(test)] +mod tests { + use std::{ + sync::{ + Arc, + atomic::{AtomicUsize, Ordering}, + }, + thread, + }; + + use super::*; + struct TestFactory { + drops: Arc, + } + + #[derive(Debug)] + struct TestInstance { + id: Uuid, + drops: Arc, + } + + impl ManagedInstance for TestInstance { + fn instance_id(&self) -> Uuid { + self.id + } + } + + impl Drop for TestInstance { + fn drop(&mut self) { + self.drops.fetch_add(1, Ordering::SeqCst); + } + } + + impl InstanceFactory for TestFactory { + type Instance = TestInstance; + type CreateContext = (); + type Error = std::convert::Infallible; + + fn create( + &self, + config: TomlConfig, + (): Self::CreateContext, + ) -> Result, Self::Error> { + Ok(Arc::new(TestInstance { + id: config.get_id(), + drops: self.drops.clone(), + })) + } + } + + fn manager() -> (Arc>, Arc) { + let drops = Arc::new(AtomicUsize::new(0)); + ( + Arc::new(InstanceManager::new( + TestFactory { + drops: drops.clone(), + }, + None, + )), + drops, + ) + } + + fn config(instance_id: Uuid) -> TomlConfig { + let config = TomlConfig::default(); + config.set_id(instance_id); + config + } + + #[test] + fn duplicate_create_drops_losing_complete_record() { + let (manager, drops) = manager(); + let instance_id = Uuid::new_v4(); + let first = manager.create(config(instance_id), ()).unwrap(); + let error = manager.create(config(instance_id), ()).unwrap_err(); + + assert!(matches!( + error, + InstanceCreateError::AlreadyExists { + instance_id: duplicate + } if duplicate == instance_id + )); + assert_eq!(drops.load(Ordering::SeqCst), 1); + assert!(Arc::ptr_eq(&first, &manager.get(instance_id).unwrap())); + } + + #[test] + fn concurrent_duplicate_create_registers_exactly_one_instance() { + let (manager, drops) = manager(); + let instance_id = Uuid::new_v4(); + let workers = (0..2) + .map(|_| { + let manager = manager.clone(); + thread::spawn(move || manager.create(config(instance_id), ())) + }) + .collect::>(); + let results = workers + .into_iter() + .map(|worker| worker.join().unwrap()) + .collect::>(); + + assert_eq!(results.iter().filter(|result| result.is_ok()).count(), 1); + assert_eq!(results.iter().filter(|result| result.is_err()).count(), 1); + assert_eq!(drops.load(Ordering::SeqCst), 1); + assert_eq!(manager.list().len(), 1); + } + + #[test] + fn list_is_an_arc_snapshot_and_remove_returns_exact_stored_value() { + let (manager, drops) = manager(); + let first_id = Uuid::new_v4(); + let second_id = Uuid::new_v4(); + let first = manager.create(config(first_id), ()).unwrap(); + let second = manager.create(config(second_id), ()).unwrap(); + + let snapshot = manager.list(); + let removed = manager.remove(first_id).unwrap(); + + assert!(Arc::ptr_eq(&first, &removed)); + assert!(snapshot.iter().any(|item| Arc::ptr_eq(item, &removed))); + assert!(manager.get(first_id).is_none()); + assert!(Arc::ptr_eq(&second, &manager.get(second_id).unwrap())); + drop(removed); + drop(first); + assert_eq!(drops.load(Ordering::SeqCst), 0); + drop(snapshot); + assert_eq!(drops.load(Ordering::SeqCst), 1); + } +} diff --git a/easytier-core/src/instance/mod.rs b/easytier-core/src/instance/mod.rs new file mode 100644 index 00000000..c92a7120 --- /dev/null +++ b/easytier-core/src/instance/mod.rs @@ -0,0 +1,897 @@ +//! Lifecycle owner for the portable EasyTier runtime. + +mod build_capabilities; +mod config; +#[cfg(feature = "proxy-smoltcp-stack")] +mod data_plane_extension; +mod lifecycle; +mod management; +#[cfg(feature = "management")] +mod management_extension; +mod management_state; +pub mod manager; +mod packet_io; +mod packet_plane; +#[cfg(feature = "proxy-packet")] +mod packet_proxy_extension; +#[cfg(feature = "public-ipv6-provider")] +mod public_ipv6_extension; +#[cfg(feature = "vpn-portal")] +mod vpn_portal_extension; + +use std::sync::{ + Arc, + atomic::{AtomicU8, Ordering}, +}; + +use parking_lot::RwLock; +use serde::{Deserialize, Serialize}; +#[cfg(feature = "test-utils")] +use std::sync::atomic::AtomicUsize; +use tokio::sync::Mutex; +use tokio_util::sync::CancellationToken; +use url::Url; + +#[cfg(feature = "tcp-hole-punch")] +use crate::connectivity::hole_punch::tcp::TcpHolePunchConnector; +use crate::{ + config::peers::{AclRuleConfig, PeerRuntimeSnapshot}, + config::runtime::{CoreInstanceRuntimeConfig, CoreRuntimeConfig, CoreRuntimeConfigStore}, + config::toml::TomlConfig, + connectivity::hole_punch::port_mapping::UdpPortMappingPlatform, + connectivity::hole_punch::tcp::TcpHolePunchHost, + connectivity::stun::{ + StunDnsRuntime, StunInfoCollector, StunInfoProvider, StunServerConfig, StunSocketMapper, + }, + connectivity::{ + direct::{ + DirectConnectorHost, DirectConnectorManager, DirectConnectorOptions, + ForeignDirectConnectorRpcRegistrar, + }, + hole_punch::udp::CoreUdpHolePunchService, + manual::{ + ManualConnectorManager, ManualConnectorOptions, + discovery::{CoreManualEndpointResolver, ManualEndpointDiscoveryConfig}, + }, + protocol::{ + ClientProtocolUpgrader, CoreClientProtocolConfig, CoreClientProtocolUpgrader, + ServerProtocolUpgrader, + }, + }, + events::CoreEventSink, + gateway::dhcp::DhcpIpv4Host, + host::dns::{DnsRecordResolver, DnsResolver}, + listener::{ + AcceptedSocketHandler, ExternalListenerFactory, ExternalListenerRequest, ListenerFactory, + RunningListenerRegistry, + plan::{ListenerRuntimeConfig, PreparedListenerPlan, prepare_listener_plan}, + transport::{ + AcceptedTransport, CoreListenerRuntime, HostAcceptedTcpSocket, + ProtocolAcceptedTransportHandler, + }, + }, + peers::peer_center::instance::PeerCenterInstance, + peers::{ + admission::{PeerAcceptedTunnelHandler, RawAcceptedTransportHandler}, + context::PeerStunInfoSource, + create_packet_recv_chan, + credential_manager::CredentialStorage, + peer_manager::{PeerManagerCore, PortablePeerManagerConfig}, + public_ipv6::{CorePublicIpv6Runtime, PublicIpv6Host}, + }, + process_runtime::CoreProcessRuntime, + socket::{tcp::VirtualTcpSocketFactory, udp::VirtualUdpSocketFactory}, +}; + +use crate::gateway::proxy::cidr_table::{ProxyCidrSnapshot, ProxyCidrTable}; +#[cfg(feature = "proxy-packet")] +use crate::gateway::proxy::icmp_host::IcmpProxyHost; +#[cfg(feature = "wrapped-transport")] +use crate::gateway::proxy::wrapped_transport::WrappedTransportEngines; +#[cfg(feature = "vpn-portal")] +use crate::gateway::vpn_portal::VpnPortalHost; + +#[cfg(feature = "public-ipv6-provider")] +use crate::peers::public_ipv6::provider::PublicIpv6ProviderPlatform; + +#[cfg(feature = "dhcp-ipv4")] +use crate::gateway::dhcp::DhcpIpv4Runtime; +#[cfg(feature = "proxy-cidr-monitor")] +use crate::gateway::proxy::cidr_monitor::ProxyCidrMonitorRuntime; +#[cfg(feature = "proxy-packet")] +use crate::gateway::proxy::service::CoreProxyModule; +#[cfg(feature = "wrapped-transport")] +use crate::gateway::proxy::wrapped_transport::WrappedTransportProxyModule; +#[cfg(feature = "vpn-portal")] +use crate::gateway::vpn_portal::VpnPortalModule; +#[cfg(feature = "proxy-smoltcp-stack")] +use crate::gateway::{ + DataPlaneRuntime, DataPlaneSession, PortForwardAdapter, Socks5GatewayAdapter, +}; +use crate::host::packet::PacketSink; +#[cfg(feature = "public-ipv6-provider")] +use crate::peers::public_ipv6::provider::PublicIpv6ProviderRuntime; +pub use config::CoreInstanceHostConfig; +use management_state::ManagementState; +use packet_io::PacketEgress; +pub use packet_plane::CorePacketPlane; + +/// Complete Host capability set required by one portable core instance. +pub trait CoreInstanceHost: DirectConnectorHost + TcpHolePunchHost {} + +impl CoreInstanceHost for T where T: DirectConnectorHost + TcpHolePunchHost {} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +#[repr(u8)] +pub enum CoreInstanceState { + Created, + Starting, + Running, + Stopping, + Stopped, +} + +impl CoreInstanceState { + fn from_u8(value: u8) -> Self { + match value { + value if value == Self::Created as u8 => Self::Created, + value if value == Self::Starting as u8 => Self::Starting, + value if value == Self::Running as u8 => Self::Running, + value if value == Self::Stopping as u8 => Self::Stopping, + value if value == Self::Stopped as u8 => Self::Stopped, + _ => unreachable!("invalid core instance state"), + } + } +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] +pub struct CoreInstanceStartupPlan { + pub gateway: bool, +} + +impl CoreInstanceStartupPlan { + fn is_default(&self) -> bool { + self == &Self::default() + } +} + +impl Default for CoreInstanceStartupPlan { + fn default() -> Self { + Self { gateway: true } + } +} + +#[derive(Debug, Clone, Serialize, Deserialize, Default)] +pub struct CoreConnectivityConfig { + pub initial_peers: Vec, + pub listeners: Option, + pub runtime: CoreRuntimeConfig, + #[serde(default, skip_serializing_if = "CoreInstanceStartupPlan::is_default")] + pub startup_plan: CoreInstanceStartupPlan, + pub stun: StunServerConfig, + pub endpoint_discovery: ManualEndpointDiscoveryConfig, + pub manual: ManualConnectorOptions, + pub direct: DirectConnectorOptions, +} + +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct CoreInstanceConfig { + #[serde(default = "crate::config::toml::default_instance_name")] + pub instance_name: String, + pub peer: PortablePeerManagerConfig, + pub connectivity: CoreConnectivityConfig, +} + +#[cfg(any(test, feature = "test-utils"))] +#[doc(hidden)] +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub struct PeerRelaySessionSnapshot { + pub has_state: bool, + pub has_session: bool, +} + +fn validate_core_instance_config( + config: &CoreInstanceConfig, +) -> anyhow::Result> { + let acl = config.connectivity.runtime.acl.build()?; + build_capabilities::validate(config)?; + Ok(acl) +} + +fn proxy_cidr_snapshot(config: &CoreInstanceRuntimeConfig) -> ProxyCidrSnapshot { + ProxyCidrSnapshot::from_proxy_networks(&config.peer.runtime.core.routes.proxy_networks) +} + +fn retain_core_peer_identity( + peer: &mut Arc, + peer_id: crate::config::PeerId, + instance_id: Option<[u8; 16]>, +) { + let peer = Arc::make_mut(peer); + peer.runtime.core.node.peer_id = Some(peer_id); + peer.runtime.core.node.instance_id = instance_id; +} + +/// Host-owned resources that must be prepared for the complete Instance +/// lifetime, such as a native packet interface. +#[async_trait::async_trait] +pub trait InstanceRuntimeHost: std::any::Any + Send + Sync + 'static { + async fn prepare( + &self, + packet_plane: Arc, + ) -> anyhow::Result>>; + + async fn shutdown(&self); + + /// Requests prompt Host cleanup when the canonical instance owner is + /// dropped without an opportunity to await [`Self::shutdown`]. + fn request_shutdown(&self) {} + + /// Returns the bounded, serialized event journal exposed by process + /// management. Hosts that do not produce events keep the default empty + /// journal. + fn management_events(&self) -> Vec { + Vec::new() + } + + /// Applies Host-side cached views of fields already committed to the + /// shared TOML model. + #[cfg(feature = "management")] + fn synchronize_config(&self, _patch: &crate::proto::api::config::InstanceConfigPatch) {} + + #[cfg(feature = "management")] + fn publish_config_patch(&self, _patch: crate::proto::api::config::InstanceConfigPatch) {} + + fn attach_tun_fd(&self, _fd: i32) -> anyhow::Result<()> { + anyhow::bail!("external TUN attachment is not supported by this Host") + } +} + +#[async_trait::async_trait] +impl InstanceRuntimeHost for () { + async fn prepare( + &self, + _packet_plane: Arc, + ) -> anyhow::Result>> { + Ok(None) + } + + async fn shutdown(&self) {} +} + +/// Host Adapters and optional native capabilities for one core instance. +/// +/// Callers provide this bundle and one normalized [`CoreInstanceConfig`] to +/// [`CoreInstance::new`]. Core constructs and owns every portable runtime +/// Module behind that seam. +pub struct CoreHostAdapters +where + H: CoreInstanceHost, +{ + host: Arc, + /// Host OS policy and build capabilities used during configuration + /// normalization. It contains no portable TOML-derived state. + pub config: CoreInstanceHostConfig, + #[cfg(any(test, feature = "test-utils"))] + /// Optional construction-time STUN provider used by deterministic tests. + stun_override: Option::Socket>>>, + dns: Arc, + process_runtime: Arc, + packet_sink: Arc, + pub instance_runtime: Arc, + pub events: Arc, + pub credential_storage: Option>, + #[cfg(feature = "wrapped-transport")] + pub wrapped_transports: WrappedTransportEngines, + pub protocol: Option::Socket>>>, + pub external_listener_factory: + Option>>>>, + pub server_protocol: Option>>>, + /// Optional OS port-mapping adapter. STUN-only hole punching remains + /// available when the host does not provide one. + pub udp_hole_punch_platform: Option>, + #[cfg(feature = "proxy-packet")] + pub icmp_proxy_host: Option>, + #[cfg(feature = "proxy-cidr-monitor")] + pub proxy_cidr_monitor_enabled: bool, + #[cfg(feature = "public-ipv6-provider")] + pub public_ipv6_host: Option>, + #[cfg(feature = "public-ipv6-provider")] + pub public_ipv6_provider: Option>, + #[cfg(feature = "vpn-portal")] + pub vpn_portal: Option>, +} + +impl CoreHostAdapters +where + H: CoreInstanceHost, +{ + /// Creates the minimal host bundle. Optional native capabilities can be + /// installed on the returned value before constructing the instance. + pub fn new( + host: Arc, + dns: Arc, + packet_sink: Arc, + process_runtime: Arc, + ) -> Self { + Self { + host, + config: CoreInstanceHostConfig::default(), + #[cfg(any(test, feature = "test-utils"))] + stun_override: None, + dns, + process_runtime, + packet_sink, + instance_runtime: Arc::new(()), + events: Arc::new(()), + credential_storage: None, + #[cfg(feature = "wrapped-transport")] + wrapped_transports: WrappedTransportEngines::default(), + protocol: None, + external_listener_factory: None, + server_protocol: None, + udp_hole_punch_platform: None, + #[cfg(feature = "proxy-packet")] + icmp_proxy_host: None, + #[cfg(feature = "proxy-cidr-monitor")] + proxy_cidr_monitor_enabled: false, + #[cfg(feature = "public-ipv6-provider")] + public_ipv6_host: None, + #[cfg(feature = "public-ipv6-provider")] + public_ipv6_provider: None, + #[cfg(feature = "vpn-portal")] + vpn_portal: None, + } + } +} + +struct CoreStunPeerInfoSource(Arc); + +impl PeerStunInfoSource for CoreStunPeerInfoSource { + fn stun_info(&self) -> crate::proto::common::StunInfo { + self.0.get_stun_info() + } +} + +/// Owns the portable peer and connectivity runtime for one EasyTier instance. +/// +/// An instance is intentionally one-shot: after it is stopped, construct a new +/// instance rather than trying to rebuild partially consumed peer-manager state. +pub struct CoreInstance +where + H: CoreInstanceHost, +{ + instance_name: String, + #[allow(dead_code)] + management: ManagementState, + pub(super) instance_runtime: Arc, + state: AtomicU8, + latest_error: RwLock>, + pub(super) operation: Mutex<()>, + pub(super) cancel: CancellationToken, + pub(super) peer_manager: Arc, + packet_plane: Arc, + pub(super) manual: ManualConnectorManager, + pub(super) direct: DirectConnectorManager, + #[cfg(feature = "tcp-hole-punch")] + tcp_hole_punch: TcpHolePunchConnector, + pub(super) listener: Option>>, + running_listeners: Arc, + pub(super) udp_hole_punch: CoreUdpHolePunchService, + #[cfg(feature = "wrapped-transport")] + wrapped_transport: Option>, + #[cfg(feature = "proxy-smoltcp-stack")] + data_plane_runtime: Arc>, + #[cfg(feature = "proxy-smoltcp-stack")] + data_plane_session: Arc>, + #[cfg(feature = "proxy-smoltcp-stack")] + socks5_adapter: Arc>, + #[cfg(feature = "proxy-smoltcp-stack")] + port_forward_adapter: Arc>, + proxy_cidr_table: Arc, + #[cfg(feature = "proxy-packet")] + packet_proxy: Arc>, + #[cfg(feature = "proxy-cidr-monitor")] + proxy_cidr_monitor: ProxyCidrMonitorRuntime, + #[cfg(feature = "dhcp-ipv4")] + dhcp_ipv4: DhcpIpv4Runtime, + pub(super) packet_egress: Option, + pub(super) peer_center: Arc, + #[cfg(feature = "public-ipv6-provider")] + public_ipv6_provider: PublicIpv6ProviderRuntime, + #[cfg(feature = "vpn-portal")] + vpn_portal: Arc, + #[cfg(feature = "proxy-smoltcp-stack")] + pub(super) startup_plan: CoreInstanceStartupPlan, + pub(super) runtime_config: CoreRuntimeConfigStore, + #[cfg(feature = "test-utils")] + acl_reload_count: AtomicUsize, +} + +impl CoreInstance +where + H: CoreInstanceHost, +{ + fn prepare_stun( + adapters: &CoreHostAdapters, + config: &CoreConnectivityConfig, + ) -> Arc::Socket>> { + #[cfg(any(test, feature = "test-utils"))] + if let Some(stun_override) = &adapters.stun_override { + return stun_override.clone(); + } + Arc::new(StunInfoCollector::new_with_socket_contexts( + adapters.host.clone(), + adapters.dns.clone(), + config.direct.udp_bind.context.clone(), + config.direct.tcp_bind.context.clone(), + config.stun.udp_servers.clone(), + config.stun.tcp_servers.clone(), + config.stun.udp_v6_servers.clone(), + )) + } + + /// Constructs the complete portable runtime for one EasyTier instance. + /// + /// This is the only instance construction entry. The normalized config is + /// authoritative after creation; all platform behavior enters through the + /// supplied Host Adapters. + pub fn new( + config: CoreInstanceConfig, + adapters: CoreHostAdapters, + ) -> anyhow::Result> { + let host_config = adapters.config.clone(); + Self::new_inner(config, None, host_config, adapters) + } + + /// Constructs an instance from the shared TOML model and retains that + /// model as the authoritative management configuration. + pub fn from_toml( + toml_config: TomlConfig, + adapters: CoreHostAdapters, + ) -> anyhow::Result> { + let host_config = adapters.config.clone(); + let config = CoreInstanceConfig::from_toml_with_host(&toml_config, &host_config)?; + Self::new_inner(config, Some(toml_config), host_config, adapters) + } + + fn new_inner( + config: CoreInstanceConfig, + toml_config: Option, + host_config: CoreInstanceHostConfig, + mut adapters: CoreHostAdapters, + ) -> anyhow::Result> { + let initial_acl = validate_core_instance_config(&config)?; + let instance_name = config.instance_name; + let (packet_tx, packet_rx) = create_packet_recv_chan(); + let runtime_config = CoreRuntimeConfigStore::new( + config.connectivity.runtime.clone(), + Arc::new(config.peer.snapshot.clone()), + ); + let events = adapters.events.clone(); + #[cfg(feature = "public-ipv6-provider")] + let public_ipv6_host: Arc = adapters + .public_ipv6_host + .take() + .unwrap_or_else(|| Arc::new(())); + #[cfg(not(feature = "public-ipv6-provider"))] + let public_ipv6_host: Arc = Arc::new(()); + #[cfg(feature = "public-ipv6-provider")] + let public_ipv6_events = events.clone(); + #[cfg(not(feature = "public-ipv6-provider"))] + let public_ipv6_events: Arc = Arc::new(()); + let public_ipv6_runtime = CorePublicIpv6Runtime::new( + runtime_config.clone(), + public_ipv6_host, + public_ipv6_events, + ); + let stun = Self::prepare_stun(&adapters, &config.connectivity); + let peer_stun: Arc = stun.clone(); + let foreign_rpc_registrar = Arc::new(ForeignDirectConnectorRpcRegistrar::new( + adapters.host.clone(), + stun.clone(), + )); + let peer_manager = Arc::new(PeerManagerCore::new( + config.peer, + runtime_config.clone(), + Arc::new(CoreStunPeerInfoSource(peer_stun)), + packet_tx, + public_ipv6_runtime.clone(), + events.clone(), + adapters.credential_storage.take(), + foreign_rpc_registrar, + )?); + peer_manager.reload_acl(initial_acl.as_ref()); + let config = config.connectivity; + let listener_plan = prepare_listener_plan( + config.listeners.as_ref(), + peer_manager.instance_id(), + adapters.server_protocol.as_deref(), + adapters.external_listener_factory.as_deref(), + )?; + let CoreHostAdapters { + host, + config: _, + #[cfg(any(test, feature = "test-utils"))] + stun_override: _, + dns, + process_runtime, + packet_sink, + instance_runtime, + events, + credential_storage: _, + #[cfg(feature = "wrapped-transport")] + wrapped_transports, + protocol, + external_listener_factory, + server_protocol, + udp_hole_punch_platform, + #[cfg(feature = "proxy-packet")] + icmp_proxy_host, + #[cfg(feature = "proxy-cidr-monitor")] + proxy_cidr_monitor_enabled, + #[cfg(feature = "public-ipv6-provider")] + public_ipv6_host: _, + #[cfg(feature = "public-ipv6-provider")] + public_ipv6_provider, + #[cfg(feature = "vpn-portal")] + vpn_portal, + } = adapters; + let dns_records: Arc = dns.clone(); + let dns: Arc = dns; + let ring_registry = process_runtime.ring_registry(); + let protected_tcp_ports = process_runtime.protected_tcp_ports(); + let CoreConnectivityConfig { + initial_peers, + listeners: _, + runtime: _, + startup_plan, + stun: _, + endpoint_discovery, + manual: manual_options, + direct: direct_options, + } = config; + #[cfg(not(feature = "proxy-smoltcp-stack"))] + let _ = startup_plan; + let accepted_transport_handler: Arc< + dyn AcceptedSocketHandler>>, + > = match server_protocol { + Some(server_protocol) => { + let tunnel_handler = PeerAcceptedTunnelHandler::new(&peer_manager, events.clone()); + Arc::new(ProtocolAcceptedTransportHandler::new( + &tunnel_handler, + server_protocol, + )) + } + None => Arc::new(RawAcceptedTransportHandler::new(&peer_manager)), + }; + let running_listeners = Arc::new(RunningListenerRegistry::default()); + let PreparedListenerPlan { + transports, + external, + failures, + } = listener_plan; + let mut external_factories = Vec::with_capacity(external.len()); + if !external.is_empty() && external_listener_factory.is_none() { + anyhow::bail!("listener plan requires an external listener factory"); + } + for (listener, socket_context) in external { + let factory = external_listener_factory.clone().unwrap(); + let request = ExternalListenerRequest { + url: listener.url, + socket_context, + }; + external_factories.push(ListenerFactory::new( + move || factory.create(request.clone()), + listener.must_succeed, + )); + } + let has_listener_work = + !transports.is_empty() || !external_factories.is_empty() || !failures.is_empty(); + let listener = has_listener_work.then(|| { + Arc::new(CoreListenerRuntime::new_with_events( + host.clone(), + dns.clone(), + ring_registry.clone(), + transports, + external_factories, + failures, + accepted_transport_handler, + events.clone(), + running_listeners.clone(), + )) + }); + let protocol = protocol.unwrap_or_else(|| { + Arc::new(CoreClientProtocolUpgrader::new( + CoreClientProtocolConfig::default(), + )) + }); + let endpoint_resolver = Arc::new(CoreManualEndpointResolver::new( + host.clone(), + dns.clone(), + dns_records, + endpoint_discovery, + )); + let manual = ManualConnectorManager::new( + peer_manager.clone(), + host.clone(), + dns.clone(), + endpoint_resolver, + protocol.clone(), + ring_registry, + manual_options, + events.clone(), + ); + for url in initial_peers { + manual.add_connector(url)?; + } + let udp_hole_punch_socket_context = direct_options.udp_bind.context.clone(); + let udp_hole_punch = CoreUdpHolePunchService::new( + peer_manager.clone(), + host.clone(), + stun.clone(), + udp_hole_punch_platform, + events.clone(), + udp_hole_punch_socket_context, + protocol.clone(), + ); + let proxy_cidr_table = Arc::new(ProxyCidrTable::from_snapshot(proxy_cidr_snapshot( + runtime_config.snapshot().as_ref(), + ))); + #[cfg(feature = "wrapped-transport")] + let tcp_proxy_socket_context = direct_options.tcp_bind.context.clone(); + #[cfg(feature = "proxy-packet")] + let packet_proxy = CoreProxyModule::new( + peer_manager.clone(), + host.clone(), + protected_tcp_ports.clone(), + running_listeners.clone(), + runtime_config.clone(), + proxy_cidr_table.clone(), + tcp_proxy_socket_context.clone(), + direct_options.udp_bind.context.clone(), + // Raw ICMP shares the datagram/network-layer routing context. + direct_options.udp_bind.context.clone(), + icmp_proxy_host, + ); + #[cfg(feature = "wrapped-transport")] + let wrapped_transport = { + let WrappedTransportEngines { kcp, quic } = wrapped_transports; + WrappedTransportProxyModule::new( + peer_manager.clone(), + runtime_config.clone(), + kcp, + quic, + host.clone(), + protected_tcp_ports.clone(), + running_listeners.clone(), + proxy_cidr_table.clone(), + tcp_proxy_socket_context, + ) + }; + #[cfg(feature = "proxy-smoltcp-stack")] + let data_plane_runtime = DataPlaneRuntime::new( + runtime_config.clone(), + peer_manager.clone(), + wrapped_transport.as_ref(), + host.clone(), + direct_options.tcp_bind.context.clone(), + ); + #[cfg(feature = "proxy-smoltcp-stack")] + let data_plane_session = DataPlaneSession::new(&data_plane_runtime); + #[cfg(feature = "proxy-smoltcp-stack")] + let socks5_adapter = Socks5GatewayAdapter::new( + runtime_config.clone(), + data_plane_runtime.clone(), + host.clone(), + dns.clone(), + direct_options.tcp_bind.context.clone(), + ); + #[cfg(feature = "proxy-smoltcp-stack")] + let port_forward_adapter = PortForwardAdapter::new( + runtime_config.clone(), + data_plane_runtime.clone(), + host.clone(), + direct_options.tcp_bind.context.clone(), + events.clone(), + ); + #[cfg(feature = "tcp-hole-punch")] + let tcp_hole_punch = TcpHolePunchConnector::new( + peer_manager.clone(), + host.clone(), + stun.clone(), + direct_options.tcp_bind.context.clone(), + protocol.clone(), + Arc::new(crate::connectivity::protocol::CoreServerProtocolUpgrader::< + HostAcceptedTcpSocket, + >::new( + crate::connectivity::protocol::CoreServerProtocolConfig::default(), + )), + ); + let direct = DirectConnectorManager::new_with_running_listeners( + peer_manager.clone(), + host.clone(), + protected_tcp_ports, + stun.clone(), + running_listeners.clone(), + dns, + protocol, + direct_options, + ); + let peer_center = Arc::new(PeerCenterInstance::new(peer_manager.clone())); + #[cfg(feature = "public-ipv6-provider")] + let public_ipv6_provider = PublicIpv6ProviderRuntime::new( + public_ipv6_provider, + runtime_config.clone(), + public_ipv6_runtime, + ); + #[cfg(feature = "vpn-portal")] + let vpn_portal = VpnPortalModule::new( + peer_manager.clone(), + runtime_config.clone(), + vpn_portal, + events.clone(), + ); + #[cfg(feature = "proxy-cidr-monitor")] + let proxy_cidr_monitor = + ProxyCidrMonitorRuntime::new(proxy_cidr_monitor_enabled, events.clone()); + #[cfg(feature = "proxy-cidr-monitor")] + let proxy_cidr_monitor_available = proxy_cidr_monitor.is_enabled(); + #[cfg(not(feature = "proxy-cidr-monitor"))] + let proxy_cidr_monitor_available = false; + let packet_plane = Arc::new(CorePacketPlane::new( + peer_manager.clone(), + runtime_config.clone(), + proxy_cidr_monitor_available, + )); + + Ok(Arc::new(Self { + instance_name, + management: ManagementState::new(toml_config, host_config), + instance_runtime, + state: AtomicU8::new(CoreInstanceState::Created as u8), + latest_error: RwLock::new(None), + operation: Mutex::new(()), + cancel: CancellationToken::new(), + peer_manager, + packet_plane, + manual, + direct, + #[cfg(feature = "tcp-hole-punch")] + tcp_hole_punch, + listener, + running_listeners, + udp_hole_punch, + #[cfg(feature = "wrapped-transport")] + wrapped_transport, + #[cfg(feature = "proxy-smoltcp-stack")] + data_plane_runtime, + #[cfg(feature = "proxy-smoltcp-stack")] + data_plane_session, + #[cfg(feature = "proxy-smoltcp-stack")] + socks5_adapter, + #[cfg(feature = "proxy-smoltcp-stack")] + port_forward_adapter, + proxy_cidr_table, + #[cfg(feature = "proxy-packet")] + packet_proxy, + #[cfg(feature = "proxy-cidr-monitor")] + proxy_cidr_monitor, + #[cfg(feature = "dhcp-ipv4")] + dhcp_ipv4: DhcpIpv4Runtime::new(), + packet_egress: Some(PacketEgress::new(packet_rx, packet_sink)), + peer_center, + #[cfg(feature = "public-ipv6-provider")] + public_ipv6_provider, + #[cfg(feature = "vpn-portal")] + vpn_portal, + #[cfg(feature = "proxy-smoltcp-stack")] + startup_plan, + runtime_config, + #[cfg(feature = "test-utils")] + acl_reload_count: AtomicUsize::new(0), + })) + } + + pub fn state(&self) -> CoreInstanceState { + CoreInstanceState::from_u8(self.state.load(Ordering::Acquire)) + } + + fn set_state(&self, state: CoreInstanceState) { + self.state.store(state as u8, Ordering::Release); + } + + async fn reload_acl_config_inner(&self, config: &AclRuleConfig) -> anyhow::Result<()> { + let acl = config.build()?; + self.peer_manager.reload_acl(acl.as_ref()); + #[cfg(feature = "test-utils")] + self.acl_reload_count.fetch_add(1, Ordering::Relaxed); + Ok(()) + } + + fn sync_peer_runtime_state(&self, snapshot: &PeerRuntimeSnapshot) { + self.peer_manager + .set_avoid_relay_data_preference(snapshot.avoid_relay_data_preference); + } + + /// Publishes one complete instance configuration version. Host changes have + /// no effect until submitted through this method. + pub async fn update_runtime_config( + &self, + config: CoreInstanceRuntimeConfig, + ) -> anyhow::Result<()> { + let _operation = self.operation.lock().await; + self.update_runtime_config_under_operation(config).await + } + + pub(crate) async fn update_runtime_config_under_operation( + &self, + mut config: CoreInstanceRuntimeConfig, + ) -> anyhow::Result<()> { + if matches!( + self.state(), + CoreInstanceState::Stopping | CoreInstanceState::Stopped + ) { + anyhow::bail!("runtime config cannot update while instance is stopping or stopped"); + } + self.validate_runtime_config_capabilities(&config)?; + let current = self.runtime_config.snapshot(); + retain_core_peer_identity( + &mut config.peer, + self.peer_id(), + current.peer.runtime.core.node.instance_id, + ); + let refresh_acl_groups = current.peer.peer_group_memberships + != config.peer.peer_group_memberships + || current.peer.acl_group_declarations != config.peer.acl_group_declarations; + if current.services.acl != config.services.acl { + self.reload_acl_config_inner(&config.services.acl).await?; + } + self.sync_peer_runtime_state(&config.peer); + self.runtime_config.replace(config); + self.proxy_cidr_table + .update_snapshot(proxy_cidr_snapshot(self.runtime_config.snapshot().as_ref())); + if refresh_acl_groups { + self.refresh_acl_groups().await; + } + #[cfg(feature = "proxy-smoltcp-stack")] + self.port_forward_adapter + .reload( + &self + .runtime_config + .snapshot() + .services + .gateway + .port_forwards, + ) + .await?; + Ok(()) + } + + pub(crate) fn validate_runtime_config_capabilities( + &self, + config: &CoreInstanceRuntimeConfig, + ) -> anyhow::Result<()> { + build_capabilities::validate_runtime(config) + } + + pub async fn wait(&self) { + self.peer_manager.wait().await; + } +} + +impl Drop for CoreInstance +where + H: CoreInstanceHost, +{ + fn drop(&mut self) { + self.cancel.cancel(); + self.instance_runtime.request_shutdown(); + } +} + +#[cfg(any(test, feature = "test-utils"))] +mod test_utils; + +#[cfg(test)] +mod tests; diff --git a/easytier-core/src/instance/packet_io.rs b/easytier-core/src/instance/packet_io.rs new file mode 100644 index 00000000..e5e49a10 --- /dev/null +++ b/easytier-core/src/instance/packet_io.rs @@ -0,0 +1,190 @@ +use std::{ + net::{IpAddr, Ipv4Addr, Ipv6Addr}, + sync::{Arc, Mutex}, +}; + +use tokio::{sync::mpsc, task::JoinHandle}; + +use crate::host::packet::PacketSink; +use crate::packet::ZCPacket; + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub(crate) struct IpPacketMeta { + pub(crate) source: IpAddr, + pub(crate) destination: IpAddr, +} + +pub(crate) fn parse_ip_packet(packet: &[u8]) -> anyhow::Result { + let Some(version) = packet.first().map(|byte| byte >> 4) else { + anyhow::bail!("IP packet is empty"); + }; + match version { + 4 => parse_ipv4_packet(packet), + 6 => parse_ipv6_packet(packet), + _ => anyhow::bail!("unsupported IP version: {version}"), + } +} + +fn parse_ipv4_packet(packet: &[u8]) -> anyhow::Result { + if packet.len() < 20 { + anyhow::bail!("IPv4 packet is shorter than the minimum header"); + } + let header_len = usize::from(packet[0] & 0x0f) * 4; + let total_len = usize::from(u16::from_be_bytes([packet[2], packet[3]])); + if header_len < 20 || total_len < header_len || total_len != packet.len() { + anyhow::bail!("invalid IPv4 header or total length"); + } + Ok(IpPacketMeta { + source: IpAddr::V4(Ipv4Addr::new( + packet[12], packet[13], packet[14], packet[15], + )), + destination: IpAddr::V4(Ipv4Addr::new( + packet[16], packet[17], packet[18], packet[19], + )), + }) +} + +fn parse_ipv6_packet(packet: &[u8]) -> anyhow::Result { + if packet.len() < 40 { + anyhow::bail!("IPv6 packet is shorter than the fixed header"); + } + let payload_len = usize::from(u16::from_be_bytes([packet[4], packet[5]])); + if (payload_len == 0 && packet.len() != 40) + || (payload_len != 0 && 40 + payload_len != packet.len()) + { + anyhow::bail!("IPv6 payload length does not match the packet"); + } + Ok(IpPacketMeta { + source: IpAddr::V6(Ipv6Addr::from( + <[u8; 16]>::try_from(&packet[8..24]).expect("checked IPv6 header length"), + )), + destination: IpAddr::V6(Ipv6Addr::from( + <[u8; 16]>::try_from(&packet[24..40]).expect("checked IPv6 header length"), + )), + }) +} + +pub(crate) struct PacketEgress { + receiver: Mutex>>, + sink: Arc, + task: Mutex>>, +} + +impl PacketEgress { + pub(crate) fn new(receiver: mpsc::Receiver, sink: Arc) -> Self { + Self { + receiver: Mutex::new(Some(receiver)), + sink, + task: Mutex::new(None), + } + } + + pub(crate) fn start(&self) -> anyhow::Result<()> { + let mut receiver = self + .receiver + .lock() + .unwrap() + .take() + .ok_or_else(|| anyhow::anyhow!("packet egress is one-shot and already started"))?; + let sink = self.sink.clone(); + let task = tokio::spawn(async move { + while let Some(packet) = receiver.recv().await { + if let Err(error) = sink.write_packet(packet.payload().to_vec()).await { + tracing::warn!(?error, "host packet sink rejected an egress packet"); + } + } + }); + *self.task.lock().unwrap() = Some(task); + Ok(()) + } + + pub(crate) async fn stop(&self) { + let task = self.task.lock().unwrap().take(); + if let Some(task) = task { + task.abort(); + let _ = task.await; + } + self.receiver.lock().unwrap().take(); + } +} + +impl Drop for PacketEgress { + fn drop(&mut self) { + if let Some(task) = self.task.lock().unwrap().take() { + task.abort(); + } + } +} + +#[cfg(test)] +mod tests { + use crate::foundation::time::{Duration, timeout}; + + use super::*; + + #[test] + fn parses_ipv4_and_ipv6_packet_endpoints() { + let mut ipv4 = vec![0u8; 20]; + ipv4[0] = 0x45; + ipv4[2..4].copy_from_slice(&20u16.to_be_bytes()); + ipv4[12..16].copy_from_slice(&[10, 1, 0, 1]); + ipv4[16..20].copy_from_slice(&[10, 2, 0, 1]); + assert_eq!( + parse_ip_packet(&ipv4).unwrap(), + IpPacketMeta { + source: "10.1.0.1".parse().unwrap(), + destination: "10.2.0.1".parse().unwrap(), + } + ); + + let mut ipv6 = vec![0u8; 40]; + ipv6[0] = 0x60; + ipv6[8..24].copy_from_slice(&"fd00::1".parse::().unwrap().octets()); + ipv6[24..40].copy_from_slice(&"fd00::2".parse::().unwrap().octets()); + assert_eq!( + parse_ip_packet(&ipv6).unwrap(), + IpPacketMeta { + source: "fd00::1".parse().unwrap(), + destination: "fd00::2".parse().unwrap(), + } + ); + } + + #[test] + fn rejects_truncated_or_unknown_ip_packets() { + assert!(parse_ip_packet(&[]).is_err()); + assert!(parse_ip_packet(&[0x70]).is_err()); + assert!(parse_ip_packet(&[0x45; 19]).is_err()); + assert!(parse_ip_packet(&[0x60; 39]).is_err()); + + let mut ipv4_with_trailing_bytes = vec![0u8; 21]; + ipv4_with_trailing_bytes[0] = 0x45; + ipv4_with_trailing_bytes[2..4].copy_from_slice(&20u16.to_be_bytes()); + assert!(parse_ip_packet(&ipv4_with_trailing_bytes).is_err()); + + let mut unsupported_ipv6_jumbogram = vec![0u8; 41]; + unsupported_ipv6_jumbogram[0] = 0x60; + assert!(parse_ip_packet(&unsupported_ipv6_jumbogram).is_err()); + } + + #[tokio::test] + async fn packet_egress_forwards_to_host_sink_and_joins_on_stop() { + let (core_tx, core_rx) = mpsc::channel(1); + let (host_tx, mut host_rx) = mpsc::channel(1); + let egress = PacketEgress::new(core_rx, Arc::new(host_tx)); + egress.start().unwrap(); + + core_tx + .send(ZCPacket::new_with_payload(b"packet")) + .await + .unwrap(); + let packet = timeout(Duration::from_secs(1), host_rx.recv()) + .await + .expect("packet egress did not forward to the host") + .expect("host packet channel closed"); + assert_eq!(packet, b"packet"); + + egress.stop().await; + assert!(egress.start().is_err()); + } +} diff --git a/easytier-core/src/instance/packet_plane.rs b/easytier-core/src/instance/packet_plane.rs new file mode 100644 index 00000000..26c3bffc --- /dev/null +++ b/easytier-core/src/instance/packet_plane.rs @@ -0,0 +1,143 @@ +use std::{collections::BTreeSet, net::IpAddr, sync::Arc}; + +use async_trait::async_trait; + +use crate::{ + config::runtime::CoreRuntimeConfigStore, + gateway::magic_dns::{MagicDnsRouteSnapshot, MagicDnsRouteSource}, + gateway::proxy::cidr_monitor::{ProxyCidrDiff, collect_proxy_cidr_diff}, + peers::peer_manager::PeerManagerCore, +}; + +#[cfg(feature = "proxy-packet")] +use crate::foundation::stats::{LabelSet, LabelType, MetricName}; +#[cfg(feature = "proxy-packet")] +use crate::gateway::magic_dns::{ + MagicDnsQueryResolver, MagicDnsResolverRegistration, magic_dns_packet_filter, +}; +#[cfg(feature = "proxy-packet")] +use crate::gateway::udp_broadcast::UdpBroadcastRelayStats; + +use super::packet_io::parse_ip_packet; + +/// Stable packet- and route-plane projection for platform integrations. +pub struct CorePacketPlane { + peer_manager: Arc, + runtime_config: CoreRuntimeConfigStore, + proxy_cidr_monitor_available: bool, +} + +impl CorePacketPlane { + pub(super) fn new( + peer_manager: Arc, + runtime_config: CoreRuntimeConfigStore, + proxy_cidr_monitor_available: bool, + ) -> Self { + Self { + peer_manager, + runtime_config, + proxy_cidr_monitor_available, + } + } + + pub async fn send_ip_packet(&self, packet: Vec) -> anyhow::Result<()> { + let meta = parse_ip_packet(&packet)?; + let source_is_local = self.peer_manager.is_local_virtual_ip(&meta.source); + if matches!(meta.source, IpAddr::V6(ip) if ip.is_unicast_link_local()) && !source_is_local { + return Ok(()); + } + self.peer_manager + .send_msg_by_ip( + crate::packet::ZCPacket::new_with_payload(&packet), + meta.destination, + source_is_local, + ) + .await + .map_err(Into::into) + } + + pub async fn send_local_ip_packet(&self, packet: Vec) -> anyhow::Result<()> { + let destination = parse_ip_packet(&packet)?.destination; + self.peer_manager + .send_msg_by_ip( + crate::packet::ZCPacket::new_with_payload(&packet), + destination, + true, + ) + .await + .map_err(Into::into) + } + + pub async fn proxy_cidr_diff( + &self, + previous: &BTreeSet, + ) -> Option { + if !self.proxy_cidr_monitor_available { + return None; + } + Some( + collect_proxy_cidr_diff(self.peer_manager.as_ref(), &self.runtime_config, previous) + .await, + ) + } + + pub async fn public_ipv6_routes(&self) -> BTreeSet { + self.peer_manager.list_public_ipv6_routes().await + } + + pub async fn public_ipv6_addr(&self) -> Option { + self.peer_manager.public_ipv6_addr().await + } + + #[cfg(feature = "proxy-packet")] + pub fn udp_broadcast_relay_stats(&self) -> UdpBroadcastRelayStats { + let network_name = self + .runtime_config + .snapshot() + .peer + .runtime + .network_identity + .network_name + .clone(); + let labels = LabelSet::new().with_label_type(LabelType::NetworkName(network_name)); + let stats = self.peer_manager.stats_manager(); + UdpBroadcastRelayStats::new( + stats.get_counter(MetricName::UdpBroadcastRelayPacketsCaptured, labels.clone()), + stats.get_counter(MetricName::UdpBroadcastRelayPacketsIgnored, labels.clone()), + stats.get_counter( + MetricName::UdpBroadcastRelayPacketsForwarded, + labels.clone(), + ), + stats.get_counter(MetricName::UdpBroadcastRelayPacketsForwardFailed, labels), + ) + } + + #[cfg(feature = "proxy-packet")] + pub async fn register_magic_dns_resolver( + &self, + fake_ip: std::net::Ipv4Addr, + resolver: Arc, + ) -> MagicDnsResolverRegistration { + let runtime = tokio::runtime::Handle::current(); + let pipeline = self + .peer_manager + .add_managed_nic_packet_process_pipeline(magic_dns_packet_filter( + fake_ip, + self.peer_manager.my_peer_id(), + resolver, + )) + .await; + MagicDnsResolverRegistration::new(Arc::downgrade(&self.peer_manager), pipeline, runtime) + } +} + +#[async_trait] +impl MagicDnsRouteSource for CorePacketPlane { + async fn snapshot(&self) -> MagicDnsRouteSnapshot { + MagicDnsRouteSource::snapshot(self.peer_manager.as_ref()).await + } + + async fn revision(&self) -> quanta::Instant { + MagicDnsRouteSource::revision(self.peer_manager.as_ref()).await + } +} diff --git a/easytier-core/src/instance/packet_proxy_extension.rs b/easytier-core/src/instance/packet_proxy_extension.rs new file mode 100644 index 00000000..2111d76c --- /dev/null +++ b/easytier-core/src/instance/packet_proxy_extension.rs @@ -0,0 +1,41 @@ +use crate::gateway::proxy::{ + tcp_proxy_engine::TcpNatEntrySnapshot, + wrapped_transport::{WrappedTransportKind, WrappedTransportRole}, +}; + +use super::{CoreInstance, CoreInstanceHost}; + +impl CoreInstance +where + H: CoreInstanceHost, +{ + pub fn tcp_proxy_entry_snapshots(&self) -> Vec { + self.packet_proxy.tcp_entry_snapshots() + } + + pub fn wrapped_tcp_proxy_entry_snapshots( + &self, + transport: WrappedTransportKind, + role: WrappedTransportRole, + ) -> Vec { + self.wrapped_transport + .as_ref() + .map_or_else(Vec::new, |proxy| match role { + WrappedTransportRole::Source => proxy.source_entry_snapshots(transport), + WrappedTransportRole::Destination => proxy.destination_entry_snapshots(transport), + }) + } + + pub fn wrapped_transport_is_started( + &self, + transport: WrappedTransportKind, + role: WrappedTransportRole, + ) -> bool { + self.wrapped_transport + .as_ref() + .is_some_and(|proxy| match role { + WrappedTransportRole::Source => proxy.source_is_started(transport), + WrappedTransportRole::Destination => proxy.destination_is_started(transport), + }) + } +} diff --git a/easytier-core/src/instance/public_ipv6_extension.rs b/easytier-core/src/instance/public_ipv6_extension.rs new file mode 100644 index 00000000..59e8b5ec --- /dev/null +++ b/easytier-core/src/instance/public_ipv6_extension.rs @@ -0,0 +1,10 @@ +use crate::instance::{CoreInstance, CoreInstanceHost}; + +impl CoreInstance +where + H: CoreInstanceHost, +{ + pub async fn reconcile_public_ipv6_provider(&self) -> bool { + self.public_ipv6_provider.reconcile().await + } +} diff --git a/easytier-core/src/instance/test_utils.rs b/easytier-core/src/instance/test_utils.rs new file mode 100644 index 00000000..ddee05c4 --- /dev/null +++ b/easytier-core/src/instance/test_utils.rs @@ -0,0 +1,93 @@ +use std::{sync::Arc, time::Duration}; + +use crate::{ + connectivity::stun::StunSocketMapper, peers::conn::peer_conn::PeerConnId, + socket::udp::VirtualUdpSocketFactory, +}; + +use super::{ + CoreHostAdapters, CoreInstance, CoreInstanceConfig, CoreInstanceHost, PeerRelaySessionSnapshot, + build_capabilities, +}; + +impl CoreInstanceConfig { + #[doc(hidden)] + pub fn validate_build_capabilities_for_test(&self) -> anyhow::Result<()> { + build_capabilities::validate(self) + } +} + +impl CoreHostAdapters +where + H: CoreInstanceHost, +{ + #[doc(hidden)] + pub fn replace_stun_provider( + &mut self, + provider: Arc::Socket>>, + ) { + self.stun_override = Some(provider); + } +} + +impl CoreInstance +where + H: CoreInstanceHost, +{ + #[doc(hidden)] + pub async fn connected_peers(&self) -> Vec { + self.peer_manager + .get_peer_map() + .list_peers_with_conn() + .await + } + + #[doc(hidden)] + pub async fn admit_client_tunnel_for_test( + &self, + tunnel: Box, + is_directly_connected: bool, + ) -> Result<(crate::config::PeerId, PeerConnId), crate::peers::error::Error> { + self.peer_manager + .add_client_tunnel(tunnel, is_directly_connected) + .await + } + + #[doc(hidden)] + pub async fn relay_route_has_static_key_for_test( + &self, + peer_id: crate::config::PeerId, + ) -> bool { + self.peer_manager + .get_peer_map() + .get_route_peer_info(peer_id) + .await + .is_some_and(|info| !info.noise_static_pubkey.is_empty()) + } + + #[doc(hidden)] + pub fn relay_session_snapshot_for_test( + &self, + peer_id: crate::config::PeerId, + ) -> PeerRelaySessionSnapshot { + let relay = self.peer_manager.get_relay_peer_map(); + PeerRelaySessionSnapshot { + has_state: relay.has_state(peer_id), + has_session: relay.has_session_without_touch(peer_id), + } + } + + #[doc(hidden)] + pub fn evict_idle_relay_sessions_for_test(&self, idle: Duration) { + self.peer_manager + .get_relay_peer_map() + .evict_idle_sessions(idle); + } + + #[doc(hidden)] + pub fn evict_unused_peer_sessions_for_test(&self, idle: Duration) { + self.peer_manager + .get_peer_session_store() + .evict_unused_sessions_idle(idle); + } +} diff --git a/easytier-core/src/instance/tests.rs b/easytier-core/src/instance/tests.rs new file mode 100644 index 00000000..bc87ee19 --- /dev/null +++ b/easytier-core/src/instance/tests.rs @@ -0,0 +1,2364 @@ +use std::sync::Arc; + +use async_trait::async_trait; + +use super::*; +use crate::{ + config::toml::ConfigLoader as _, + listener::transport::TransportListenerConfig, + socket::{ + SocketContext, SocketListener, + udp::{UdpSessionAcceptKind, UdpSessionProtocol}, + }, +}; + +struct TestServerProtocol; + +#[async_trait] +impl ServerProtocolUpgrader<()> for TestServerProtocol { + fn supports_scheme(&self, scheme: &str) -> bool { + matches!(scheme, "ws" | "wss" | "wg" | "quic" | "faketcp" | "unix") + } + + async fn upgrade_tcp( + &self, + _socket: (), + _local_url: Url, + ) -> anyhow::Result { + unreachable!() + } + + async fn upgrade_udp( + &self, + _session: crate::socket::udp::UdpSession, + _local_url: Url, + _admission: Option, + ) -> anyhow::Result { + unreachable!() + } + + async fn upgrade_byte_stream( + &self, + _socket: (), + _local_url: Url, + _remote_url: Option, + ) -> anyhow::Result { + unreachable!() + } +} + +struct TestExternalListenerFactory; + +impl ExternalListenerFactory<()> for TestExternalListenerFactory { + fn supports_scheme(&self, scheme: &str) -> bool { + matches!(scheme, "faketcp" | "unix") + } + + fn create(&self, _request: ExternalListenerRequest) -> Box> { + unreachable!() + } +} + +#[test] +fn runtime_updates_retain_core_owned_peer_identity() { + let mut snapshot = Arc::new(PeerRuntimeSnapshot::default()); + Arc::make_mut(&mut snapshot).runtime.core.node.peer_id = Some(17); + Arc::make_mut(&mut snapshot).runtime.core.node.instance_id = Some([1; 16]); + let submitted = snapshot.clone(); + + retain_core_peer_identity(&mut snapshot, 23, Some([2; 16])); + + assert_eq!(snapshot.runtime.core.node.peer_id, Some(23)); + assert_eq!(snapshot.runtime.core.node.instance_id, Some([2; 16])); + assert_eq!(submitted.runtime.core.node.peer_id, Some(17)); + assert_eq!(submitted.runtime.core.node.instance_id, Some([1; 16])); +} + +#[test] +fn core_plans_transport_and_external_listener_capabilities() { + let self_id = uuid::Uuid::new_v4(); + let config = ListenerRuntimeConfig::new( + [ + "tcp://127.0.0.1:1", + "udp://127.0.0.1:2", + "ws://127.0.0.1:3", + "wg://127.0.0.1:4", + "quic://127.0.0.1:5", + "faketcp://127.0.0.1:6", + "unix:///tmp/easytier-test", + "http://127.0.0.1:7", + ] + .into_iter() + .map(str::parse) + .collect::, _>>() + .unwrap(), + false, + SocketContext::default().with_socket_mark(Some(7)), + ); + + let plan = prepare_listener_plan::<(), ()>( + Some(&config), + self_id, + Some(&TestServerProtocol), + Some(&TestExternalListenerFactory), + ) + .unwrap(); + + assert_eq!(plan.transports.len(), 6); + assert_eq!(plan.external.len(), 2); + assert_eq!(plan.failures.len(), 1); + assert_eq!( + plan.transports[0].url(), + &crate::listener::plan::ring_listener_url(self_id) + ); + assert!(matches!( + &plan.transports[4], + TransportListenerConfig::Udp { + accept_kind: UdpSessionAcceptKind::Classified(UdpSessionProtocol::WireGuard), + .. + } + )); + assert!(matches!( + &plan.transports[5], + TransportListenerConfig::Udp { + accept_kind: UdpSessionAcceptKind::Classified(UdpSessionProtocol::Quic), + .. + } + )); + assert_eq!(plan.external[0].0.url.scheme(), "faketcp"); + assert_eq!(plan.external[0].1.socket_mark, Some(7)); +} + +#[test] +fn unsupported_protocol_listener_becomes_a_plan_failure() { + let config = ListenerRuntimeConfig::new( + vec!["wg://127.0.0.1:11011".parse().unwrap()], + false, + SocketContext::default(), + ); + + let plan = + prepare_listener_plan::<(), ()>(Some(&config), uuid::Uuid::new_v4(), None, None).unwrap(); + + assert_eq!(plan.transports.len(), 1); + assert!(plan.external.is_empty()); + assert_eq!(plan.failures.len(), 1); +} + +#[test] +fn raw_unix_listener_does_not_require_a_server_protocol() { + let config = ListenerRuntimeConfig::new( + vec!["unix:///tmp/easytier-test".parse().unwrap()], + false, + SocketContext::default(), + ); + + let plan = prepare_listener_plan::<(), ()>( + Some(&config), + uuid::Uuid::new_v4(), + None, + Some(&TestExternalListenerFactory), + ) + .unwrap(); + + assert_eq!(plan.transports.len(), 1); + assert_eq!(plan.external.len(), 1); + assert!(plan.failures.is_empty()); + assert_eq!(plan.external[0].0.url.scheme(), "unix"); +} + +#[test] +fn core_instance_config_round_trips_as_normalized_json() { + let mut core = crate::config::CoreConfig::default(); + core.peer_policy.encryption_required = false; + core.peer_policy.p2p_enabled = false; + let peer = crate::peers::peer_manager::PortablePeerManagerConfig::new( + crate::config::peers::PeerRuntimeConfig { + core, + network_identity: crate::config::NetworkIdentity { + network_name: "default".to_owned(), + network_secret: Some("test".to_owned()), + network_secret_digest: None, + }, + stun_info: crate::proto::common::StunInfo::default(), + feature_flags: crate::proto::common::PeerFeatureFlag::default(), + secure_mode: None, + host_routing: crate::config::peers::HostRoutingPolicy::default(), + }, + ); + let config = CoreInstanceConfig { + instance_name: String::new(), + peer, + connectivity: CoreConnectivityConfig::default(), + }; + + let mut config = config; + config.connectivity.direct.testing = true; + let encoded = serde_json::to_value(&config).unwrap(); + assert!(encoded["connectivity"]["direct"].get("testing").is_none()); + let decoded: CoreInstanceConfig = serde_json::from_value(encoded.clone()).unwrap(); + + assert!(!decoded.connectivity.direct.testing); + assert!(decoded.connectivity.startup_plan.gateway); + assert_eq!(serde_json::to_value(&decoded).unwrap(), encoded); + + let mut legacy = encoded; + legacy.as_object_mut().unwrap().remove("instance_name"); + let decoded: CoreInstanceConfig = serde_json::from_value(legacy).unwrap(); + assert_eq!(decoded.instance_name, "default"); +} + +#[test] +fn wasi_create_config_uses_shared_toml() { + let fixture = include_bytes!("../../testdata/wasi_core_instance_create.json"); + let mut create = + serde_json::from_slice::(fixture) + .unwrap(); + create.validate().unwrap(); + let config = create.parse_config().unwrap(); + let normalized = CoreInstanceConfig::from_toml(&config).unwrap(); + + assert_eq!( + normalized.peer.snapshot.runtime.core.node.instance_id, + Some(*config.get_id().as_bytes()) + ); + assert!(normalized.peer.snapshot.flags.disable_p2p); + create.version += 1; + assert!(create.validate().is_err()); +} + +#[test] +fn core_instance_config_validation_rejects_invalid_acl_whitelist() { + let peer = crate::peers::peer_manager::PortablePeerManagerConfig::new( + crate::config::peers::PeerRuntimeConfig { + core: crate::config::CoreConfig::default(), + network_identity: crate::config::NetworkIdentity { + network_name: "default".to_owned(), + network_secret: Some("test".to_owned()), + network_secret_digest: None, + }, + stun_info: crate::proto::common::StunInfo::default(), + feature_flags: crate::proto::common::PeerFeatureFlag::default(), + secure_mode: None, + host_routing: crate::config::peers::HostRoutingPolicy::default(), + }, + ); + let mut config = CoreInstanceConfig { + instance_name: String::new(), + peer, + connectivity: CoreConnectivityConfig::default(), + }; + config.connectivity.runtime.acl.tcp_whitelist = vec!["9000-8000".to_owned()]; + + let error = validate_core_instance_config(&config).unwrap_err(); + + assert!(error.to_string().contains("Start port must be <= end port")); +} + +mod portable_runtime { + use std::{ + net::{IpAddr, Ipv6Addr, SocketAddr}, + sync::{ + Arc, + atomic::{AtomicBool, AtomicUsize, Ordering}, + }, + time::Duration, + }; + + use tokio::sync::Notify; + use tokio_util::task::AbortOnDropHandle; + + #[cfg(feature = "proxy-packet")] + use std::sync::Mutex as StdMutex; + + use super::*; + use crate::{ + config::peers::{HostRoutingPolicy, PeerRuntimeConfig}, + config::runtime::CoreInstanceRuntimeConfig, + config::{ + CoreConfig, IpPrefix, NetworkIdentity, ProxyNetworkConfig, gateway::PortForwardConfig, + }, + connectivity::manual::{ManualConnectorHost, ManualInterfaceAddrs}, + gateway::proxy::wrapped_transport::{ + WrappedTransportEngine, WrappedTransportEngineStart, WrappedTransportEngines, + WrappedTransportRole, + }, + host::testkit::{TestDns, TestHost, TestTcpSocket}, + listener::transport::AcceptedTransport, + peers::peer_manager::PortablePeerManagerConfig, + proto::{common::StunInfo, peer_rpc::GetIpListResponse}, + socket::{SocketContext, udp::PreferredIpv6Source}, + }; + + #[cfg(feature = "proxy-packet")] + use crate::gateway::proxy::wrapped_transport::WrappedTransportKind; + + #[async_trait] + impl ManualConnectorHost for TestHost { + async fn local_addr_for_remote( + &self, + remote_addr: SocketAddr, + _context: SocketContext, + ) -> anyhow::Result { + Ok(match remote_addr { + SocketAddr::V4(_) => "127.0.0.1:0".parse().unwrap(), + SocketAddr::V6(_) => "[::1]:0".parse().unwrap(), + }) + } + + async fn interface_addrs(&self) -> anyhow::Result { + Ok(ManualInterfaceAddrs { + interface_ipv4s: vec![], + interface_ipv6s: vec![], + public_ipv6: None, + }) + } + } + + #[async_trait] + impl DirectConnectorHost for TestHost { + async fn collect_ip_addrs(&self, _context: &SocketContext) -> GetIpListResponse { + GetIpListResponse::default() + } + + fn mapped_listeners(&self) -> Vec { + Vec::new() + } + + fn is_local_ip(&self, _ip: &IpAddr) -> bool { + false + } + + async fn preferred_ipv6_source( + &self, + _ip: Ipv6Addr, + _context: SocketContext, + ) -> Option { + None + } + } + + fn test_config(network_name: &str) -> CoreInstanceConfig { + let mut core = CoreConfig::default(); + core.node.network_name = network_name.to_owned(); + core.peer_policy.encryption_required = false; + let peer = PortablePeerManagerConfig::new(PeerRuntimeConfig { + core, + network_identity: NetworkIdentity { + network_name: network_name.to_owned(), + network_secret: Some(String::new()), + network_secret_digest: None, + }, + stun_info: StunInfo::default(), + feature_flags: Default::default(), + secure_mode: None, + host_routing: HostRoutingPolicy::default(), + }); + let connectivity = CoreConnectivityConfig::default(); + CoreInstanceConfig { + instance_name: network_name.to_owned(), + peer, + connectivity, + } + } + + fn runtime_snapshot(config: &CoreInstanceConfig) -> CoreInstanceRuntimeConfig { + CoreInstanceRuntimeConfig { + services: config.connectivity.runtime.clone(), + peer: Arc::new(config.peer.snapshot.clone()), + } + } + + fn proxy_network(real: &str, mapped: Option<&str>) -> ProxyNetworkConfig { + fn prefix(value: &str) -> IpPrefix { + let (address, prefix_len) = value.split_once('/').unwrap(); + IpPrefix { + address: address.parse().unwrap(), + prefix_len: prefix_len.parse().unwrap(), + } + } + + ProxyNetworkConfig { + real: prefix(real), + mapped: mapped.map(prefix), + } + } + + fn adapters( + external_listener_factory: Option< + Arc>>, + >, + packet_sink: Arc, + ) -> CoreHostAdapters { + adapters_with_process_runtime( + external_listener_factory, + packet_sink, + CoreProcessRuntime::new(), + ) + } + + fn adapters_with_process_runtime( + external_listener_factory: Option< + Arc>>, + >, + packet_sink: Arc, + process_runtime: Arc, + ) -> CoreHostAdapters { + adapters_with_host_and_process_runtime( + Arc::new(TestHost::default()), + external_listener_factory, + packet_sink, + process_runtime, + ) + } + + #[cfg(any(feature = "proxy-packet", feature = "proxy-smoltcp-stack"))] + fn adapters_with_host( + host: Arc, + external_listener_factory: Option< + Arc>>, + >, + packet_sink: Arc, + ) -> CoreHostAdapters { + adapters_with_host_and_process_runtime( + host, + external_listener_factory, + packet_sink, + CoreProcessRuntime::new(), + ) + } + + fn adapters_with_host_and_process_runtime( + host: Arc, + external_listener_factory: Option< + Arc>>, + >, + packet_sink: Arc, + process_runtime: Arc, + ) -> CoreHostAdapters { + let dns = Arc::new(TestDns); + let mut adapters = CoreHostAdapters::new(host, dns, packet_sink, process_runtime); + adapters.external_listener_factory = external_listener_factory; + adapters + } + + fn build_with_engines( + config: CoreInstanceConfig, + engines: WrappedTransportEngines, + ) -> anyhow::Result>> { + build_with_engines_and_listener(config, engines, None) + } + + fn build_with_engines_and_listener( + config: CoreInstanceConfig, + engines: WrappedTransportEngines, + external_listener_factory: Option< + Arc>>, + >, + ) -> anyhow::Result>> { + let (packet_sink, _packet_receiver) = tokio::sync::mpsc::channel(16); + let mut adapters = adapters(external_listener_factory, Arc::new(packet_sink)); + adapters.wrapped_transports = engines; + CoreInstance::new(config, adapters) + } + + fn build_instance(config: CoreInstanceConfig) -> anyhow::Result>> { + build_with_engines(config, WrappedTransportEngines::default()) + } + + #[tokio::test] + async fn core_instance_is_a_direct_managed_record() { + let instance = build_instance(test_config("managed-directly")).unwrap(); + + assert_eq!(instance.instance_name(), "managed-directly"); + assert_eq!( + instance.instance_id(), + crate::instance::manager::ManagedInstance::instance_id(instance.as_ref()) + ); + } + + #[tokio::test] + async fn instance_start_ignores_configured_peer_id() { + let mut config = test_config("fresh-peer-id"); + let instance_id = uuid::Uuid::from_bytes([7; 16]); + config.peer.snapshot.runtime.core.node.instance_id = Some(*instance_id.as_bytes()); + config.peer.snapshot.runtime.core.node.peer_id = Some(0); + + let instance = build_instance(config).unwrap(); + let generated_peer_id = instance.peer_id(); + + assert_eq!(instance.instance_id(), instance_id); + assert_ne!(generated_peer_id, 0); + assert_eq!( + instance + .runtime_config + .snapshot() + .peer + .runtime + .core + .node + .peer_id, + Some(generated_peer_id) + ); + } + + #[cfg(all(feature = "proxy-packet", feature = "management"))] + #[tokio::test] + async fn process_management_rpc_resolves_and_calls_core_instance_directly() { + use crate::{ + config::toml::TomlConfig, + instance::manager::{InstanceFactory, InstanceManager}, + management::InstanceManagementRpc, + }; + use easytier_proto::{ + api::config::{ConfigRpc, GetConfigRequest, InstanceConfigPatch, PatchConfigRequest}, + api::instance::{ + PeerManageRpc, ShowNodeInfoRequest, + instance_identifier::{InstanceSelector, Selector}, + }, + rpc_types::controller::BaseController, + }; + + struct ManagementTestFactory; + + impl InstanceFactory for ManagementTestFactory { + type Instance = CoreInstance; + type CreateContext = (); + type Error = anyhow::Error; + + fn create( + &self, + config: TomlConfig, + (): Self::CreateContext, + ) -> Result, Self::Error> { + let (packet_sink, _packet_receiver) = tokio::sync::mpsc::channel(16); + let mut adapters = adapters(None, Arc::new(packet_sink)); + adapters.config.force_exit_node = true; + adapters.config.public_ipv6_provider_supported = true; + adapters.config.easytier_version = "host-version".to_owned(); + CoreInstance::from_toml(config, adapters) + } + } + + let manager = Arc::new(InstanceManager::new(ManagementTestFactory, None)); + let config = TomlConfig::new_from_str( + r#" +instance_name = "managed-by-name" +hostname = "core-owned-config" +"#, + ) + .unwrap(); + let instance = manager.create(config, ()).unwrap(); + instance.start().await.unwrap(); + let rpc = InstanceManagementRpc::::new(manager); + + let response = rpc + .show_node_info( + BaseController::default(), + ShowNodeInfoRequest { + instance: Some(easytier_proto::api::instance::InstanceIdentifier { + selector: Some(Selector::InstanceSelector(InstanceSelector { + name: Some("managed-by-name".to_owned()), + })), + }), + }, + ) + .await + .unwrap(); + + let node = response.node_info.unwrap(); + assert_eq!(node.hostname, "core-owned-config"); + assert!(node.config.contains("instance_name = \"managed-by-name\"")); + + let selector = || easytier_proto::api::instance::InstanceIdentifier { + selector: Some(Selector::InstanceSelector(InstanceSelector { + name: Some("managed-by-name".to_owned()), + })), + }; + rpc.patch_config( + BaseController::default(), + PatchConfigRequest { + patch: Some(InstanceConfigPatch { + hostname: Some("patched-in-core".to_owned()), + ..Default::default() + }), + instance: Some(selector()), + }, + ) + .await + .unwrap(); + let response = rpc + .get_config( + BaseController::default(), + GetConfigRequest { + instance: Some(selector()), + }, + ) + .await + .unwrap(); + assert_eq!( + response.config.unwrap().hostname.as_deref(), + Some("patched-in-core") + ); + let runtime = instance.runtime_config_snapshot(); + assert!(runtime.services.proxy.enable_exit_node); + assert!(runtime.services.public_ipv6_provider.provider_supported); + assert_eq!(runtime.peer.easytier_version, "host-version"); + } + + #[cfg(feature = "management")] + #[tokio::test] + async fn config_patch_rejects_instance_before_host_startup_is_complete() { + let instance = build_instance(test_config("not-ready-for-patch")).unwrap(); + + let error = crate::management::apply_config_patch( + &instance, + easytier_proto::api::config::InstanceConfigPatch { + hostname: Some("too-early".to_owned()), + ..Default::default() + }, + ) + .await + .unwrap_err(); + + assert!(error.to_string().contains("instance is not ready")); + } + + #[cfg(all(feature = "management", not(feature = "proxy-smoltcp-stack")))] + #[tokio::test] + async fn unavailable_gateway_patch_does_not_commit_shared_toml() { + use easytier_proto::{ + api::config::{ConfigPatchAction, InstanceConfigPatch, PortForwardPatch}, + common::{PortForwardConfigPb, SocketType}, + }; + + let (packet_sink, _packet_receiver) = tokio::sync::mpsc::channel(16); + let instance = CoreInstance::from_toml( + crate::config::toml::TomlConfig::new_from_str( + "instance_name = \"rejected-gateway-patch\"", + ) + .unwrap(), + adapters(None, Arc::new(packet_sink)), + ) + .unwrap(); + instance.start().await.unwrap(); + + let error = crate::management::apply_config_patch( + &instance, + InstanceConfigPatch { + port_forwards: vec![PortForwardPatch { + action: ConfigPatchAction::Add as i32, + cfg: Some(PortForwardConfigPb { + bind_addr: Some( + "127.0.0.1:18080" + .parse::() + .unwrap() + .into(), + ), + dst_addr: Some( + "10.144.144.2:8080" + .parse::() + .unwrap() + .into(), + ), + socket_type: SocketType::Tcp as i32, + }), + }], + ..Default::default() + }, + ) + .await + .unwrap_err(); + + assert!( + error + .to_string() + .contains("does not include the smoltcp gateway"), + "unexpected patch error: {error:#}" + ); + assert!( + instance + .toml_config() + .unwrap() + .get_port_forwards() + .is_empty() + ); + assert!( + instance + .runtime_config_snapshot() + .services + .gateway + .port_forwards + .is_empty() + ); + instance.stop().await; + } + + #[tokio::test] + async fn dropping_core_instance_requests_host_shutdown() { + struct DropAwareRuntimeHost(Arc); + + #[async_trait] + impl InstanceRuntimeHost for DropAwareRuntimeHost { + async fn prepare( + &self, + _packet_plane: Arc, + ) -> anyhow::Result>> { + Ok(None) + } + + async fn shutdown(&self) {} + + fn request_shutdown(&self) { + self.0.store(true, Ordering::Release); + } + } + + let shutdown_requested = Arc::new(AtomicBool::new(false)); + let (packet_sink, _packet_receiver) = tokio::sync::mpsc::channel(16); + let mut adapters = adapters(None, Arc::new(packet_sink)); + adapters.instance_runtime = Arc::new(DropAwareRuntimeHost(shutdown_requested.clone())); + let instance = CoreInstance::new(test_config("drop-cleanup"), adapters).unwrap(); + + drop(instance); + + assert!(shutdown_requested.load(Ordering::Acquire)); + } + + #[tokio::test] + async fn host_prepare_failure_runs_unified_cleanup() { + #[derive(Default)] + struct FailingRuntimeHost { + shutdown_calls: AtomicUsize, + } + + #[async_trait] + impl InstanceRuntimeHost for FailingRuntimeHost { + async fn prepare( + &self, + _packet_plane: Arc, + ) -> anyhow::Result>> { + anyhow::bail!("host prepare failed") + } + + async fn shutdown(&self) { + self.shutdown_calls.fetch_add(1, Ordering::Relaxed); + } + } + + let runtime_host = Arc::new(FailingRuntimeHost::default()); + let (packet_sink, _packet_receiver) = tokio::sync::mpsc::channel(16); + let mut adapters = adapters(None, Arc::new(packet_sink)); + adapters.instance_runtime = runtime_host.clone(); + let instance = CoreInstance::new(test_config("host-prepare-failure"), adapters).unwrap(); + + let error = instance.start().await.unwrap_err(); + + assert!(error.to_string().contains("host prepare failed")); + assert_eq!(instance.state(), CoreInstanceState::Stopped); + assert!(!instance.is_ready()); + assert_eq!(runtime_host.shutdown_calls.load(Ordering::Relaxed), 1); + assert!( + instance + .latest_error() + .unwrap() + .contains("host prepare failed") + ); + } + + #[tokio::test] + async fn aborting_host_prepare_runs_unified_cleanup() { + #[derive(Default)] + struct BlockingPrepareRuntimeHost { + prepare_started: Notify, + prepare_release: Notify, + shutdown_calls: AtomicUsize, + } + + #[async_trait] + impl InstanceRuntimeHost for BlockingPrepareRuntimeHost { + async fn prepare( + &self, + _packet_plane: Arc, + ) -> anyhow::Result>> { + self.prepare_started.notify_one(); + self.prepare_release.notified().await; + Ok(None) + } + + async fn shutdown(&self) { + self.shutdown_calls.fetch_add(1, Ordering::Relaxed); + } + } + + let runtime_host = Arc::new(BlockingPrepareRuntimeHost::default()); + let (packet_sink, _packet_receiver) = tokio::sync::mpsc::channel(16); + let mut adapters = adapters(None, Arc::new(packet_sink)); + adapters.instance_runtime = runtime_host.clone(); + let instance = CoreInstance::new(test_config("aborted-host-prepare"), adapters).unwrap(); + let start = tokio::spawn({ + let instance = instance.clone(); + async move { instance.start().await } + }); + runtime_host.prepare_started.notified().await; + + start.abort(); + assert!(start.await.unwrap_err().is_cancelled()); + tokio::time::timeout(Duration::from_secs(2), async { + while instance.state() != CoreInstanceState::Stopped { + tokio::task::yield_now().await; + } + }) + .await + .expect("aborted Host prepare should stop the instance"); + + assert!(!instance.is_ready()); + assert_eq!(runtime_host.shutdown_calls.load(Ordering::Relaxed), 1); + } + + #[cfg(feature = "management")] + #[tokio::test] + async fn cancelled_delete_finishes_stop_and_keeps_wait_blocked() { + use std::{path::Path, path::PathBuf}; + + use crate::{ + config::toml::TomlConfig, + instance::manager::InstanceFactory, + management::{ + ConfigFileControl, ConfigFilePermission, ConfigFileStorage, InstanceManager, + InstanceMutationHooks, ProcessManagementRpc, + }, + }; + use easytier_proto::{ + api::manage::{DeleteNetworkInstanceRequest, WebClientService}, + rpc_types::controller::BaseController, + }; + + #[derive(Default)] + struct BlockingRuntimeHost { + shutdown_started: Notify, + shutdown_release: Notify, + } + + #[async_trait] + impl InstanceRuntimeHost for BlockingRuntimeHost { + async fn prepare( + &self, + _packet_plane: Arc, + ) -> anyhow::Result>> { + Ok(None) + } + + async fn shutdown(&self) { + self.shutdown_started.notify_one(); + self.shutdown_release.notified().await; + } + } + + struct BlockingFactory { + runtime_host: Arc, + process_runtime: Arc, + } + + impl InstanceFactory for BlockingFactory { + type Instance = CoreInstance; + type CreateContext = (); + type Error = anyhow::Error; + + fn create( + &self, + config: TomlConfig, + (): Self::CreateContext, + ) -> Result, Self::Error> { + let (packet_sink, _packet_receiver) = tokio::sync::mpsc::channel(16); + let mut adapters = adapters_with_process_runtime( + None, + Arc::new(packet_sink), + self.process_runtime.clone(), + ); + adapters.instance_runtime = self.runtime_host.clone(); + CoreInstance::from_toml(config, adapters) + } + } + + impl crate::management::ProcessRuntimeProvider for BlockingFactory { + fn process_runtime(&self) -> Arc { + self.process_runtime.clone() + } + } + + #[derive(Default)] + struct RecordingStorage(AtomicBool); + + #[async_trait] + impl ConfigFileStorage for RecordingStorage { + async fn inspect(&self, path: &Path) -> ConfigFileControl { + ConfigFileControl::new(Some(path.to_owned()), ConfigFilePermission::default()) + } + + async fn read(&self, _path: &Path) -> anyhow::Result>> { + Ok(None) + } + + async fn write(&self, _path: &Path, _contents: &[u8]) -> anyhow::Result<()> { + Ok(()) + } + + async fn remove(&self, _path: &Path) -> anyhow::Result<()> { + self.0.store(true, Ordering::Release); + Ok(()) + } + } + + #[derive(Default)] + struct RecordingHooks(AtomicBool); + + #[async_trait] + impl InstanceMutationHooks for RecordingHooks { + async fn post_remove_network_instances( + &self, + _instance_ids: &[uuid::Uuid], + ) -> Result<(), String> { + self.0.store(true, Ordering::Release); + Ok(()) + } + } + + let runtime_host = Arc::new(BlockingRuntimeHost::default()); + let process_runtime = CoreProcessRuntime::new(); + let instances = Arc::new(InstanceManager::new( + BlockingFactory { + runtime_host: runtime_host.clone(), + process_runtime, + }, + Some(tokio::runtime::Handle::current()), + )); + let config = TomlConfig::default(); + config.set_listeners(Vec::new()); + let instance_id = config.get_id(); + instances + .run_network_instance( + config, + ConfigFileControl::new( + Some(PathBuf::from("cancelled-delete.toml")), + ConfigFilePermission::default(), + ), + ) + .unwrap(); + let storage = Arc::new(RecordingStorage::default()); + let hooks = Arc::new(RecordingHooks::default()); + let rpc = ProcessManagementRpc::::new( + instances.clone(), + hooks.clone(), + storage.clone(), + ); + + let deletion = tokio::spawn(async move { + rpc.delete_network_instance( + BaseController::default(), + DeleteNetworkInstanceRequest { + inst_ids: vec![instance_id.into()], + }, + ) + .await + }); + runtime_host.shutdown_started.notified().await; + deletion.abort(); + assert!(deletion.await.unwrap_err().is_cancelled()); + + let mut wait = tokio::spawn({ + let instances = instances.clone(); + async move { instances.wait().await } + }); + assert!( + tokio::time::timeout(Duration::from_millis(20), &mut wait) + .await + .is_err() + ); + + runtime_host.shutdown_release.notify_one(); + tokio::time::timeout(Duration::from_secs(1), async { + while !storage.0.load(Ordering::Acquire) { + tokio::task::yield_now().await; + } + }) + .await + .unwrap(); + tokio::time::timeout(Duration::from_secs(1), wait) + .await + .unwrap() + .unwrap(); + assert!(hooks.0.load(Ordering::Acquire)); + assert!(instances.instances().is_empty()); + } + + #[cfg(feature = "management")] + #[tokio::test] + async fn process_management_rpc_owns_instance_create_list_and_delete() { + use crate::{ + config::toml::TomlConfig, + instance::manager::InstanceFactory, + management::{InstanceManager, ProcessManagementRpc, UnsupportedConfigFileStorage}, + }; + use easytier_proto::{ + api::manage::{ + DeleteNetworkInstanceRequest, ListNetworkInstanceRequest, NetworkConfig, + NetworkingMethod, RunNetworkInstanceRequest, WebClientService, + }, + rpc_types::controller::BaseController, + }; + + struct ManagementTestFactory(Arc); + + impl InstanceFactory for ManagementTestFactory { + type Instance = CoreInstance; + type CreateContext = (); + type Error = anyhow::Error; + + fn create( + &self, + config: TomlConfig, + (): Self::CreateContext, + ) -> Result, Self::Error> { + let (packet_sink, _packet_receiver) = tokio::sync::mpsc::channel(16); + CoreInstance::from_toml( + config, + adapters_with_process_runtime(None, Arc::new(packet_sink), self.0.clone()), + ) + } + } + + impl crate::management::ProcessRuntimeProvider for ManagementTestFactory { + fn process_runtime(&self) -> Arc { + self.0.clone() + } + } + + let instances = Arc::new(InstanceManager::new( + ManagementTestFactory(CoreProcessRuntime::new()), + Some(tokio::runtime::Handle::current()), + )); + let rpc = ProcessManagementRpc::::new( + instances.clone(), + Arc::new(()), + Arc::new(UnsupportedConfigFileStorage), + ); + let created = rpc + .run_network_instance( + BaseController::default(), + RunNetworkInstanceRequest { + config: Some(NetworkConfig { + network_name: Some("managed-process-rpc".to_owned()), + networking_method: Some(NetworkingMethod::Standalone.into()), + ..Default::default() + }), + overwrite: true, + ..Default::default() + }, + ) + .await + .unwrap() + .inst_id + .unwrap(); + + let listed = rpc + .list_network_instance(BaseController::default(), ListNetworkInstanceRequest {}) + .await + .unwrap(); + assert_eq!(listed.inst_ids, vec![created]); + + let deleted = rpc + .delete_network_instance( + BaseController::default(), + DeleteNetworkInstanceRequest { + inst_ids: vec![created], + }, + ) + .await + .unwrap(); + assert!(deleted.remain_inst_ids.is_empty()); + assert!(instances.instances().is_empty()); + } + + #[cfg(feature = "management")] + #[tokio::test] + async fn owned_selection_and_cleanup_share_the_canonical_transaction() { + use crate::{ + config::toml::TomlConfig, + instance::manager::InstanceFactory, + management::{ + ConfigFileControl, InstanceManager, InstanceMutationHooks, ProcessManagement, + UnsupportedConfigFileStorage, + }, + }; + + struct ManagementTestFactory(Arc); + + impl InstanceFactory for ManagementTestFactory { + type Instance = CoreInstance; + type CreateContext = (); + type Error = anyhow::Error; + + fn create( + &self, + config: TomlConfig, + (): Self::CreateContext, + ) -> Result, Self::Error> { + let (packet_sink, _packet_receiver) = tokio::sync::mpsc::channel(16); + CoreInstance::from_toml( + config, + adapters_with_process_runtime(None, Arc::new(packet_sink), self.0.clone()), + ) + } + } + + impl crate::management::ProcessRuntimeProvider for ManagementTestFactory { + fn process_runtime(&self) -> Arc { + self.0.clone() + } + } + + #[derive(Default)] + struct BlockingRemovalHook { + entered: Notify, + release: Notify, + removed: std::sync::Mutex>>, + } + + #[async_trait] + impl InstanceMutationHooks for BlockingRemovalHook { + async fn post_remove_network_instances( + &self, + instance_ids: &[uuid::Uuid], + ) -> Result<(), String> { + self.removed.lock().unwrap().push(instance_ids.to_vec()); + self.entered.notify_one(); + self.release.notified().await; + Ok(()) + } + } + + let instances = Arc::new(InstanceManager::new( + ManagementTestFactory(CoreProcessRuntime::new()), + Some(tokio::runtime::Handle::current()), + )); + let config = TomlConfig::default(); + config.set_listeners(Vec::new()); + let instance_id = config.get_id(); + instances + .run_network_instance(config, ConfigFileControl::STATIC_CONFIG) + .unwrap(); + let hooks = Arc::new(BlockingRemovalHook::default()); + let management = ProcessManagement::::new( + instances.clone(), + hooks.clone(), + Arc::new(UnsupportedConfigFileStorage), + ); + + let deletion_management = management.clone(); + let deletion = tokio::spawn(async move { + deletion_management + .delete_owned_network_instances(vec![instance_id, uuid::Uuid::new_v4()]) + .await + }); + hooks.entered.notified().await; + assert!(instances.mutation_lock().try_lock().is_err()); + assert_eq!( + hooks.removed.lock().unwrap().as_slice(), + &[vec![instance_id]] + ); + hooks.release.notify_one(); + + let result = deletion.await.unwrap().unwrap(); + assert_eq!(result.removed_instance_ids, vec![instance_id]); + assert_eq!(hooks.removed.lock().unwrap().len(), 1); + + let old_name = format!("old-{instance_id}"); + let new_name = format!("new-{instance_id}"); + let old_config = TomlConfig::default(); + old_config.set_id(instance_id); + old_config.set_inst_name(old_name.clone()); + old_config.set_listeners(Vec::new()); + instances + .run_network_instance(old_config, ConfigFileControl::STATIC_CONFIG) + .unwrap(); + + let mutation_guard = instances.mutation_lock().lock_owned().await; + let deletion = management.delete_owned_network_instances_by_name(vec![old_name]); + tokio::pin!(deletion); + assert!(matches!( + futures::poll!(deletion.as_mut()), + std::task::Poll::Pending + )); + + instances + .delete_network_instances([instance_id]) + .await + .unwrap(); + let new_config = TomlConfig::default(); + new_config.set_id(instance_id); + new_config.set_inst_name(new_name.clone()); + new_config.set_listeners(Vec::new()); + instances + .run_network_instance(new_config, ConfigFileControl::STATIC_CONFIG) + .unwrap(); + + drop(mutation_guard); + hooks.release.notify_one(); + let result = deletion.await.unwrap(); + assert!(result.removed_instance_ids.is_empty()); + assert_eq!( + instances + .instance(instance_id) + .map(|instance| instance.instance_name().to_owned()) + .as_deref(), + Some(new_name.as_str()) + ); + assert_eq!( + hooks.removed.lock().unwrap().as_slice(), + &[vec![instance_id], Vec::new()] + ); + instances + .delete_network_instances([instance_id]) + .await + .unwrap(); + + let selection_ran = Arc::new(std::sync::atomic::AtomicBool::new(false)); + let selection_ran_for_call = selection_ran.clone(); + let mutation_guard = instances.mutation_lock().lock_owned().await; + let deletion = management.delete_owned_network_instances_selected_by(move || { + selection_ran_for_call.store(true, std::sync::atomic::Ordering::Release); + Vec::new() + }); + tokio::pin!(deletion); + assert!(matches!( + futures::poll!(deletion.as_mut()), + std::task::Poll::Pending + )); + assert!(!selection_ran.load(std::sync::atomic::Ordering::Acquire)); + drop(mutation_guard); + hooks.release.notify_one(); + deletion.await.unwrap(); + assert!(selection_ran.load(std::sync::atomic::Ordering::Acquire)); + } + + #[cfg(feature = "management")] + #[tokio::test] + async fn process_management_rpc_rolls_back_instance_and_file_on_hook_failure() { + use std::{ + collections::HashMap, + path::{Path, PathBuf}, + sync::Mutex as StdMutex, + }; + + use crate::{ + config::toml::TomlConfig, + instance::manager::InstanceFactory, + management::{ + ConfigFileControl, ConfigFilePermission, ConfigFileStorage, InstanceManager, + InstanceMutationHooks, ProcessManagementRpc, + }, + }; + use easytier_proto::{ + api::manage::{ + NetworkConfig, NetworkingMethod, RunNetworkInstanceRequest, WebClientService, + }, + rpc_types::controller::BaseController, + }; + + struct ManagementTestFactory(Arc); + + impl InstanceFactory for ManagementTestFactory { + type Instance = CoreInstance; + type CreateContext = (); + type Error = anyhow::Error; + + fn create( + &self, + config: TomlConfig, + (): Self::CreateContext, + ) -> Result, Self::Error> { + let (packet_sink, _packet_receiver) = tokio::sync::mpsc::channel(16); + CoreInstance::from_toml( + config, + adapters_with_process_runtime(None, Arc::new(packet_sink), self.0.clone()), + ) + } + } + + impl crate::management::ProcessRuntimeProvider for ManagementTestFactory { + fn process_runtime(&self) -> Arc { + self.0.clone() + } + } + + #[derive(Default)] + struct MemoryStorage(StdMutex>>); + + #[async_trait] + impl ConfigFileStorage for MemoryStorage { + async fn inspect(&self, path: &Path) -> ConfigFileControl { + ConfigFileControl::new(Some(path.to_owned()), ConfigFilePermission::default()) + } + + async fn read(&self, path: &Path) -> anyhow::Result>> { + Ok(self.0.lock().unwrap().get(path).cloned()) + } + + async fn write(&self, path: &Path, contents: &[u8]) -> anyhow::Result<()> { + self.0 + .lock() + .unwrap() + .insert(path.to_owned(), contents.to_vec()); + Ok(()) + } + + async fn remove(&self, path: &Path) -> anyhow::Result<()> { + self.0.lock().unwrap().remove(path); + Ok(()) + } + } + + struct RejectPostRun; + + #[async_trait] + impl InstanceMutationHooks for RejectPostRun { + fn manages_remote_config_instances(&self) -> bool { + true + } + + async fn post_run_network_instance( + &self, + _instance_id: &uuid::Uuid, + ) -> Result<(), String> { + Err("rejected for rollback test".to_owned()) + } + } + + let instances = Arc::new(InstanceManager::new( + ManagementTestFactory(CoreProcessRuntime::new()), + Some(tokio::runtime::Handle::current()), + )); + let storage = Arc::new(MemoryStorage::default()); + let instance_id = uuid::Uuid::new_v4(); + let config_path = PathBuf::from("managed-process-rpc.toml"); + let original_file = b"original configuration".to_vec(); + storage + .0 + .lock() + .unwrap() + .insert(config_path.clone(), original_file.clone()); + let original = TomlConfig::default(); + original.set_id(instance_id); + original.set_inst_name("original-instance".to_owned()); + original.set_listeners(Vec::new()); + instances + .run_network_instance( + original, + ConfigFileControl::new(Some(config_path.clone()), ConfigFilePermission::default()), + ) + .unwrap(); + let rpc = ProcessManagementRpc::::new( + instances.clone(), + Arc::new(RejectPostRun), + storage.clone(), + ); + + let result = rpc + .run_network_instance( + BaseController::default(), + RunNetworkInstanceRequest { + inst_id: Some(instance_id.into()), + config: Some(NetworkConfig { + network_name: Some("replacement".to_owned()), + networking_method: Some(NetworkingMethod::Standalone.into()), + ..Default::default() + }), + overwrite: true, + ..Default::default() + }, + ) + .await; + + assert!(result.is_err()); + assert_eq!( + instances + .instance(instance_id) + .map(|instance| instance.instance_name().to_owned()) + .as_deref(), + Some("original-instance") + ); + assert_eq!( + storage.0.lock().unwrap().get(&config_path), + Some(&original_file) + ); + instances + .delete_network_instances([instance_id]) + .await + .unwrap(); + } + + #[tokio::test] + async fn packet_plane_does_not_retain_core_instance() { + let instance = build_instance(test_config("packet-plane-ownership")).unwrap(); + let weak = Arc::downgrade(&instance); + let packet_plane = instance.packet_plane(); + + drop(instance); + + assert!(weak.upgrade().is_none()); + drop(packet_plane); + } + + #[derive(Default)] + struct RecordingProxyService { + start_calls: AtomicUsize, + stop_calls: AtomicUsize, + start_gate: Option>, + #[cfg(feature = "proxy-packet")] + destination_ingress: StdMutex< + Option, + >, + } + + #[derive(Default)] + struct ProxyStartGate { + entered: Notify, + release: Notify, + } + + impl RecordingProxyService { + fn blocking() -> (Arc, Arc) { + let gate = Arc::new(ProxyStartGate::default()); + ( + Arc::new(Self { + start_gate: Some(gate.clone()), + ..Default::default() + }), + gate, + ) + } + + #[cfg(feature = "proxy-packet")] + fn destination_ingress( + &self, + ) -> Option + { + self.destination_ingress.lock().unwrap().clone() + } + } + + #[async_trait] + impl WrappedTransportEngine for RecordingProxyService { + async fn prepare(&self, options: WrappedTransportEngineStart) -> anyhow::Result<()> { + self.start_calls.fetch_add(1, Ordering::Relaxed); + #[cfg(feature = "proxy-packet")] + { + *self.destination_ingress.lock().unwrap() = options.destination_ingress; + } + #[cfg(not(feature = "proxy-packet"))] + let _ = options; + if let Some(gate) = &self.start_gate { + gate.entered.notify_one(); + gate.release.notified().await; + } + Ok(()) + } + + async fn activate(&self) -> anyhow::Result<()> { + Ok(()) + } + + async fn inject_peer_datagram( + &self, + _role: WrappedTransportRole, + _from_peer_id: u32, + _payload: bytes::Bytes, + ) -> anyhow::Result<()> { + Ok(()) + } + + #[cfg(feature = "proxy-packet")] + async fn connect_source( + &self, + _request: crate::gateway::proxy::wrapped_transport::WrappedTransportConnect, + ) -> anyhow::Result> { + anyhow::bail!("recording engine does not open streams") + } + + async fn stop(&self) { + self.stop_calls.fetch_add(1, Ordering::Relaxed); + } + } + + #[derive(Debug, Default)] + struct BlockingListenerState { + start_entered: Notify, + drop_calls: AtomicUsize, + } + + #[derive(Debug)] + struct BlockingSocketListener { + url: Url, + state: Arc, + } + + #[async_trait] + impl SocketListener for BlockingSocketListener { + type Accepted = AcceptedTransport; + + async fn listen(&mut self) -> anyhow::Result<()> { + self.state.start_entered.notify_one(); + std::future::pending().await + } + + async fn accept(&mut self) -> anyhow::Result { + std::future::pending().await + } + + fn local_url(&self) -> Url { + self.url.clone() + } + } + + impl Drop for BlockingSocketListener { + fn drop(&mut self) { + self.state.drop_calls.fetch_add(1, Ordering::Relaxed); + } + } + + struct BlockingExternalListenerFactory { + state: Arc, + } + + impl ExternalListenerFactory> for BlockingExternalListenerFactory { + fn supports_scheme(&self, scheme: &str) -> bool { + scheme == "unix" + } + + fn create( + &self, + request: ExternalListenerRequest, + ) -> Box>> { + Box::new(BlockingSocketListener { + url: request.url, + state: self.state.clone(), + }) + } + } + + #[derive(Debug)] + struct ReadySocketListener(Url); + + #[async_trait] + impl SocketListener for ReadySocketListener { + type Accepted = AcceptedTransport; + + async fn listen(&mut self) -> anyhow::Result<()> { + Ok(()) + } + + async fn accept(&mut self) -> anyhow::Result { + std::future::pending().await + } + + fn local_url(&self) -> Url { + self.0.clone() + } + } + + struct ReadyExternalListenerFactory; + + impl ExternalListenerFactory> for ReadyExternalListenerFactory { + fn supports_scheme(&self, scheme: &str) -> bool { + scheme == "unix" + } + + fn create( + &self, + request: ExternalListenerRequest, + ) -> Box>> { + Box::new(ReadySocketListener(request.url)) + } + } + + #[tokio::test] + async fn runtime_updates_refresh_avoid_relay_preference() { + let config = test_config("portable-runtime-update"); + let instance = build_instance(config.clone()).unwrap(); + + assert!( + !instance + .node_snapshot() + .await + .feature_flags + .avoid_relay_data + ); + + let mut enabled = runtime_snapshot(&config); + Arc::make_mut(&mut enabled.peer).avoid_relay_data_preference = true; + instance.update_runtime_config(enabled).await.unwrap(); + assert!( + instance + .node_snapshot() + .await + .feature_flags + .avoid_relay_data + ); + + let mut disabled = runtime_snapshot(&config); + Arc::make_mut(&mut disabled.peer).avoid_relay_data_preference = false; + instance.update_runtime_config(disabled).await.unwrap(); + assert!( + !instance + .node_snapshot() + .await + .feature_flags + .avoid_relay_data + ); + } + + #[cfg(all(feature = "test-utils", feature = "dhcp-ipv4"))] + #[cfg_attr( + not(target_os = "wasi"), + tokio::test(flavor = "multi_thread", worker_threads = 2) + )] + #[cfg_attr(target_os = "wasi", tokio::test)] + async fn concurrent_runtime_updates_keep_snapshot_and_derived_state_coherent() { + let config = test_config("concurrent-runtime-update"); + let instance = build_instance(config.clone()).unwrap(); + instance.start().await.unwrap(); + + let original = instance.node_snapshot().await; + let mut full = runtime_snapshot(&config); + full.services.dhcp_ipv4 = true; + full.services.acl.tcp_whitelist = vec!["80".to_owned()]; + { + let peer = Arc::make_mut(&mut full.peer); + peer.runtime.core.node.hostname = Some("full".to_owned()); + peer.runtime.core.routes.proxy_networks = + vec![proxy_network("192.0.2.0/24", Some("198.51.100.0/24"))]; + } + let mut peer_update = full.clone(); + { + let peer = Arc::make_mut(&mut peer_update.peer); + peer.runtime.core.node.hostname = Some("peer".to_owned()); + peer.runtime.core.routes.proxy_networks = + vec![proxy_network("203.0.113.0/24", Some("10.20.30.0/24"))]; + } + + let start = Arc::new(tokio::sync::Barrier::new(3)); + let full_update = tokio::spawn({ + let instance = instance.clone(); + let start = start.clone(); + async move { + start.wait().await; + instance.update_runtime_config(full).await + } + }); + let peer_update = tokio::spawn({ + let instance = instance.clone(); + let start = start.clone(); + async move { + start.wait().await; + instance.update_runtime_config(peer_update).await + } + }); + start.wait().await; + full_update.await.unwrap().unwrap(); + peer_update.await.unwrap().unwrap(); + + let final_config = instance.runtime_config.snapshot(); + assert!(final_config.services.dhcp_ipv4); + assert_eq!(instance.acl_whitelist_snapshot().tcp_ports, ["80"]); + assert_eq!(instance.acl_reload_count.load(Ordering::Relaxed), 1); + let node = instance.node_snapshot().await; + assert_eq!(node.peer_id, original.peer_id); + assert_eq!(node.instance_id, original.instance_id); + assert_eq!( + final_config.peer.runtime.core.node.hostname.as_deref(), + Some(node.hostname.as_str()) + ); + match node.hostname.as_str() { + "full" => assert_eq!( + instance + .proxy_cidr_table + .lookup_v4("198.51.100.42".parse().unwrap()), + Some("192.0.2.42".parse().unwrap()) + ), + "peer" => assert_eq!( + instance + .proxy_cidr_table + .lookup_v4("10.20.30.42".parse().unwrap()), + Some("203.0.113.42".parse().unwrap()) + ), + hostname => panic!("unexpected final hostname: {hostname:?}"), + } + + instance.stop().await; + } + + #[cfg(all(feature = "test-utils", feature = "dhcp-ipv4"))] + #[tokio::test] + async fn active_runtime_update_skips_unchanged_and_rejects_invalid_acl() { + let config = test_config("invalid-active-acl-update"); + let instance = build_instance(config.clone()).unwrap(); + instance.start().await.unwrap(); + + let mut unrelated = runtime_snapshot(&config); + Arc::make_mut(&mut unrelated.peer) + .runtime + .core + .node + .hostname = Some("accepted".to_owned()); + instance.update_runtime_config(unrelated).await.unwrap(); + assert_eq!(instance.acl_reload_count.load(Ordering::Relaxed), 0); + let before = instance.node_snapshot().await; + + let mut rejected = runtime_snapshot(&config); + rejected.services.dhcp_ipv4 = true; + rejected.services.acl.tcp_whitelist = vec!["invalid".to_owned()]; + Arc::make_mut(&mut rejected.peer).runtime.core.node.hostname = Some("rejected".to_owned()); + + let error = instance.update_runtime_config(rejected).await.unwrap_err(); + assert!(error.to_string().contains("Invalid port number")); + assert!(!instance.runtime_config.snapshot().services.dhcp_ipv4); + assert!(instance.acl_whitelist_snapshot().tcp_ports.is_empty()); + assert_eq!(instance.acl_reload_count.load(Ordering::Relaxed), 0); + assert_eq!(instance.node_snapshot().await.hostname, before.hostname); + instance.stop().await; + } + + #[cfg(feature = "proxy-packet")] + #[tokio::test] + async fn runtime_core_instance_owns_connectivity_lifecycle() { + let mut config = test_config("connectivity-lifecycle"); + config.peer.snapshot.runtime.core.routes.proxy_networks = + vec![proxy_network("10.1.2.0/24", None)]; + let initial_peer: Url = "tcp://127.0.0.1:29999".parse().unwrap(); + config.connectivity.initial_peers = vec![initial_peer.clone()]; + let proxy = Arc::new(RecordingProxyService::default()); + let instance = build_with_engines( + config, + WrappedTransportEngines { + kcp: Some(proxy.clone()), + quic: None, + }, + ) + .unwrap(); + + assert_eq!(instance.state(), CoreInstanceState::Created); + assert_eq!(instance.list_connectors().len(), 1); + assert_eq!(instance.list_connectors()[0].url, initial_peer); + instance.start().await.unwrap(); + assert_eq!(instance.state(), CoreInstanceState::Running); + assert!(instance.is_ready()); + assert!(instance.start().await.is_err()); + assert_eq!(proxy.start_calls.load(Ordering::Relaxed), 1); + + instance.stop().await; + instance.stop().await; + assert_eq!(instance.state(), CoreInstanceState::Stopped); + assert_eq!(proxy.stop_calls.load(Ordering::Relaxed), 1); + } + + #[cfg(feature = "proxy-smoltcp-stack")] + #[tokio::test] + async fn startup_plan_controls_gateway_for_initial_and_updated_config() { + fn build(config: CoreInstanceConfig) -> (Arc>, Arc) { + let host = Arc::new(TestHost { + reject_socks5_listener: true, + ..Default::default() + }); + let (packet_sink, _packet_receiver) = tokio::sync::mpsc::channel(16); + let adapters = adapters_with_host(host.clone(), None, Arc::new(packet_sink)); + (CoreInstance::new(config, adapters).unwrap(), host) + } + + let mut enabled_config = test_config("gateway-enabled-by-default"); + enabled_config.connectivity.runtime.gateway.socks5_bind = + Some("127.0.0.1:1080".parse().unwrap()); + let (enabled, _) = build(enabled_config); + let error = enabled.start().await.unwrap_err(); + assert!(error.to_string().contains("rejected SOCKS5 listener")); + assert_eq!(enabled.state(), CoreInstanceState::Stopped); + + let mut disabled_config = test_config("gateway-disabled-by-plan"); + disabled_config.connectivity.startup_plan.gateway = false; + disabled_config.peer.snapshot.runtime.core.routes.ipv4 = + Some(IpPrefix::new("10.144.0.1".parse().unwrap(), 24).unwrap()); + let (disabled, _) = build(disabled_config.clone()); + let mut updated = runtime_snapshot(&disabled_config); + updated.services.gateway.socks5_bind = Some("127.0.0.1:1080".parse().unwrap()); + disabled.update_runtime_config(updated).await.unwrap(); + disabled.start().await.unwrap(); + assert_eq!(disabled.state(), CoreInstanceState::Running); + let socket = disabled + .data_plane_udp_bind(0, Duration::from_secs(1)) + .await + .expect("startup plan must not disable the data-plane runtime"); + drop(socket); + disabled.stop().await; + } + + #[cfg(feature = "proxy-smoltcp-stack")] + #[tokio::test] + async fn failed_port_forward_start_releases_started_listeners() { + let mut config = test_config("port-forward-start-rollback"); + config.connectivity.runtime.gateway.port_forwards = vec![ + PortForwardConfig { + bind_addr: "127.0.0.1:18080".parse().unwrap(), + dst_addr: "10.144.0.2:80".parse().unwrap(), + proto: "tcp".to_owned(), + }, + PortForwardConfig { + bind_addr: "127.0.0.1:18081".parse().unwrap(), + dst_addr: "10.144.0.2:81".parse().unwrap(), + proto: "unsupported".to_owned(), + }, + ]; + let host = Arc::new(TestHost::default()); + let (packet_sink, _packet_receiver) = tokio::sync::mpsc::channel(16); + let adapters = adapters_with_host(host.clone(), None, Arc::new(packet_sink)); + let instance = CoreInstance::new(config, adapters).unwrap(); + + let error = instance.start().await.unwrap_err(); + + assert!(error.to_string().contains("unsupported protocol")); + assert_eq!(instance.state(), CoreInstanceState::Stopped); + assert_eq!(host.active_tcp_listeners.load(Ordering::Relaxed), 0); + } + + #[cfg(feature = "proxy-packet")] + #[tokio::test] + async fn runtime_core_instance_owns_wrapped_transport_source_nat() { + let mut config = test_config("wrapped-source"); + config.peer.snapshot.flags.enable_kcp_proxy = true; + config.peer.snapshot.flags.disable_kcp_input = true; + let engine = Arc::new(RecordingProxyService::default()); + let instance = build_with_engines( + config, + WrappedTransportEngines { + kcp: Some(engine.clone()), + quic: None, + }, + ) + .unwrap(); + + instance.start().await.unwrap(); + assert!( + instance.wrapped_transport_is_started( + WrappedTransportKind::Kcp, + WrappedTransportRole::Source, + ) + ); + assert!( + instance + .wrapped_tcp_proxy_entry_snapshots( + WrappedTransportKind::Kcp, + WrappedTransportRole::Source, + ) + .is_empty() + ); + + instance.stop().await; + assert!( + !instance.wrapped_transport_is_started( + WrappedTransportKind::Kcp, + WrappedTransportRole::Source, + ) + ); + assert_eq!(engine.stop_calls.load(Ordering::Relaxed), 1); + } + + #[cfg(feature = "proxy-packet")] + #[tokio::test] + async fn runtime_core_instance_owns_wrapped_transport_destination_sessions() { + let mut config = test_config("wrapped-destination"); + config.peer.snapshot.flags.enable_kcp_proxy = false; + config.peer.snapshot.flags.disable_kcp_input = false; + let engine = Arc::new(RecordingProxyService::default()); + let (connections, mut connection_receiver) = tokio::sync::mpsc::unbounded_channel(); + let host = Arc::new(TestHost { + proxy_nat_connections: Some(connections), + ..Default::default() + }); + let (packet_sink, _packet_receiver) = tokio::sync::mpsc::channel(16); + let mut adapters = adapters_with_host(host, None, Arc::new(packet_sink)); + adapters.wrapped_transports = WrappedTransportEngines { + kcp: Some(engine.clone()), + quic: None, + }; + let instance = CoreInstance::new(config, adapters).unwrap(); + + instance.start().await.unwrap(); + assert!(instance.wrapped_transport_is_started( + WrappedTransportKind::Kcp, + WrappedTransportRole::Destination, + )); + + let destination: SocketAddr = "127.0.0.1:20100".parse().unwrap(); + let ingress = engine + .destination_ingress() + .expect("core should inject a destination ingress"); + let (core_stream, peer_stream) = tokio::io::duplex(1024); + ingress + .submit( + crate::gateway::proxy::wrapped_transport::WrappedTransportAcceptedStream { + src: "10.0.0.2:40000".parse().unwrap(), + dst: destination, + initial_acl_packet_size: 16, + stream: Box::new(core_stream), + }, + ) + .await + .unwrap(); + let (connected_destination, destination_stream) = tokio::time::timeout( + std::time::Duration::from_secs(2), + connection_receiver.recv(), + ) + .await + .expect("core should request the destination socket") + .unwrap(); + assert_eq!(connected_destination, destination); + tokio::time::timeout(std::time::Duration::from_secs(2), async { + loop { + let entries = instance.wrapped_tcp_proxy_entry_snapshots( + WrappedTransportKind::Kcp, + WrappedTransportRole::Destination, + ); + if entries.iter().any(|entry| { + entry.state + == crate::gateway::proxy::tcp_proxy_engine::TcpNatEntryState::Connected + }) { + break; + } + tokio::task::yield_now().await; + } + }) + .await + .expect("core should own the connected destination entry"); + + drop(peer_stream); + drop(destination_stream); + tokio::time::timeout(std::time::Duration::from_secs(2), async { + while !instance + .wrapped_tcp_proxy_entry_snapshots( + WrappedTransportKind::Kcp, + WrappedTransportRole::Destination, + ) + .is_empty() + { + tokio::task::yield_now().await; + } + }) + .await + .expect("completed destination entry should be removed"); + + let (core_stream, _blocked_peer_stream) = tokio::io::duplex(1024); + ingress + .submit( + crate::gateway::proxy::wrapped_transport::WrappedTransportAcceptedStream { + src: "10.0.0.2:40001".parse().unwrap(), + dst: destination, + initial_acl_packet_size: 16, + stream: Box::new(core_stream), + }, + ) + .await + .unwrap(); + let (connected_destination, _blocked_destination_stream) = tokio::time::timeout( + std::time::Duration::from_secs(2), + connection_receiver.recv(), + ) + .await + .expect("second destination session should request a socket") + .unwrap(); + assert_eq!(connected_destination, destination); + tokio::time::timeout(std::time::Duration::from_secs(2), async { + while instance + .wrapped_tcp_proxy_entry_snapshots( + WrappedTransportKind::Kcp, + WrappedTransportRole::Destination, + ) + .is_empty() + { + tokio::task::yield_now().await; + } + }) + .await + .expect("blocked destination session should be visible"); + + tokio::time::timeout(std::time::Duration::from_secs(2), instance.stop()) + .await + .expect("stop should cancel core-owned destination sessions"); + assert!( + instance + .wrapped_tcp_proxy_entry_snapshots( + WrappedTransportKind::Kcp, + WrappedTransportRole::Destination, + ) + .is_empty() + ); + assert!( + ingress + .submit( + crate::gateway::proxy::wrapped_transport::WrappedTransportAcceptedStream { + src: "10.0.0.2:40002".parse().unwrap(), + dst: destination, + initial_acl_packet_size: 16, + stream: Box::new(tokio::io::duplex(64).0), + }, + ) + .await + .is_err() + ); + } + + #[tokio::test] + async fn runtime_core_instance_owns_the_transport_proxy_cidr_table() { + let mut config = test_config("transport-proxy-cidr"); + config.connectivity.runtime.proxy.forward_by_system = true; + config.peer.snapshot.runtime.core.routes.proxy_networks = + vec![proxy_network("192.0.2.0/24", Some("198.51.100.0/24"))]; + let mut updated = runtime_snapshot(&config); + let proxy = Arc::new(RecordingProxyService::default()); + let instance = build_with_engines( + config, + WrappedTransportEngines { + kcp: Some(proxy.clone()), + quic: None, + }, + ) + .unwrap(); + + assert_eq!(instance.node_snapshot().await.proxy_networks.len(), 1); + instance.start().await.unwrap(); + assert_eq!(proxy.start_calls.load(Ordering::Relaxed), 1); + + Arc::make_mut(&mut updated.peer) + .runtime + .core + .routes + .proxy_networks = vec![proxy_network("203.0.113.0/24", Some("10.20.30.0/24"))]; + instance.update_runtime_config(updated).await.unwrap(); + let proxy_networks = instance.node_snapshot().await.proxy_networks; + assert_eq!(proxy_networks.len(), 1); + assert_eq!( + proxy_networks[0].real.address, + "203.0.113.0".parse::().unwrap() + ); + assert_eq!( + proxy_networks[0].mapped.as_ref().unwrap().address, + "10.20.30.0".parse::().unwrap() + ); + + instance.stop().await; + assert!(instance.start().await.is_err()); + assert_eq!(proxy.stop_calls.load(Ordering::Relaxed), 1); + } + + #[tokio::test] + async fn runtime_core_rejects_invalid_acl_runtime_snapshot() { + let config = test_config("explicit-acl"); + let instance = build_instance(config.clone()).unwrap(); + assert_eq!(instance.acl_whitelist_snapshot(), Default::default()); + + let mut updated = runtime_snapshot(&config); + updated.services.acl.tcp_whitelist = vec!["invalid".to_owned()]; + let error = instance.update_runtime_config(updated).await.unwrap_err(); + assert!(error.to_string().contains("Invalid port number")); + assert!(instance.acl_whitelist_snapshot().tcp_ports.is_empty()); + instance.start().await.unwrap(); + instance.stop().await; + } + + #[cfg(feature = "dhcp-ipv4")] + #[tokio::test] + async fn runtime_core_accepts_explicit_dhcp_runtime_snapshot() { + let config = test_config("explicit-dhcp"); + let instance = build_instance(config.clone()).unwrap(); + let mut updated = runtime_snapshot(&config); + updated.services.dhcp_ipv4 = true; + instance.update_runtime_config(updated).await.unwrap(); + let error = instance.start().await.unwrap_err(); + assert!(error.to_string().contains("no host adapter was provided")); + assert_eq!(instance.state(), CoreInstanceState::Stopped); + } + + #[cfg(feature = "public-ipv6-provider")] + #[tokio::test] + async fn runtime_core_accepts_explicit_public_ipv6_runtime_snapshot() { + let config = test_config("explicit-public-ipv6"); + let instance = build_instance(config.clone()).unwrap(); + let mut updated = runtime_snapshot(&config); + updated.services.public_ipv6_provider.provider_enabled = true; + updated.services.public_ipv6_provider.provider_supported = true; + updated.services.public_ipv6_provider.configured_prefix = + Some("fd00::/64".parse().unwrap()); + instance.update_runtime_config(updated).await.unwrap(); + + let error = instance.start().await.unwrap_err(); + assert!(error.to_string().contains("not a valid global unicast")); + assert_eq!(instance.state(), CoreInstanceState::Stopped); + } + + #[cfg(not(feature = "proxy-packet"))] + #[test] + fn runtime_core_rejects_packet_proxy_requests_when_unavailable() { + let mut config = test_config("packet-proxy-unavailable"); + config.connectivity.runtime.proxy.enable_exit_node = true; + + let error = match build_instance(config) { + Ok(_) => panic!("packet proxy request unexpectedly succeeded"), + Err(error) => error, + }; + + assert!( + error + .to_string() + .contains("does not include packet proxy services") + ); + } + + #[cfg(not(feature = "proxy-smoltcp-stack"))] + #[tokio::test] + async fn runtime_core_rejects_unavailable_gateway_updates() { + let config = test_config("smoltcp-gateway-update-unavailable"); + let instance = build_instance(config.clone()).unwrap(); + let mut updated = runtime_snapshot(&config); + updated.services.gateway.socks5_bind = Some("127.0.0.1:1080".parse().unwrap()); + + let error = instance.update_runtime_config(updated).await.unwrap_err(); + + assert!( + error + .to_string() + .contains("does not include the smoltcp gateway") + ); + } + + #[tokio::test] + async fn stopping_while_transport_proxy_starts_rolls_back_once() { + let mut config = test_config("blocking-transport-proxy"); + config.connectivity.runtime.proxy.forward_by_system = true; + config.peer.snapshot.runtime.core.routes.proxy_networks = + vec![proxy_network("10.1.2.0/24", None)]; + config.connectivity.initial_peers = vec!["tcp://127.0.0.1:29998".parse().unwrap()]; + let (proxy, start_gate) = RecordingProxyService::blocking(); + let instance = build_with_engines( + config, + WrappedTransportEngines { + kcp: Some(proxy.clone()), + quic: None, + }, + ) + .unwrap(); + + let start_task = tokio::spawn({ + let instance = instance.clone(); + async move { instance.start().await } + }); + start_gate.entered.notified().await; + let stop_task = tokio::spawn({ + let instance = instance.clone(); + async move { instance.stop().await } + }); + while !instance.cancel.is_cancelled() { + tokio::task::yield_now().await; + } + + assert!(start_task.await.unwrap().is_err()); + stop_task.await.unwrap(); + assert_eq!(instance.list_connectors().len(), 1); + assert_eq!(instance.state(), CoreInstanceState::Stopped); + assert!(!instance.is_ready()); + assert_eq!(proxy.start_calls.load(Ordering::Relaxed), 1); + assert_eq!(proxy.stop_calls.load(Ordering::Relaxed), 1); + } + + #[tokio::test] + async fn start_serializes_runtime_updates() { + let mut config = test_config("serialized-start"); + config.connectivity.runtime.proxy.forward_by_system = true; + config.peer.snapshot.runtime.core.routes.proxy_networks = + vec![proxy_network("10.1.4.0/24", None)]; + let mut updated = runtime_snapshot(&config); + Arc::make_mut(&mut updated.peer).runtime.core.node.hostname = + Some("updated-after-start".to_owned()); + let (proxy, start_gate) = RecordingProxyService::blocking(); + let instance = build_with_engines( + config, + WrappedTransportEngines { + kcp: Some(proxy), + quic: None, + }, + ) + .unwrap(); + + let start = tokio::spawn({ + let instance = instance.clone(); + async move { instance.start().await } + }); + start_gate.entered.notified().await; + + let update = instance.update_runtime_config(updated); + tokio::pin!(update); + assert!(matches!( + futures::poll!(update.as_mut()), + std::task::Poll::Pending + )); + + start_gate.release.notify_one(); + start.await.unwrap().unwrap(); + update.await.unwrap(); + assert_eq!( + instance.node_snapshot().await.hostname, + "updated-after-start" + ); + instance.stop().await; + } + + #[tokio::test] + async fn aborting_start_stops_partial_runtime() { + let mut config = test_config("aborted-start"); + config.connectivity.runtime.proxy.forward_by_system = true; + config.peer.snapshot.runtime.core.routes.proxy_networks = + vec![proxy_network("10.1.3.0/24", None)]; + let (proxy, start_gate) = RecordingProxyService::blocking(); + let instance = build_with_engines( + config, + WrappedTransportEngines { + kcp: Some(proxy.clone()), + quic: None, + }, + ) + .unwrap(); + + let start = tokio::spawn({ + let instance = instance.clone(); + async move { instance.start().await } + }); + start_gate.entered.notified().await; + start.abort(); + assert!(start.await.unwrap_err().is_cancelled()); + + tokio::time::timeout(Duration::from_secs(2), async { + while instance.state() != CoreInstanceState::Stopped { + tokio::task::yield_now().await; + } + }) + .await + .expect("aborted activation should recover the instance"); + assert!(!instance.is_ready()); + assert_eq!(proxy.start_calls.load(Ordering::Relaxed), 1); + assert_eq!(proxy.stop_calls.load(Ordering::Relaxed), 1); + } + + #[tokio::test] + async fn invalid_initial_peer_fails_during_construction() { + let mut config = test_config("invalid-initial-peer"); + config.connectivity.initial_peers = + vec!["unsupported://peer.example:1234".parse().unwrap()]; + + let error = build_instance(config) + .err() + .expect("invalid initial peer should fail construction"); + assert!( + error + .to_string() + .contains("unsupported core manual connector URL"), + "unexpected construction error: {error:#}" + ); + } + + #[tokio::test] + async fn runtime_core_instances_keep_lifecycle_and_connectors_isolated() { + let instance_a = build_instance(test_config("instance-a")).unwrap(); + let instance_b = build_instance(test_config("instance-b")).unwrap(); + let connector_a: Url = "tcp://127.0.0.1:21001".parse().unwrap(); + let connector_b: Url = "udp://127.0.0.1:21002".parse().unwrap(); + + instance_a.add_connector(connector_a.clone()).unwrap(); + instance_b.add_connector(connector_b.clone()).unwrap(); + assert_eq!(instance_a.list_connectors()[0].url, connector_a); + assert_eq!(instance_b.list_connectors()[0].url, connector_b); + instance_a.clear_connectors(); + instance_b.clear_connectors(); + + let (start_a, start_b) = tokio::join!(instance_a.start(), instance_b.start()); + start_a.unwrap(); + start_b.unwrap(); + assert_eq!(instance_a.state(), CoreInstanceState::Running); + assert_eq!(instance_b.state(), CoreInstanceState::Running); + + instance_a.stop().await; + assert_eq!(instance_a.state(), CoreInstanceState::Stopped); + assert_eq!(instance_b.state(), CoreInstanceState::Running); + instance_b.stop().await; + assert_eq!(instance_b.state(), CoreInstanceState::Stopped); + } + + #[cfg(unix)] + #[tokio::test] + async fn stop_cancels_pending_listener_start() { + let state = Arc::new(BlockingListenerState::default()); + let mut config = test_config("pending-listener"); + config.connectivity.listeners = Some(ListenerRuntimeConfig::new( + vec!["unix:///tmp/easytier-pending-listener".parse().unwrap()], + false, + SocketContext::default(), + )); + let instance = build_with_engines_and_listener( + config, + WrappedTransportEngines::default(), + Some(Arc::new(BlockingExternalListenerFactory { + state: state.clone(), + })), + ) + .unwrap(); + let start_instance = instance.clone(); + let start_task = + AbortOnDropHandle::new(tokio::spawn(async move { start_instance.start().await })); + let start_result = tokio::time::timeout(Duration::from_secs(1), async { + state.start_entered.notified().await; + instance.stop().await; + start_task.await.unwrap() + }) + .await + .expect("listener cancellation should complete promptly"); + + assert!(start_result.is_err()); + assert_eq!(instance.state(), CoreInstanceState::Stopped); + assert_eq!(state.drop_calls.load(Ordering::Relaxed), 1); + } + + #[tokio::test] + async fn external_listener_uses_core_running_listener_registry() { + let external_url: Url = "unix:///tmp/easytier-external-listener-test" + .parse() + .unwrap(); + let mut config = test_config("external-listener-registry"); + config.connectivity.listeners = Some(ListenerRuntimeConfig::new( + vec![external_url.clone()], + false, + SocketContext::default(), + )); + let instance = build_with_engines_and_listener( + config, + WrappedTransportEngines::default(), + Some(Arc::new(ReadyExternalListenerFactory)), + ) + .unwrap(); + + instance.start().await.unwrap(); + let running = instance.running_listeners(); + assert_eq!(running.len(), 2); + assert!(running.iter().any(|url| url.scheme() == "ring")); + assert!(running.contains(&external_url)); + assert_eq!(instance.node_snapshot().await.listeners, running); + + instance.stop().await; + assert!(instance.running_listeners().is_empty()); + } +} diff --git a/easytier-core/src/instance/vpn_portal_extension.rs b/easytier-core/src/instance/vpn_portal_extension.rs new file mode 100644 index 00000000..10e26eef --- /dev/null +++ b/easytier-core/src/instance/vpn_portal_extension.rs @@ -0,0 +1,13 @@ +use crate::{ + gateway::vpn_portal::VpnPortalInfoSnapshot, + instance::{CoreInstance, CoreInstanceHost}, +}; + +impl CoreInstance +where + H: CoreInstanceHost, +{ + pub async fn vpn_portal_info(&self) -> VpnPortalInfoSnapshot { + self.vpn_portal.info_snapshot().await + } +} diff --git a/easytier-core/src/lib.rs b/easytier-core/src/lib.rs new file mode 100644 index 00000000..440e33df --- /dev/null +++ b/easytier-core/src/lib.rs @@ -0,0 +1,21 @@ +pub mod config; +pub mod connectivity; +pub mod events; +pub mod foundation; +pub mod gateway; +pub mod host; +pub mod instance; +pub mod listener; +#[cfg(feature = "management-rpc")] +pub mod management; +pub mod packet; +pub mod peers; +pub mod process_runtime; +pub mod rpc; +pub mod socket; +pub mod tunnel; + +#[cfg(any(test, target_os = "wasi"))] +pub mod wasi; + +pub(crate) use easytier_proto as proto; diff --git a/easytier-core/src/listener/mod.rs b/easytier-core/src/listener/mod.rs new file mode 100644 index 00000000..4f6955d0 --- /dev/null +++ b/easytier-core/src/listener/mod.rs @@ -0,0 +1,1016 @@ +use std::{fmt::Debug, future::Future, pin::Pin, sync::Arc, time::Duration}; + +use anyhow::Context as _; +use async_trait::async_trait; +use tokio::{ + sync::{Mutex, mpsc}, + task::JoinSet, +}; +use tokio_util::sync::CancellationToken; +use url::Url; + +use crate::{ + events::{CoreEvent, CoreEventSink}, + socket::{SocketContext, SocketListener}, +}; + +pub mod plan; +pub mod transport; + +pub trait ExternalListenerFactory: Send + Sync + 'static +where + Accepted: Send + 'static, +{ + fn supports_scheme(&self, scheme: &str) -> bool; + + fn create( + &self, + request: ExternalListenerRequest, + ) -> Box>; +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct ExternalListenerRequest { + pub url: Url, + pub socket_context: SocketContext, +} + +#[async_trait] +pub trait AcceptedSocketHandler: Send + Sync { + async fn handle_accepted_socket(&self, accepted: Accepted) -> anyhow::Result<()>; +} + +#[async_trait] +impl AcceptedSocketHandler for F +where + Accepted: Send + 'static, + F: Fn(Accepted) -> Fut + Send + Sync, + Fut: Future> + Send, +{ + async fn handle_accepted_socket(&self, accepted: Accepted) -> anyhow::Result<()> { + self(accepted).await + } +} + +#[derive(Debug, Default)] +pub(crate) struct RunningListenerRegistry { + listeners: std::sync::Mutex>, +} + +impl RunningListenerRegistry { + pub(crate) fn running_listeners(&self) -> Vec { + self.listeners + .lock() + .unwrap() + .iter() + .map(|(url, _)| url.clone()) + .collect() + } + + fn register(self: &Arc, url: Url) -> RunningListenerRegistration { + let mut listeners = self.listeners.lock().unwrap(); + if let Some((_, count)) = listeners.iter_mut().find(|(listener, _)| listener == &url) { + *count += 1; + } else { + listeners.push((url.clone(), 1)); + } + RunningListenerRegistration { + registry: self.clone(), + url, + } + } + + fn unregister(&self, url: &Url) { + let mut listeners = self.listeners.lock().unwrap(); + let Some(index) = listeners.iter().position(|(listener, _)| listener == url) else { + return; + }; + if listeners[index].1 == 1 { + listeners.remove(index); + } else { + listeners[index].1 -= 1; + } + } +} + +struct RunningListenerRegistration { + registry: Arc, + url: Url, +} + +impl Drop for RunningListenerRegistration { + fn drop(&mut self) { + self.registry.unregister(&self.url); + } +} + +impl crate::connectivity::LocalListenerUrls for RunningListenerRegistry { + fn local_listener_urls(&self) -> Vec { + self.running_listeners() + } +} + +type ListenerCreatorArc = + Arc Box> + Send + Sync>; + +#[derive(Clone)] +pub(crate) struct ListenerFactory { + creator: ListenerCreatorArc, + must_succeed: bool, +} + +impl ListenerFactory { + pub(crate) fn new(creator: C, must_succeed: bool) -> Self + where + C: Fn() -> Box> + Send + Sync + 'static, + { + Self { + creator: Arc::new(creator), + must_succeed, + } + } +} + +#[derive(Debug, Clone)] +pub struct ListenerManagerOptions { + pub max_listen_retries: usize, + pub listen_retry_delay: Duration, + pub accept_retry_delay: Duration, +} + +impl Default for ListenerManagerOptions { + fn default() -> Self { + Self { + max_listen_retries: 5, + listen_retry_delay: Duration::from_secs(1), + accept_retry_delay: Duration::from_secs(1), + } + } +} + +pub struct ListenerManager { + factories: Vec>, + handler: Arc, + events: Arc, + registry: Arc, + options: ListenerManagerOptions, + operation: Mutex<()>, + cancel: CancellationToken, + tasks: Mutex>, + handler_tasks: Arc>>, + accepted_tasks: AcceptedTaskSpawner, + accepted_task_rx: std::sync::Mutex>>, +} + +type AcceptedTask = Pin + Send + 'static>>; + +#[derive(Clone)] +struct AcceptedTaskSpawner { + tx: mpsc::UnboundedSender, +} + +impl AcceptedTaskSpawner { + fn new() -> (Self, mpsc::UnboundedReceiver) { + let (tx, rx) = mpsc::unbounded_channel(); + (Self { tx }, rx) + } + + fn spawn(&self, future: F) + where + F: Future + Send + 'static, + { + if self.tx.send(Box::pin(future)).is_err() { + tracing::warn!("accepted socket handler task runner stopped"); + } + } +} + +impl ListenerManager +where + Accepted: Send + 'static, + H: AcceptedSocketHandler + ?Sized + 'static, +{ + pub fn new_with_events(handler: Arc, events: Arc) -> Self { + Self::new_with_options(handler, events, ListenerManagerOptions::default()) + } + + pub fn new_with_options( + handler: Arc, + events: Arc, + options: ListenerManagerOptions, + ) -> Self { + Self::new_with_registry( + handler, + events, + Arc::new(RunningListenerRegistry::default()), + options, + ) + } + + pub(crate) fn new_with_registry( + handler: Arc, + events: Arc, + registry: Arc, + options: ListenerManagerOptions, + ) -> Self { + let (accepted_tasks, accepted_task_rx) = AcceptedTaskSpawner::new(); + Self { + factories: Vec::new(), + handler, + events, + registry, + options, + operation: Mutex::new(()), + cancel: CancellationToken::new(), + tasks: Mutex::new(JoinSet::new()), + handler_tasks: Arc::new(Mutex::new(JoinSet::new())), + accepted_tasks, + accepted_task_rx: std::sync::Mutex::new(Some(accepted_task_rx)), + } + } + + pub fn add_listener(&mut self, creator: C, must_succeed: bool) + where + C: Fn() -> Box> + Send + Sync + 'static, + { + self.factories + .push(ListenerFactory::new(creator, must_succeed)); + } + + pub(crate) fn add_factory(&mut self, factory: ListenerFactory) { + self.factories.push(factory); + } + + pub async fn run(&self) -> anyhow::Result<()> { + let _operation = self.operation.lock().await; + if self.cancel.is_cancelled() { + anyhow::bail!("listener manager is stopped"); + } + let accepted_task_rx = self + .accepted_task_rx + .lock() + .unwrap() + .take() + .ok_or_else(|| anyhow::anyhow!("listener manager is one-shot and already ran"))?; + let mut initial_listeners = Vec::with_capacity(self.factories.len()); + for factory in &self.factories { + initial_listeners.push(if factory.must_succeed { + let listener = tokio::select! { + _ = self.cancel.cancelled() => { + anyhow::bail!("listener manager stopped during startup") + } + result = listen_once(factory.creator.clone()) => { + result.with_context(|| "required listener failed to start")? + } + }; + Some(listener) + } else { + None + }); + } + if self.cancel.is_cancelled() { + anyhow::bail!("listener manager stopped during startup"); + } + let initial_listeners = initial_listeners + .into_iter() + .map(|listener| { + listener.map(|listener| { + RegisteredListener::new(listener, self.events.clone(), self.registry.clone()) + }) + }) + .collect::>(); + + let mut tasks = self.tasks.lock().await; + tasks.spawn(run_accepted_task_runner( + accepted_task_rx, + self.handler_tasks.clone(), + self.cancel.clone(), + )); + + for (factory, initial_listener) in self.factories.iter().zip(initial_listeners) { + let cancel = self.cancel.clone(); + let listener = run_listener( + factory.creator.clone(), + self.handler.clone(), + self.events.clone(), + self.registry.clone(), + self.options.clone(), + self.accepted_tasks.clone(), + initial_listener, + ); + tasks.spawn(async move { + tokio::select! { + _ = cancel.cancelled() => {} + _ = listener => {} + } + }); + } + + Ok(()) + } + + pub async fn stop(&self) { + self.cancel.cancel(); + let _operation = self.operation.lock().await; + let mut tasks = self.tasks.lock().await; + tasks.abort_all(); + while tasks.join_next().await.is_some() {} + drop(tasks); + + let mut handler_tasks = self.handler_tasks.lock().await; + handler_tasks.abort_all(); + while handler_tasks.join_next().await.is_some() {} + } +} + +async fn listen_once( + creator: ListenerCreatorArc, +) -> anyhow::Result>> +where + Accepted: Send + 'static, +{ + let mut listener = creator(); + match listener.listen().await { + Ok(()) => Ok(listener), + Err(error) => Err(error), + } +} + +async fn run_listener( + creator: ListenerCreatorArc, + handler: Arc, + events: Arc, + registry: Arc, + options: ListenerManagerOptions, + accepted_tasks: AcceptedTaskSpawner, + mut initial_listener: Option>, +) where + Accepted: Send + 'static, + H: AcceptedSocketHandler + ?Sized + 'static, +{ + let mut listen_error_count = 0; + loop { + let registered_listener = match initial_listener.take() { + Some(listener) => listener, + None => { + let mut listener = creator(); + match listener.listen().await { + Ok(()) => { + listen_error_count = 0; + RegisteredListener::new(listener, events.clone(), registry.clone()) + } + Err(error) => { + listen_error_count += 1; + let will_retry = listen_error_count <= options.max_listen_retries; + events.emit(CoreEvent::ListenerAddFailed { + url: listener.local_url(), + error: format!("{error:?}"), + retry_count: listen_error_count, + will_retry, + }); + tracing::error!(?error, ?listener, "listener listen error"); + if !will_retry { + return; + } + crate::foundation::time::sleep(options.listen_retry_delay).await; + continue; + } + } + } + }; + let mut listener = registered_listener.listener; + let _registration = registered_listener.registration; + + loop { + let listener_url = listener.local_url(); + let accepted = match listener.accept().await { + Ok(accepted) => accepted, + Err(error) => { + events.emit(CoreEvent::ListenerAcceptFailed { + url: listener_url.clone(), + error: format!("{error:?}"), + }); + tracing::error!(?error, ?listener, "listener accept error"); + crate::foundation::time::sleep(options.accept_retry_delay).await; + break; + } + }; + + events.emit(CoreEvent::ListenerSocketAccepted { + url: listener_url.clone(), + }); + let handler = handler.clone(); + let events = events.clone(); + accepted_tasks.spawn(async move { + if let Err(error) = handler.handle_accepted_socket(accepted).await { + events.emit(CoreEvent::ListenerAcceptedSocketHandleFailed { + url: listener_url, + error: format!("{error:?}"), + }); + } + }); + } + } +} + +async fn run_accepted_task_runner( + mut accepted_task_rx: mpsc::UnboundedReceiver, + handler_tasks: Arc>>, + cancel: CancellationToken, +) { + loop { + tokio::select! { + _ = cancel.cancelled() => break, + maybe_task = accepted_task_rx.recv() => { + match maybe_task { + Some(task) => { + handler_tasks.lock().await.spawn(task); + } + None => break, + } + } + _ = crate::foundation::time::sleep(Duration::from_secs(1)) => { + let mut handler_tasks = handler_tasks.lock().await; + while let Some(task) = handler_tasks.try_join_next() { + if let Err(error) = task { + tracing::error!(?error, "accepted socket handler task failed"); + } + } + } + } + } + + let mut handler_tasks = handler_tasks.lock().await; + if cancel.is_cancelled() { + handler_tasks.abort_all(); + } + while let Some(task) = handler_tasks.join_next().await { + if let Err(error) = task { + tracing::error!(?error, "accepted socket handler task failed"); + } + } +} + +struct RegisteredListener +where + Accepted: Send + 'static, +{ + listener: Box>, + registration: ListenerRegistration, +} + +impl RegisteredListener +where + Accepted: Send + 'static, +{ + fn new( + listener: Box>, + events: Arc, + registry: Arc, + ) -> Self { + let url = listener.local_url(); + let registry = registry.register(url.clone()); + events.emit(CoreEvent::ListenerAdded { + url: url.clone(), + connection_counter: listener.connection_counter(), + }); + Self { + listener, + registration: ListenerRegistration { + url, + events, + registry: Some(registry), + }, + } + } +} + +struct ListenerRegistration { + url: Url, + events: Arc, + registry: Option, +} + +impl Drop for ListenerRegistration { + fn drop(&mut self) { + drop(self.registry.take()); + self.events.emit(CoreEvent::ListenerRemoved { + url: self.url.clone(), + }); + } +} + +#[cfg(test)] +mod tests { + use std::{ + collections::VecDeque, + fmt, + sync::{ + Mutex, + atomic::{AtomicUsize, Ordering}, + }, + }; + + use super::*; + + #[derive(Debug)] + struct MockListener { + url: Url, + listen_results: Arc>>>, + accepts: Arc>>>, + listen_count: Arc, + drop_count: Arc, + } + + impl Drop for MockListener { + fn drop(&mut self) { + self.drop_count.fetch_add(1, Ordering::Relaxed); + } + } + + #[async_trait] + impl SocketListener for MockListener { + type Accepted = usize; + + async fn listen(&mut self) -> anyhow::Result<()> { + self.listen_count.fetch_add(1, Ordering::Relaxed); + self.listen_results + .lock() + .unwrap() + .pop_front() + .unwrap_or(Ok(())) + } + + async fn accept(&mut self) -> anyhow::Result { + let next_accept = self.accepts.lock().unwrap().pop_front(); + match next_accept { + Some(ret) => ret, + None => std::future::pending().await, + } + } + + fn local_url(&self) -> Url { + self.url.clone() + } + } + + #[derive(Debug)] + struct BlockingListenListener { + started: Arc, + } + + #[async_trait] + impl SocketListener for BlockingListenListener { + type Accepted = usize; + + async fn listen(&mut self) -> anyhow::Result<()> { + self.started.notify_one(); + std::future::pending().await + } + + async fn accept(&mut self) -> anyhow::Result { + std::future::pending().await + } + + fn local_url(&self) -> Url { + "mock://blocking-start".parse().unwrap() + } + } + + #[derive(Debug)] + struct MockHandler { + accepted: Mutex>, + } + + #[async_trait] + impl AcceptedSocketHandler for MockHandler { + async fn handle_accepted_socket(&self, accepted: usize) -> anyhow::Result<()> { + self.accepted.lock().unwrap().push(accepted); + Ok(()) + } + } + + #[derive(Default)] + struct Events { + events: Mutex>, + } + + impl Debug for Events { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.debug_struct("Events").finish() + } + } + + impl CoreEventSink for Events { + fn emit(&self, event: CoreEvent) { + self.events.lock().unwrap().push(event); + } + } + + #[test] + fn running_listener_registry_reference_counts_duplicate_urls() { + let registry = Arc::new(RunningListenerRegistry::default()); + let url: Url = "tcp://127.0.0.1:11010".parse().unwrap(); + let first = registry.register(url.clone()); + let second = registry.register(url.clone()); + assert_eq!(registry.running_listeners(), vec![url.clone()]); + + drop(first); + assert_eq!(registry.running_listeners(), vec![url.clone()]); + drop(second); + assert!(registry.running_listeners().is_empty()); + } + + #[tokio::test] + async fn required_listener_reuses_successful_initial_listen() { + let handler = Arc::new(MockHandler { + accepted: Mutex::new(Vec::new()), + }); + let events = Arc::new(Events::default()); + let listen_count = Arc::new(AtomicUsize::new(0)); + let drop_count = Arc::new(AtomicUsize::new(0)); + let accepts = Arc::new(Mutex::new(VecDeque::from([Ok::<_, anyhow::Error>(7)]))); + let mut manager = ListenerManager::new_with_options( + handler.clone(), + events.clone(), + ListenerManagerOptions { + accept_retry_delay: Duration::from_millis(1), + listen_retry_delay: Duration::from_millis(1), + max_listen_retries: 0, + }, + ); + + let listener_accepts = accepts.clone(); + let listener_listen_count = listen_count.clone(); + let listener_drop_count = drop_count.clone(); + manager.add_listener( + move || { + Box::new(MockListener { + url: "mock://required".parse().unwrap(), + listen_results: Arc::new(Mutex::new(VecDeque::from([Ok(())]))), + accepts: listener_accepts.clone(), + listen_count: listener_listen_count.clone(), + drop_count: listener_drop_count.clone(), + }) + }, + true, + ); + + manager.run().await.unwrap(); + crate::foundation::time::sleep(Duration::from_millis(20)).await; + + assert_eq!(listen_count.load(Ordering::Relaxed), 1); + assert_eq!(handler.accepted.lock().unwrap().as_slice(), &[7]); + assert!( + events + .events + .lock() + .unwrap() + .iter() + .any(|event| matches!(event, CoreEvent::ListenerAdded { .. })) + ); + } + + #[tokio::test] + async fn optional_listener_retries_until_listen_succeeds() { + let handler = Arc::new(MockHandler { + accepted: Mutex::new(Vec::new()), + }); + let events = Arc::new(Events::default()); + let listen_count = Arc::new(AtomicUsize::new(0)); + let listen_results = Arc::new(Mutex::new(VecDeque::from([ + Err(anyhow::anyhow!("not ready")), + Ok(()), + ]))); + let accepts = Arc::new(Mutex::new(VecDeque::from([Ok::<_, anyhow::Error>(3)]))); + let mut manager = ListenerManager::new_with_options( + handler.clone(), + events.clone(), + ListenerManagerOptions { + accept_retry_delay: Duration::from_millis(1), + listen_retry_delay: Duration::from_millis(1), + max_listen_retries: 2, + }, + ); + + let listener_results = listen_results.clone(); + let listener_accepts = accepts.clone(); + let listener_listen_count = listen_count.clone(); + manager.add_listener( + move || { + Box::new(MockListener { + url: "mock://optional".parse().unwrap(), + listen_results: listener_results.clone(), + accepts: listener_accepts.clone(), + listen_count: listener_listen_count.clone(), + drop_count: Arc::new(AtomicUsize::new(0)), + }) + }, + false, + ); + + manager.run().await.unwrap(); + crate::foundation::time::sleep(Duration::from_millis(20)).await; + + assert!(listen_count.load(Ordering::Relaxed) >= 2); + assert_eq!(handler.accepted.lock().unwrap().as_slice(), &[3]); + assert!(events.events.lock().unwrap().iter().any(|event| matches!( + event, + CoreEvent::ListenerAddFailed { + will_retry: true, + .. + } + ))); + } + + #[tokio::test] + async fn required_listener_failure_does_not_leave_partial_tasks_running() { + let handler = Arc::new(MockHandler { + accepted: Mutex::new(Vec::new()), + }); + let events = Arc::new(Events::default()); + let first_accepts = Arc::new(Mutex::new(VecDeque::from([Ok::<_, anyhow::Error>(9)]))); + let mut manager = ListenerManager::new_with_options( + handler.clone(), + events.clone(), + ListenerManagerOptions { + accept_retry_delay: Duration::from_millis(1), + listen_retry_delay: Duration::from_millis(1), + max_listen_retries: 0, + }, + ); + + let first_accepts_clone = first_accepts.clone(); + manager.add_listener( + move || { + Box::new(MockListener { + url: "mock://first".parse().unwrap(), + listen_results: Arc::new(Mutex::new(VecDeque::from([Ok(())]))), + accepts: first_accepts_clone.clone(), + listen_count: Arc::new(AtomicUsize::new(0)), + drop_count: Arc::new(AtomicUsize::new(0)), + }) + }, + true, + ); + manager.add_listener( + move || { + Box::new(MockListener { + url: "mock://second".parse().unwrap(), + listen_results: Arc::new(Mutex::new(VecDeque::from([Err(anyhow::anyhow!( + "bind failed" + ))]))), + accepts: Arc::new(Mutex::new(VecDeque::new())), + listen_count: Arc::new(AtomicUsize::new(0)), + drop_count: Arc::new(AtomicUsize::new(0)), + }) + }, + true, + ); + + assert!(manager.run().await.is_err()); + crate::foundation::time::sleep(Duration::from_millis(20)).await; + assert!(handler.accepted.lock().unwrap().is_empty()); + assert!( + events + .events + .lock() + .unwrap() + .iter() + .all(|event| !matches!(event, CoreEvent::ListenerAdded { .. })) + ); + } + + #[tokio::test] + async fn manager_owned_closure_handler_handles_accepted_socket() { + let accepted = Arc::new(Mutex::new(Vec::new())); + let mut manager = ListenerManager::new_with_options( + Arc::new({ + let accepted = accepted.clone(); + move |value| { + let accepted = accepted.clone(); + async move { + accepted.lock().unwrap().push(value); + Ok(()) + } + } + }), + Arc::new(Events::default()), + ListenerManagerOptions { + accept_retry_delay: Duration::from_millis(1), + listen_retry_delay: Duration::from_millis(1), + max_listen_retries: 0, + }, + ); + + manager.add_listener( + move || { + Box::new(MockListener { + url: "mock://closure".parse().unwrap(), + listen_results: Arc::new(Mutex::new(VecDeque::from([Ok(())]))), + accepts: Arc::new(Mutex::new(VecDeque::from([Ok::<_, anyhow::Error>(11)]))), + listen_count: Arc::new(AtomicUsize::new(0)), + drop_count: Arc::new(AtomicUsize::new(0)), + }) + }, + true, + ); + + manager.run().await.unwrap(); + crate::foundation::time::sleep(Duration::from_millis(20)).await; + + assert_eq!(accepted.lock().unwrap().as_slice(), &[11]); + } + + #[derive(Debug)] + struct DropSignal(Option>); + + impl Drop for DropSignal { + fn drop(&mut self) { + if let Some(tx) = self.0.take() { + let _ = tx.send(()); + } + } + } + + #[derive(Debug)] + struct DropSignalListener { + url: Url, + listen_results: Arc>>>, + accepts: Arc>>>, + } + + #[async_trait] + impl SocketListener for DropSignalListener { + type Accepted = DropSignal; + + async fn listen(&mut self) -> anyhow::Result<()> { + self.listen_results + .lock() + .unwrap() + .pop_front() + .unwrap_or(Ok(())) + } + + async fn accept(&mut self) -> anyhow::Result { + let next_accept = self.accepts.lock().unwrap().pop_front(); + match next_accept { + Some(ret) => ret, + None => std::future::pending().await, + } + } + + fn local_url(&self) -> Url { + self.url.clone() + } + } + + #[derive(Debug)] + struct PendingHandler; + + #[async_trait] + impl AcceptedSocketHandler for PendingHandler { + async fn handle_accepted_socket(&self, accepted: DropSignal) -> anyhow::Result<()> { + let _accepted = accepted; + std::future::pending::<()>().await; + Ok(()) + } + } + + #[tokio::test] + async fn stop_joins_in_flight_handler_tasks_and_is_one_shot() { + let (drop_tx, drop_rx) = tokio::sync::oneshot::channel(); + let mut manager = ListenerManager::new_with_options( + Arc::new(PendingHandler), + Arc::new(Events::default()), + ListenerManagerOptions { + accept_retry_delay: Duration::from_millis(1), + listen_retry_delay: Duration::from_millis(1), + max_listen_retries: 0, + }, + ); + let accepts = Arc::new(Mutex::new(VecDeque::from([Ok::<_, anyhow::Error>( + DropSignal(Some(drop_tx)), + )]))); + let listener_accepts = accepts.clone(); + manager.add_listener( + move || { + Box::new(DropSignalListener { + url: "mock://drop".parse().unwrap(), + listen_results: Arc::new(Mutex::new(VecDeque::from([Ok(())]))), + accepts: listener_accepts.clone(), + }) + }, + true, + ); + + manager.run().await.unwrap(); + crate::foundation::time::sleep(Duration::from_millis(20)).await; + manager.stop().await; + + crate::foundation::time::timeout(Duration::from_secs(1), drop_rx) + .await + .unwrap() + .unwrap(); + assert!(manager.run().await.is_err()); + } + + #[tokio::test] + async fn stop_interrupts_required_listener_startup() { + let started = Arc::new(tokio::sync::Notify::new()); + let mut manager = ListenerManager::new_with_events( + Arc::new(MockHandler { + accepted: Mutex::new(Vec::new()), + }), + Arc::new(Events::default()), + ); + let listener_started = started.clone(); + manager.add_listener( + move || { + Box::new(BlockingListenListener { + started: listener_started.clone(), + }) + }, + true, + ); + let manager = Arc::new(manager); + let run_task = { + let manager = manager.clone(); + tokio::spawn(async move { manager.run().await }) + }; + + crate::foundation::time::timeout(Duration::from_secs(1), started.notified()) + .await + .unwrap(); + manager.stop().await; + + assert!(run_task.await.unwrap().is_err()); + assert!(manager.tasks.lock().await.is_empty()); + } + + #[tokio::test] + async fn listen_retry_exhaustion_keeps_in_flight_handler_tasks_until_manager_drop() { + let (drop_tx, mut drop_rx) = tokio::sync::oneshot::channel(); + let listen_results = Arc::new(Mutex::new(VecDeque::from([ + Ok(()), + Err(anyhow::anyhow!("bind failed")), + ]))); + let accepts = Arc::new(Mutex::new(VecDeque::from([ + Ok::<_, anyhow::Error>(DropSignal(Some(drop_tx))), + Err(anyhow::anyhow!("accept failed")), + ]))); + let events = Arc::new(Events::default()); + let mut manager = ListenerManager::new_with_options( + Arc::new(PendingHandler), + events.clone(), + ListenerManagerOptions { + accept_retry_delay: Duration::from_millis(1), + listen_retry_delay: Duration::from_millis(1), + max_listen_retries: 0, + }, + ); + let listener_results = listen_results.clone(); + let listener_accepts = accepts.clone(); + manager.add_listener( + move || { + Box::new(DropSignalListener { + url: "mock://retry-exhausted".parse().unwrap(), + listen_results: listener_results.clone(), + accepts: listener_accepts.clone(), + }) + }, + true, + ); + + manager.run().await.unwrap(); + crate::foundation::time::timeout(Duration::from_secs(1), async { + loop { + if events.events.lock().unwrap().iter().any(|event| { + matches!( + event, + CoreEvent::ListenerAddFailed { + will_retry: false, + .. + } + ) + }) { + break; + } + crate::foundation::time::sleep(Duration::from_millis(1)).await; + } + }) + .await + .unwrap(); + + assert!(matches!( + drop_rx.try_recv(), + Err(tokio::sync::oneshot::error::TryRecvError::Empty) + )); + + drop(manager); + crate::foundation::time::timeout(Duration::from_secs(1), drop_rx) + .await + .unwrap() + .unwrap(); + } +} diff --git a/easytier-core/src/listener/plan.rs b/easytier-core/src/listener/plan.rs new file mode 100644 index 00000000..6d044ed7 --- /dev/null +++ b/easytier-core/src/listener/plan.rs @@ -0,0 +1,475 @@ +use std::{ + collections::{BTreeMap, BTreeSet}, + net::IpAddr, + str::FromStr, +}; + +use percent_encoding::percent_decode_str; +use serde::{Deserialize, Serialize}; +use url::Url; + +use crate::{ + connectivity::{ + protocol::{ProtocolTransport, ServerProtocolUpgrader, protocol_transport}, + transport::UdpSessionMode, + }, + socket::{ + SocketContext, + tcp::{TcpBindOptions, TcpListenOptions}, + udp::{UdpBindOptions, UdpSessionAcceptKind, UdpSessionListenRequest}, + }, +}; + +use super::{ExternalListenerFactory, transport::TransportListenerConfig}; + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub(crate) enum ListenerKind { + Ring, + TcpStream, + UdpSession, + External, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub(crate) enum ListenerPlanSource { + Ring, + Configured, + Ipv6Shadow { original: Url }, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub(crate) struct PlannedListener { + pub url: Url, + pub kind: ListenerKind, + pub must_succeed: bool, + pub source: ListenerPlanSource, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub(crate) struct ListenerPlanFailure { + pub url: Url, + pub message: String, +} + +#[derive(Debug, Clone, Default, PartialEq, Eq)] +pub(crate) struct ListenerPlan { + pub listeners: Vec, + pub failures: Vec, +} + +#[derive(Debug, Clone, Default, PartialEq, Eq)] +pub(crate) struct ListenerSchemeRegistry { + schemes: BTreeMap, + no_ipv6_shadow: BTreeSet, +} + +impl ListenerSchemeRegistry { + pub(crate) fn new() -> Self { + Self::default() + } + + pub(crate) fn support(mut self, scheme: impl Into, kind: ListenerKind) -> Self { + self.schemes.insert(normalize_scheme(scheme), kind); + self + } + + pub(crate) fn disable_ipv6_shadow(mut self, scheme: impl Into) -> Self { + self.no_ipv6_shadow.insert(normalize_scheme(scheme)); + self + } + + pub(crate) fn classify(&self, url: &Url) -> Option { + self.schemes.get(url.scheme()).copied() + } + + fn allows_ipv6_shadow(&self, url: &Url) -> bool { + !self.no_ipv6_shadow.contains(url.scheme()) + } +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub(crate) struct ListenerPlanRequest { + pub self_id: uuid::Uuid, + pub listeners: Vec, + pub enable_ipv6: bool, +} + +/// Normalized listener inputs owned and planned by one core instance. +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +pub struct ListenerRuntimeConfig { + pub urls: Vec, + pub enable_ipv6: bool, + pub socket_context: SocketContext, +} + +impl ListenerRuntimeConfig { + pub fn new(urls: Vec, enable_ipv6: bool, socket_context: SocketContext) -> Self { + Self { + urls, + enable_ipv6, + socket_context, + } + } + + pub(crate) fn request(&self, self_id: uuid::Uuid) -> ListenerPlanRequest { + ListenerPlanRequest::new(self_id, self.urls.clone(), self.enable_ipv6) + } +} + +pub(crate) struct PreparedListenerPlan { + pub transports: Vec, + pub external: Vec<(PlannedListener, SocketContext)>, + pub failures: Vec, +} + +pub(crate) fn prepare_listener_plan( + config: Option<&ListenerRuntimeConfig>, + self_id: uuid::Uuid, + server_protocol: Option<&dyn ServerProtocolUpgrader>, + external_factory: Option<&dyn ExternalListenerFactory>, +) -> anyhow::Result +where + Accepted: Send + 'static, +{ + let Some(config) = config else { + return Ok(PreparedListenerPlan { + transports: Vec::new(), + external: Vec::new(), + failures: Vec::new(), + }); + }; + let mut schemes = ListenerSchemeRegistry::new() + .support("tcp", ListenerKind::TcpStream) + .support("udp", ListenerKind::UdpSession); + for (scheme, kind) in [ + ("ws", ListenerKind::TcpStream), + ("wss", ListenerKind::TcpStream), + ("wg", ListenerKind::UdpSession), + ("quic", ListenerKind::UdpSession), + ] { + if server_protocol.is_some_and(|protocol| protocol.supports_scheme(scheme)) { + schemes = schemes.support(scheme, kind); + } + } + schemes = schemes.disable_ipv6_shadow("quic"); + if server_protocol.is_some_and(|protocol| protocol.supports_scheme("faketcp")) + && external_factory.is_some_and(|factory| factory.supports_scheme("faketcp")) + { + schemes = schemes.support("faketcp", ListenerKind::External); + } + if external_factory.is_some_and(|factory| factory.supports_scheme("unix")) { + schemes = schemes.support("unix", ListenerKind::External); + } + schemes = schemes.disable_ipv6_shadow("faketcp"); + + let plan = plan_listeners(config.request(self_id), &schemes); + let mut transports = Vec::new(); + let mut external = Vec::new(); + for listener in plan.listeners { + let must_succeed = listener.must_succeed; + match listener.kind { + ListenerKind::Ring => transports.push(TransportListenerConfig::Ring { + url: listener.url, + must_succeed, + }), + ListenerKind::TcpStream => { + let max_pending_upgrades = server_protocol + .and_then(|protocol| protocol.max_pending_tcp_upgrades(listener.url.scheme())); + transports.push(TransportListenerConfig::Tcp { + url: listener.url, + options: unresolved_tcp_listener_options(config.socket_context.clone()), + max_pending_upgrades, + must_succeed, + }); + } + ListenerKind::UdpSession => { + let accept_kind = match protocol_transport(listener.url.scheme()) { + Some(ProtocolTransport::Udp(UdpSessionMode::EasyTierMux)) => { + UdpSessionAcceptKind::EasyTierMux + } + Some(ProtocolTransport::Udp(UdpSessionMode::Classified(protocol))) => { + UdpSessionAcceptKind::Classified(protocol) + } + _ => anyhow::bail!( + "listener scheme {} cannot produce a core UDP session", + listener.url.scheme() + ), + }; + let request = unresolved_udp_session_listen_request( + &listener.url, + config.socket_context.clone(), + ); + transports.push(TransportListenerConfig::Udp { + url: listener.url, + request, + accept_kind, + must_succeed, + }); + } + ListenerKind::External => external.push((listener, config.socket_context.clone())), + } + } + validate_listener_protocols(&transports, server_protocol.is_some())?; + + Ok(PreparedListenerPlan { + transports, + external, + failures: plan.failures, + }) +} + +fn validate_listener_protocols( + listeners: &[TransportListenerConfig], + has_server_protocol: bool, +) -> anyhow::Result<()> { + if has_server_protocol { + return Ok(()); + } + if let Some(listener) = listeners + .iter() + .find(|listener| !listener.supports_raw_handler()) + { + anyhow::bail!( + "listener {} requires a server protocol upgrader", + listener.url() + ); + } + Ok(()) +} + +impl ListenerPlanRequest { + pub(crate) fn new(self_id: uuid::Uuid, listeners: Vec, enable_ipv6: bool) -> Self { + Self { + self_id, + listeners, + enable_ipv6, + } + } +} + +pub(crate) fn plan_listeners( + request: ListenerPlanRequest, + registry: &ListenerSchemeRegistry, +) -> ListenerPlan { + let mut plan = ListenerPlan::default(); + plan.listeners.push(PlannedListener { + url: ring_listener_url(request.self_id), + kind: ListenerKind::Ring, + must_succeed: true, + source: ListenerPlanSource::Ring, + }); + + for url in request.listeners { + let Some(kind) = registry.classify(&url) else { + plan.failures.push(unsupported_listener(&url)); + continue; + }; + + plan.listeners.push(PlannedListener { + url: url.clone(), + kind, + must_succeed: true, + source: ListenerPlanSource::Configured, + }); + + if should_add_ipv6_shadow_listener(&url, request.enable_ipv6, registry) { + match ipv6_shadow_listener(&url) { + Ok(ipv6_url) => plan.listeners.push(PlannedListener { + url: ipv6_url, + kind, + must_succeed: false, + source: ListenerPlanSource::Ipv6Shadow { original: url }, + }), + Err(message) => plan.failures.push(ListenerPlanFailure { url, message }), + } + } + } + + plan +} + +pub(crate) fn ring_listener_url(self_id: uuid::Uuid) -> Url { + format!("ring://{self_id}") + .parse() + .expect("ring listener url should be valid") +} + +pub(crate) fn listener_default_port(scheme: &str) -> Option { + crate::connectivity::protocol::protocol_default_port(scheme) +} + +pub(crate) fn listener_url_bind_device(url: &Url) -> Option { + url.path().strip_prefix('/').and_then(|path| { + if path.is_empty() { + None + } else { + Some(String::from_utf8(percent_decode_str(path).collect()).unwrap()) + } + }) +} + +pub(crate) fn unresolved_udp_session_listen_request( + url: &Url, + context: SocketContext, +) -> UdpSessionListenRequest { + UdpSessionListenRequest::new( + UdpBindOptions::port_bound_listener("0.0.0.0:0".parse().unwrap()) + .with_local_addr(None) + .with_context(context) + .with_bind_device(listener_url_bind_device(url)), + ) +} + +pub(crate) fn unresolved_tcp_listener_options(context: SocketContext) -> TcpListenOptions { + TcpListenOptions::direct_connect("0.0.0.0:0".parse().unwrap()) + .with_bind(TcpBindOptions::default().with_context(context)) +} + +pub(crate) fn is_url_host_ipv6(url: &Url) -> bool { + url.host_str().is_some_and(|h| h.contains(':')) +} + +pub(crate) fn is_url_host_unspecified(url: &Url) -> bool { + if let Ok(ip) = IpAddr::from_str(url.host_str().unwrap_or_default()) { + ip.is_unspecified() + } else { + false + } +} + +fn should_add_ipv6_shadow_listener( + url: &Url, + enable_ipv6: bool, + registry: &ListenerSchemeRegistry, +) -> bool { + enable_ipv6 + && registry.allows_ipv6_shadow(url) + && !is_url_host_ipv6(url) + && is_url_host_unspecified(url) +} + +fn ipv6_shadow_listener(url: &Url) -> Result { + let mut ipv6_url = url.clone(); + ipv6_url + .set_host(Some("[::]")) + .map_err(|_| format!("failed to set ipv6 host for listener: {url}"))?; + Ok(ipv6_url) +} + +fn unsupported_listener(url: &Url) -> ListenerPlanFailure { + ListenerPlanFailure { + url: url.clone(), + message: format!("failed to get listener by url: {url}, maybe not supported"), + } +} + +fn normalize_scheme(scheme: impl Into) -> String { + scheme.into().to_ascii_lowercase() +} + +#[cfg(test)] +mod tests { + use super::*; + + fn registry() -> ListenerSchemeRegistry { + ListenerSchemeRegistry::new() + .support("tcp", ListenerKind::TcpStream) + .support("udp", ListenerKind::UdpSession) + .support("quic", ListenerKind::External) + .support("faketcp", ListenerKind::External) + .disable_ipv6_shadow("quic") + .disable_ipv6_shadow("faketcp") + } + + #[test] + fn listener_plan_adds_ring_and_configured_listener() { + let self_id = uuid::Uuid::parse_str("00000000-0000-0000-0000-000000000001").unwrap(); + let plan = plan_listeners( + ListenerPlanRequest::new( + self_id, + vec!["udp://127.0.0.1:11010".parse().unwrap()], + false, + ), + ®istry(), + ); + + assert_eq!(plan.failures, Vec::new()); + assert_eq!( + plan.listeners + .iter() + .map(|entry| (&entry.url, entry.kind, entry.must_succeed)) + .collect::>(), + vec![ + ( + &"ring://00000000-0000-0000-0000-000000000001" + .parse() + .unwrap(), + ListenerKind::Ring, + true + ), + ( + &"udp://127.0.0.1:11010".parse().unwrap(), + ListenerKind::UdpSession, + true + ), + ] + ); + } + + #[test] + fn listener_plan_adds_ipv6_shadow_for_unspecified_ip_listener() { + let plan = plan_listeners( + ListenerPlanRequest::new( + uuid::Uuid::new_v4(), + vec!["tcp://0.0.0.0:11010".parse().unwrap()], + true, + ), + ®istry(), + ); + + assert_eq!(plan.failures, Vec::new()); + assert_eq!(plan.listeners.len(), 3); + assert_eq!(plan.listeners[2].url, "tcp://[::]:11010".parse().unwrap()); + assert!(!plan.listeners[2].must_succeed); + assert!(matches!( + plan.listeners[2].source, + ListenerPlanSource::Ipv6Shadow { .. } + )); + } + + #[test] + fn listener_plan_skips_ipv6_shadow_for_excluded_schemes() { + for url in ["quic://0.0.0.0:11012", "faketcp://0.0.0.0:11013"] { + let plan = plan_listeners( + ListenerPlanRequest::new(uuid::Uuid::new_v4(), vec![url.parse().unwrap()], true), + ®istry(), + ); + + assert_eq!(plan.failures, Vec::new()); + assert_eq!(plan.listeners.len(), 2); + } + } + + #[test] + fn listener_plan_reports_unsupported_scheme() { + let url = "http://0.0.0.0:8080".parse().unwrap(); + let plan = plan_listeners( + ListenerPlanRequest::new(uuid::Uuid::new_v4(), vec![url], true), + ®istry(), + ); + + assert_eq!(plan.listeners.len(), 1); + assert_eq!(plan.failures.len(), 1); + assert_eq!( + plan.failures[0].message, + "failed to get listener by url: http://0.0.0.0:8080/, maybe not supported" + ); + } + + #[test] + fn listener_url_bind_device_decodes_url_path() { + let url = "udp://0.0.0.0:11010/eth%2Btest".parse().unwrap(); + + assert_eq!(listener_url_bind_device(&url), Some("eth+test".to_owned())); + } +} diff --git a/easytier-core/src/listener/transport.rs b/easytier-core/src/listener/transport.rs new file mode 100644 index 00000000..16c5d706 --- /dev/null +++ b/easytier-core/src/listener/transport.rs @@ -0,0 +1,1581 @@ +use std::{fmt, marker::PhantomData, sync::Arc}; + +use async_trait::async_trait; +use rand::seq::SliceRandom as _; +use tokio::sync::{OwnedSemaphorePermit, Semaphore}; +use url::Url; + +use crate::{ + connectivity::{ + manual::resolve_url_addrs, + protocol::{ + ServerProtocolAdmissionController, ServerProtocolUpgrade, ServerProtocolUpgrader, raw, + }, + }, + events::{CoreEvent, CoreEventSink}, + host::dns::DnsResolver, + listener::{ + AcceptedSocketHandler, ListenerFactory, ListenerManager, RunningListenerRegistry, + plan::ListenerPlanFailure, + }, + socket::{ + IpVersion, ListenerConnectionCounter, SocketContext, SocketListener, + tcp::{TcpListenOptions, TcpSocketListener, VirtualTcpListener, VirtualTcpListenerFactory}, + udp::{ + UdpSession, UdpSessionAcceptKind, UdpSessionListenRequest, UdpSessionSocket, + UdpSessionSocketListener, VirtualUdpSocketFactory, + }, + }, + tunnel::{Tunnel, ring::RingTunnelRegistry}, +}; + +pub type HostAcceptedTcpSocket = + <::Listener as VirtualTcpListener>::Socket; + +/// The transport boundary of a core listener. +/// +/// Socket creation and binding belong to the host factory. Protocol handling +/// consumes this value and may turn it into one or more EasyTier tunnels. +#[allow(clippy::large_enum_variant)] +pub enum AcceptedTransport { + Tunnel { + tunnel: Box, + local_url: Url, + }, + Tcp { + socket: TcpSocket, + local_url: Url, + upgrade_permit: Option, + }, + Udp { + session: UdpSession, + local_url: Url, + admission: Option, + }, + ByteStream { + socket: TcpSocket, + local_url: Url, + remote_url: Option, + }, +} + +impl AcceptedTransport { + pub fn local_url(&self) -> &Url { + match self { + Self::Tunnel { local_url, .. } + | Self::Tcp { local_url, .. } + | Self::Udp { local_url, .. } + | Self::ByteStream { local_url, .. } => local_url, + } + } +} + +#[async_trait] +pub trait AcceptedTunnelHandler: Send + Sync + 'static { + async fn handle_tunnel(&self, tunnel: Box) -> anyhow::Result<()>; +} + +/// Upgrades accepted sockets or sessions in core before peer admission. +pub struct ProtocolAcceptedTransportHandler { + tunnel_handler: Arc, + protocol: Arc>, +} + +impl ProtocolAcceptedTransportHandler { + pub fn new( + tunnel_handler: &Arc, + protocol: Arc>, + ) -> Self { + Self { + tunnel_handler: tunnel_handler.clone(), + protocol, + } + } +} + +impl ProtocolAcceptedTransportHandler +where + H: AcceptedTunnelHandler, +{ + async fn handle_tunnel(&self, tunnel: Box) -> anyhow::Result<()> { + self.tunnel_handler.handle_tunnel(tunnel).await + } + + async fn handle_upgrade(&self, upgrade: ServerProtocolUpgrade) -> anyhow::Result<()> { + match upgrade { + ServerProtocolUpgrade::Tunnel(tunnel) => self.handle_tunnel(tunnel).await, + ServerProtocolUpgrade::Acceptor(mut acceptor) => loop { + let tunnel = acceptor.accept().await?; + let _ = self.handle_tunnel(tunnel).await; + }, + } + } +} + +#[async_trait] +impl AcceptedSocketHandler> + for ProtocolAcceptedTransportHandler +where + TcpSocket: crate::socket::tcp::VirtualTcpSocket, + H: AcceptedTunnelHandler, +{ + async fn handle_accepted_socket( + &self, + accepted: AcceptedTransport, + ) -> anyhow::Result<()> { + let (upgrade, tcp_upgrade_permit) = match accepted { + AcceptedTransport::Tunnel { tunnel, .. } => { + return self.handle_tunnel(tunnel).await; + } + AcceptedTransport::Tcp { + socket, + local_url, + upgrade_permit, + } => ( + self.protocol.upgrade_tcp(socket, local_url).await?, + upgrade_permit, + ), + AcceptedTransport::Udp { + session, + local_url, + admission, + } => ( + self.protocol + .upgrade_udp(session, local_url, admission) + .await?, + None, + ), + AcceptedTransport::ByteStream { + socket, + local_url, + remote_url, + } if local_url.scheme() == "unix" => { + let tunnel = raw::upgrade_accepted_byte_stream(socket, local_url, remote_url)?; + return self.handle_tunnel(tunnel).await; + } + AcceptedTransport::ByteStream { + socket, + local_url, + remote_url, + } => ( + self.protocol + .upgrade_byte_stream(socket, local_url, remote_url) + .await?, + None, + ), + }; + drop(tcp_upgrade_permit); + self.handle_upgrade(upgrade).await + } +} + +#[derive(Debug, Clone)] +pub(crate) enum TransportListenerConfig { + Ring { + url: Url, + must_succeed: bool, + }, + Tcp { + url: Url, + options: TcpListenOptions, + max_pending_upgrades: Option, + must_succeed: bool, + }, + Udp { + url: Url, + request: UdpSessionListenRequest, + accept_kind: UdpSessionAcceptKind, + must_succeed: bool, + }, +} + +impl TransportListenerConfig { + pub(crate) fn must_succeed(&self) -> bool { + match self { + Self::Ring { must_succeed, .. } + | Self::Tcp { must_succeed, .. } + | Self::Udp { must_succeed, .. } => *must_succeed, + } + } + + pub(crate) fn url(&self) -> &Url { + match self { + Self::Ring { url, .. } | Self::Tcp { url, .. } | Self::Udp { url, .. } => url, + } + } + + pub(crate) fn supports_raw_handler(&self) -> bool { + matches!(self, Self::Ring { url, .. } if url.scheme() == "ring") + || matches!(self, Self::Tcp { url, .. } if url.scheme() == "tcp") + || matches!( + self, + Self::Udp { + url, + accept_kind: UdpSessionAcceptKind::EasyTierMux, + .. + } if url.scheme() == "udp" + ) + } +} + +struct RingTransportListener { + url: Url, + registry: Arc, + inner: Option, + tcp_socket: PhantomData TcpSocket>, +} + +impl RingTransportListener { + fn new(url: Url, registry: Arc) -> Self { + Self { + url, + registry, + inner: None, + tcp_socket: PhantomData, + } + } +} + +impl fmt::Debug for RingTransportListener { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter + .debug_struct("RingTransportListener") + .field("url", &self.url) + .field("listening", &self.inner.is_some()) + .finish() + } +} + +#[async_trait] +impl SocketListener for RingTransportListener +where + TcpSocket: Send + 'static, +{ + type Accepted = AcceptedTransport; + + async fn listen(&mut self) -> anyhow::Result<()> { + if self.inner.is_none() { + if self.url.scheme() != "ring" { + anyhow::bail!("Ring listener requires ring URL: {}", self.url); + } + let local_id = self + .url + .host_str() + .ok_or_else(|| anyhow::anyhow!("ring listener URL has no peer id: {}", self.url))? + .parse()?; + self.inner = Some(self.registry.bind(local_id)?); + } + Ok(()) + } + + async fn accept(&mut self) -> anyhow::Result { + let accepted = self + .inner + .as_mut() + .ok_or_else(|| anyhow::anyhow!("Ring transport listener is not started"))? + .accept() + .await?; + Ok(AcceptedTransport::Tunnel { + tunnel: accepted.into_tunnel(), + local_url: self.url.clone(), + }) + } + + fn local_url(&self) -> Url { + self.url.clone() + } +} + +struct TcpTransportListener +where + H: VirtualTcpListenerFactory, +{ + url: Url, + options: TcpListenOptions, + host: Arc, + dns: Arc, + upgrade_slots: Option>, + inner: Option>, +} + +impl TcpTransportListener +where + H: VirtualTcpListenerFactory, +{ + fn new( + url: Url, + options: TcpListenOptions, + max_pending_upgrades: Option, + host: Arc, + dns: Arc, + ) -> Self { + let upgrade_slots = max_pending_upgrades.map(|limit| Arc::new(Semaphore::new(limit.get()))); + Self { + url, + options, + host, + dns, + upgrade_slots, + inner: None, + } + } + + fn inner(&mut self) -> anyhow::Result<&mut TcpSocketListener> { + self.inner + .as_mut() + .ok_or_else(|| anyhow::anyhow!("TCP transport listener is not started")) + } +} + +impl fmt::Debug for TcpTransportListener +where + H: VirtualTcpListenerFactory, +{ + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter + .debug_struct("TcpTransportListener") + .field("url", &self.url) + .field("listening", &self.inner.is_some()) + .finish() + } +} + +#[async_trait] +impl SocketListener for TcpTransportListener +where + H: VirtualTcpListenerFactory, +{ + type Accepted = AcceptedTransport>; + + async fn listen(&mut self) -> anyhow::Result<()> { + if self.inner.is_some() { + return Ok(()); + } + let mut options = self.options.clone(); + if options.bind.local_addr.is_none() { + options.bind.local_addr = Some( + resolve_listener_addr(&self.url, options.bind.context.clone(), self.dns.as_ref()) + .await?, + ); + } + if let Some(local_addr) = options.bind.local_addr { + options.bind.context.ip_version = if local_addr.is_ipv4() { + IpVersion::V4 + } else { + IpVersion::V6 + }; + options.bind.only_v6 = local_addr.is_ipv6(); + } + let mut inner = + TcpSocketListener::new_with_options(self.url.clone(), options, self.host.clone()); + inner.listen().await?; + self.inner = Some(inner); + Ok(()) + } + + async fn accept(&mut self) -> anyhow::Result { + let upgrade_permit = match &self.upgrade_slots { + Some(slots) => Some(slots.clone().acquire_owned().await?), + None => None, + }; + let local_url = self.local_url(); + let socket = self.inner()?.accept().await?; + Ok(AcceptedTransport::Tcp { + socket, + local_url, + upgrade_permit, + }) + } + + fn local_url(&self) -> Url { + self.inner + .as_ref() + .map(SocketListener::local_url) + .unwrap_or_else(|| self.url.clone()) + } + + fn connection_counter(&self) -> Arc { + self.inner + .as_ref() + .map(SocketListener::connection_counter) + .unwrap_or_else(|| Arc::new(EmptyTransportConnectionCounter)) + } +} + +struct UdpTransportListener +where + H: VirtualUdpSocketFactory, +{ + url: Url, + request: UdpSessionListenRequest, + accept_kind: UdpSessionAcceptKind, + host: Arc, + dns: Arc, + inner: Option>, + protocol_admission: Option, + tcp_socket: PhantomData TcpSocket>, +} + +impl UdpTransportListener +where + H: VirtualUdpSocketFactory, +{ + fn new( + url: Url, + request: UdpSessionListenRequest, + accept_kind: UdpSessionAcceptKind, + host: Arc, + dns: Arc, + ) -> Self { + let protocol_admission = matches!( + accept_kind, + UdpSessionAcceptKind::Classified(crate::socket::udp::UdpSessionProtocol::Quic) + ) + .then(ServerProtocolAdmissionController::quic); + Self { + url, + request, + accept_kind, + host, + dns, + inner: None, + protocol_admission, + tcp_socket: PhantomData, + } + } + + fn inner(&mut self) -> anyhow::Result<&mut UdpSessionSocketListener> { + self.inner + .as_mut() + .ok_or_else(|| anyhow::anyhow!("UDP transport listener is not started")) + } +} + +impl fmt::Debug for UdpTransportListener +where + H: VirtualUdpSocketFactory, +{ + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter + .debug_struct("UdpTransportListener") + .field("url", &self.url) + .field("accept_kind", &self.accept_kind) + .field("listening", &self.inner.is_some()) + .finish() + } +} + +#[async_trait] +impl SocketListener for UdpTransportListener +where + H: VirtualUdpSocketFactory, + TcpSocket: Send + 'static, +{ + type Accepted = AcceptedTransport; + + async fn listen(&mut self) -> anyhow::Result<()> { + if self.inner.is_some() { + return Ok(()); + } + let mut request = self.request.clone(); + if request.bind.local_addr.is_none() { + request.bind.local_addr = Some( + resolve_listener_addr(&self.url, request.bind.context.clone(), self.dns.as_ref()) + .await?, + ); + } + if let Some(local_addr) = request.bind.local_addr { + request.bind.context.ip_version = if local_addr.is_ipv4() { + IpVersion::V4 + } else { + IpVersion::V6 + }; + request.bind.only_v6 = local_addr.is_ipv6(); + } + let mut inner = UdpSessionSocketListener::new_with_request( + self.url.clone(), + request, + self.accept_kind, + self.host.clone(), + ); + inner.listen().await?; + self.inner = Some(inner); + Ok(()) + } + + async fn accept(&mut self) -> anyhow::Result { + loop { + let session = self.inner()?.accept().await?; + let admission = match &self.protocol_admission { + Some(controller) => match controller.try_admit() { + Some(admission) => Some(admission), + None => { + tracing::debug!( + peer_addr = ?session.peer_addr(), + "drop UDP session after protocol admission limit" + ); + continue; + } + }, + None => None, + }; + return Ok(AcceptedTransport::Udp { + session, + local_url: self.local_url(), + admission, + }); + } + } + + fn local_url(&self) -> Url { + self.inner + .as_ref() + .map(SocketListener::local_url) + .unwrap_or_else(|| self.url.clone()) + } + + fn connection_counter(&self) -> Arc { + self.inner + .as_ref() + .map(SocketListener::connection_counter) + .unwrap_or_else(|| Arc::new(EmptyTransportConnectionCounter)) + } +} + +#[derive(Debug)] +struct EmptyTransportConnectionCounter; + +impl ListenerConnectionCounter for EmptyTransportConnectionCounter { + fn get(&self) -> Option { + None + } +} + +async fn resolve_listener_addr( + url: &Url, + context: SocketContext, + dns: &dyn DnsResolver, +) -> anyhow::Result { + let default_port = super::plan::listener_default_port(url.scheme()) + .ok_or_else(|| anyhow::anyhow!("listener has no default port: {url}"))?; + resolve_url_addrs( + url, + default_port, + context.with_ip_version(IpVersion::Both), + dns, + ) + .await? + .choose(&mut rand::thread_rng()) + .copied() + .ok_or_else(|| anyhow::anyhow!("listener has no resolved address: {url}")) +} + +type HostTransportListenerManager = ListenerManager< + AcceptedTransport>, + dyn AcceptedSocketHandler>>, +>; + +/// Owns all listeners planned by core, including host-backed external sockets. +pub(crate) struct CoreListenerRuntime +where + H: VirtualTcpListenerFactory + VirtualUdpSocketFactory, +{ + manager: HostTransportListenerManager, + plan_failures: Vec, + events: Arc, +} + +impl CoreListenerRuntime +where + H: VirtualTcpListenerFactory + VirtualUdpSocketFactory, +{ + #[allow(clippy::too_many_arguments)] + pub(crate) fn new_with_events( + host: Arc, + dns: Arc, + ring_registry: Arc, + configs: Vec, + external_factories: Vec>>>, + plan_failures: Vec, + handler: Arc>>>, + events: Arc, + registry: Arc, + ) -> Self { + let mut manager = ListenerManager::new_with_registry( + handler, + events.clone(), + registry, + crate::listener::ListenerManagerOptions::default(), + ); + + for config in configs { + let must_succeed = config.must_succeed(); + match config { + TransportListenerConfig::Ring { url, .. } => { + let ring_registry = ring_registry.clone(); + manager.add_listener( + move || { + Box::new(RingTransportListener::new( + url.clone(), + ring_registry.clone(), + )) + }, + must_succeed, + ); + } + TransportListenerConfig::Tcp { + url, + options, + max_pending_upgrades, + .. + } => { + let host = host.clone(); + let dns = dns.clone(); + manager.add_listener( + move || { + Box::new(TcpTransportListener::new( + url.clone(), + options.clone(), + max_pending_upgrades, + host.clone(), + dns.clone(), + )) + }, + must_succeed, + ); + } + TransportListenerConfig::Udp { + url, + request, + accept_kind, + .. + } => { + let host = host.clone(); + let dns = dns.clone(); + manager.add_listener( + move || { + Box::new(UdpTransportListener::new( + url.clone(), + request.clone(), + accept_kind, + host.clone(), + dns.clone(), + )) + }, + must_succeed, + ); + } + } + } + + for factory in external_factories { + manager.add_factory(factory); + } + + Self { + manager, + plan_failures, + events, + } + } + + pub async fn start(&self) -> anyhow::Result<()> { + for failure in &self.plan_failures { + self.events.emit(CoreEvent::ListenerPlanFailed { + url: failure.url.clone(), + error: failure.message.clone(), + }); + } + self.manager.run().await + } + + pub async fn stop(&self) { + self.manager.stop().await; + } +} + +#[cfg(test)] +mod tests { + use std::{ + collections::VecDeque, + io, + net::SocketAddr, + pin::Pin, + sync::{ + Arc, Mutex as StdMutex, + atomic::{AtomicU16, AtomicUsize, Ordering}, + }, + task::{Context, Poll}, + time::Duration, + }; + + use futures::{SinkExt as _, StreamExt as _}; + use tokio::{ + io::{AsyncRead, AsyncWrite, DuplexStream, ReadBuf}, + sync::{Mutex, Notify, Semaphore, mpsc}, + }; + + use super::*; + use crate::{ + connectivity::protocol::{CoreServerProtocolUpgrader, ServerTunnelAcceptor}, + host::dns::{DnsQuery, DnsResolver}, + packet::ZCPacket, + socket::{ + tcp::{VirtualTcpListener, VirtualTcpSocket}, + udp::{ + UdpBindOptions, UdpSessionKind, UdpSessionProtocol, UdpSessionSocket, + VirtualUdpSocket, new_syn_packet, + }, + }, + }; + + struct MockDns; + + #[async_trait] + impl DnsResolver for MockDns { + async fn resolve(&self, _query: DnsQuery) -> anyhow::Result> { + Ok(vec!["127.0.0.1".parse().unwrap()]) + } + } + + struct MockTcpSocket { + stream: DuplexStream, + local_addr: SocketAddr, + peer_addr: SocketAddr, + } + + impl AsyncRead for MockTcpSocket { + fn poll_read( + mut self: Pin<&mut Self>, + cx: &mut Context<'_>, + buf: &mut ReadBuf<'_>, + ) -> Poll> { + Pin::new(&mut self.stream).poll_read(cx, buf) + } + } + + impl AsyncWrite for MockTcpSocket { + fn poll_write( + mut self: Pin<&mut Self>, + cx: &mut Context<'_>, + buf: &[u8], + ) -> Poll> { + Pin::new(&mut self.stream).poll_write(cx, buf) + } + + fn poll_flush(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + Pin::new(&mut self.stream).poll_flush(cx) + } + + fn poll_shutdown(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + Pin::new(&mut self.stream).poll_shutdown(cx) + } + } + + impl VirtualTcpSocket for MockTcpSocket { + fn local_addr(&self) -> io::Result { + Ok(self.local_addr) + } + + fn peer_addr(&self) -> io::Result { + Ok(self.peer_addr) + } + } + + struct MockTcpListener { + local_addr: SocketAddr, + accepted_tx: mpsc::UnboundedSender, + accepted_rx: Mutex>, + } + + impl MockTcpListener { + fn new(local_addr: SocketAddr) -> Self { + let (accepted_tx, accepted_rx) = mpsc::unbounded_channel(); + Self { + local_addr, + accepted_tx, + accepted_rx: Mutex::new(accepted_rx), + } + } + + fn accept_from(&self, peer_addr: SocketAddr) { + let (stream, _remote) = tokio::io::duplex(64); + self.accepted_tx + .send(MockTcpSocket { + stream, + local_addr: self.local_addr, + peer_addr, + }) + .unwrap(); + } + } + + #[async_trait] + impl VirtualTcpListener for MockTcpListener { + type Socket = MockTcpSocket; + + fn local_addr(&self) -> io::Result { + Ok(self.local_addr) + } + + async fn accept(&self) -> io::Result<(Self::Socket, SocketAddr)> { + let socket = + self.accepted_rx.lock().await.recv().await.ok_or_else(|| { + io::Error::new(io::ErrorKind::BrokenPipe, "accept queue closed") + })?; + let peer_addr = socket.peer_addr; + Ok((socket, peer_addr)) + } + } + + struct MockUdpSocket { + local_addr: SocketAddr, + incoming: StdMutex, SocketAddr)>>, + incoming_notify: Notify, + } + + impl MockUdpSocket { + fn new(local_addr: SocketAddr) -> Self { + Self { + local_addr, + incoming: StdMutex::new(VecDeque::new()), + incoming_notify: Notify::new(), + } + } + + fn receive_from(&self, data: Vec, peer_addr: SocketAddr) { + self.incoming.lock().unwrap().push_back((data, peer_addr)); + self.incoming_notify.notify_one(); + } + } + + #[async_trait] + impl VirtualUdpSocket for MockUdpSocket { + fn local_addr(&self) -> io::Result { + Ok(self.local_addr) + } + + async fn send_to(&self, data: &[u8], _addr: SocketAddr) -> io::Result { + Ok(data.len()) + } + + async fn recv_from(&self, buf: &mut [u8]) -> io::Result<(usize, SocketAddr)> { + loop { + if let Some((data, peer_addr)) = self.incoming.lock().unwrap().pop_front() { + let len = data.len().min(buf.len()); + buf[..len].copy_from_slice(&data[..len]); + return Ok((len, peer_addr)); + } + self.incoming_notify.notified().await; + } + } + } + + struct MockHost { + next_tcp_port: AtomicU16, + next_udp_port: AtomicU16, + tcp_bind_options: StdMutex>, + udp_bind_options: StdMutex>, + tcp_listeners: StdMutex>>, + udp_sockets: StdMutex>>, + } + + impl MockHost { + fn new() -> Self { + Self { + next_tcp_port: AtomicU16::new(21000), + next_udp_port: AtomicU16::new(22000), + tcp_bind_options: StdMutex::new(Vec::new()), + udp_bind_options: StdMutex::new(Vec::new()), + tcp_listeners: StdMutex::new(Vec::new()), + udp_sockets: StdMutex::new(Vec::new()), + } + } + + fn tcp_listener(&self, index: usize) -> Arc { + self.tcp_listeners.lock().unwrap()[index].clone() + } + + fn udp_socket(&self, index: usize) -> Arc { + self.udp_sockets.lock().unwrap()[index].clone() + } + + fn tcp_bind_options(&self, index: usize) -> TcpListenOptions { + self.tcp_bind_options.lock().unwrap()[index].clone() + } + + fn udp_bind_options(&self, index: usize) -> UdpBindOptions { + self.udp_bind_options.lock().unwrap()[index].clone() + } + } + + fn assigned_addr(requested: Option, next_port: &AtomicU16) -> SocketAddr { + let mut addr = requested.unwrap_or_else(|| "127.0.0.1:0".parse().unwrap()); + if addr.port() == 0 { + addr.set_port(next_port.fetch_add(1, Ordering::Relaxed)); + } + addr + } + + #[async_trait] + impl VirtualTcpListenerFactory for MockHost { + type Listener = MockTcpListener; + + async fn bind_tcp(&self, options: TcpListenOptions) -> anyhow::Result> { + let local_addr = assigned_addr(options.bind.local_addr, &self.next_tcp_port); + self.tcp_bind_options.lock().unwrap().push(options); + let listener = Arc::new(MockTcpListener::new(local_addr)); + self.tcp_listeners.lock().unwrap().push(listener.clone()); + Ok(listener) + } + } + + #[async_trait] + impl VirtualUdpSocketFactory for MockHost { + type Socket = MockUdpSocket; + + async fn bind_udp(&self, options: UdpBindOptions) -> anyhow::Result> { + let local_addr = assigned_addr(options.local_addr, &self.next_udp_port); + self.udp_bind_options.lock().unwrap().push(options); + let socket = Arc::new(MockUdpSocket::new(local_addr)); + self.udp_sockets.lock().unwrap().push(socket.clone()); + Ok(socket) + } + } + + #[derive(Debug, PartialEq, Eq)] + enum AcceptedEvent { + Tcp { port: u16 }, + Udp { port: u16, kind: UdpSessionKind }, + } + + struct ActiveHandlerGuard(Arc); + + impl ActiveHandlerGuard { + fn new(active: Arc) -> Self { + active.fetch_add(1, Ordering::Relaxed); + Self(active) + } + } + + impl Drop for ActiveHandlerGuard { + fn drop(&mut self) { + self.0.fetch_sub(1, Ordering::Relaxed); + } + } + + struct RecordingHandler { + events: mpsc::UnboundedSender, + blocked: Arc, + active: Arc, + } + + #[derive(Default)] + struct RecordingListenerEvents { + events: StdMutex>, + } + + impl CoreEventSink for RecordingListenerEvents { + fn emit(&self, event: CoreEvent) { + self.events.lock().unwrap().push(event); + } + } + + struct QueueTunnelAcceptor { + tunnels: VecDeque>, + } + + #[async_trait] + impl ServerTunnelAcceptor for QueueTunnelAcceptor { + async fn accept(&mut self) -> anyhow::Result> { + self.tunnels + .pop_front() + .ok_or_else(|| anyhow::anyhow!("server tunnel acceptor finished")) + } + } + + struct RecordingServerProtocolUpgrader { + tcp_calls: AtomicUsize, + udp_calls: AtomicUsize, + byte_stream_calls: AtomicUsize, + } + + impl RecordingServerProtocolUpgrader { + fn new() -> Self { + Self { + tcp_calls: AtomicUsize::new(0), + udp_calls: AtomicUsize::new(0), + byte_stream_calls: AtomicUsize::new(0), + } + } + } + + #[async_trait] + impl ServerProtocolUpgrader for RecordingServerProtocolUpgrader { + fn supports_scheme(&self, scheme: &str) -> bool { + matches!(scheme, "tcp" | "udp") + } + + async fn upgrade_tcp( + &self, + socket: MockTcpSocket, + _local_url: Url, + ) -> anyhow::Result { + self.tcp_calls.fetch_add(1, Ordering::Relaxed); + let first = raw::tests::upgrade_accepted_tcp(socket)?; + let (stream, _remote) = tokio::io::duplex(64); + let second = raw::tests::upgrade_accepted_tcp(MockTcpSocket { + stream, + local_addr: "127.0.0.1:21001".parse().unwrap(), + peer_addr: "127.0.0.1:31001".parse().unwrap(), + })?; + Ok(ServerProtocolUpgrade::Acceptor(Box::new( + QueueTunnelAcceptor { + tunnels: VecDeque::from([first, second]), + }, + ))) + } + + async fn upgrade_udp( + &self, + session: UdpSession, + _local_url: Url, + _admission: Option, + ) -> anyhow::Result { + self.udp_calls.fetch_add(1, Ordering::Relaxed); + Ok(ServerProtocolUpgrade::Tunnel( + raw::tests::upgrade_accepted_udp(session)?, + )) + } + + async fn upgrade_byte_stream( + &self, + socket: MockTcpSocket, + local_url: Url, + remote_url: Option, + ) -> anyhow::Result { + self.byte_stream_calls.fetch_add(1, Ordering::Relaxed); + Ok(ServerProtocolUpgrade::Tunnel( + raw::upgrade_accepted_byte_stream(socket, local_url, remote_url)?, + )) + } + } + + struct RecordingTunnelHandler { + calls: AtomicUsize, + } + + struct BlockingTunnelHandler { + entered: Notify, + release: Notify, + } + + #[async_trait] + impl AcceptedTunnelHandler for RecordingTunnelHandler { + async fn handle_tunnel( + &self, + _tunnel: Box, + ) -> anyhow::Result<()> { + let call = self.calls.fetch_add(1, Ordering::Relaxed); + if call == 0 { + anyhow::bail!("first admission rejected"); + } + Ok(()) + } + } + + #[async_trait] + impl AcceptedTunnelHandler for BlockingTunnelHandler { + async fn handle_tunnel( + &self, + _tunnel: Box, + ) -> anyhow::Result<()> { + self.entered.notify_one(); + self.release.notified().await; + Ok(()) + } + } + + #[async_trait] + impl AcceptedSocketHandler> for RecordingHandler { + async fn handle_accepted_socket( + &self, + accepted: AcceptedTransport, + ) -> anyhow::Result<()> { + let _active = ActiveHandlerGuard::new(self.active.clone()); + match accepted { + AcceptedTransport::Tunnel { tunnel, .. } => { + drop(tunnel); + } + AcceptedTransport::Tcp { + socket, local_url, .. + } => { + self.events.send(AcceptedEvent::Tcp { + port: local_url.port().unwrap(), + })?; + self.blocked.notified().await; + drop(socket); + } + AcceptedTransport::Udp { + session, local_url, .. + } => { + self.events.send(AcceptedEvent::Udp { + port: local_url.port().unwrap(), + kind: session.kind(), + })?; + self.blocked.notified().await; + drop(session); + } + AcceptedTransport::ByteStream { + socket, local_url, .. + } => { + self.events.send(AcceptedEvent::Tcp { + port: local_url.port().unwrap_or_default(), + })?; + self.blocked.notified().await; + drop(socket); + } + } + Ok(()) + } + } + + fn wireguard_packet() -> Vec { + let mut packet = vec![0; 32]; + packet[..4].copy_from_slice(&4u32.to_le_bytes()); + packet + } + + fn quic_initial_packet(dcid: u8) -> Vec { + let mut packet = vec![0; 1200]; + packet[0] = 0xc0; + packet[4] = 1; + packet[5] = 1; + packet[6] = dcid; + packet[7] = 0; + packet[8] = 0; + packet[9] = 1; + packet + } + + #[tokio::test] + async fn ring_listener_delivers_packet_native_tunnel() -> anyhow::Result<()> { + let local_id = uuid::Uuid::new_v4(); + let registry = Arc::new(RingTunnelRegistry::default()); + let mut listener = RingTransportListener::::new( + format!("ring://{local_id}").parse()?, + registry.clone(), + ); + listener.listen().await?; + let client = registry.connect(local_id)?.into_tunnel(); + let AcceptedTransport::Tunnel { tunnel: server, .. } = listener.accept().await? else { + anyhow::bail!("Ring listener did not produce a Tunnel"); + }; + let (_client_stream, mut client_sink) = client.split(); + let (mut server_stream, _server_sink) = server.split(); + + client_sink + .send(ZCPacket::new_with_payload(b"packet-native-listener")) + .await?; + let packet = server_stream.next().await.transpose()?.unwrap(); + assert_eq!(packet.payload(), b"packet-native-listener"); + Ok(()) + } + + #[tokio::test] + async fn ring_listener_rejects_non_ring_url() { + let registry = Arc::new(RingTunnelRegistry::default()); + let mut listener = RingTransportListener::::new( + format!("tcp://{}", uuid::Uuid::new_v4()).parse().unwrap(), + registry, + ); + + assert!(listener.listen().await.is_err()); + } + + #[tokio::test] + async fn unresolved_listener_urls_use_core_dns_and_protocol_default_ports() -> anyhow::Result<()> + { + let host = Arc::new(MockHost::new()); + let mut tcp = TcpTransportListener::new( + "tcp://listener.example".parse()?, + super::super::plan::unresolved_tcp_listener_options( + SocketContext::default().with_socket_mark(Some(7)), + ), + None, + host.clone(), + Arc::new(MockDns), + ); + tcp.listen().await?; + assert_eq!(host.tcp_listener(0).local_addr()?.port(), 11010); + assert_eq!( + host.tcp_bind_options(0).bind.context.ip_version, + IpVersion::V4 + ); + assert!(!host.tcp_bind_options(0).bind.only_v6); + + let mut udp = UdpTransportListener::::new( + "wg://listener.example".parse()?, + super::super::plan::unresolved_udp_session_listen_request( + &"wg://listener.example".parse()?, + SocketContext::default().with_socket_mark(Some(7)), + ), + UdpSessionAcceptKind::Classified(UdpSessionProtocol::WireGuard), + host.clone(), + Arc::new(MockDns), + ); + udp.listen().await?; + assert_eq!(host.udp_socket(0).local_addr()?.port(), 11011); + assert_eq!(host.udp_bind_options(0).context.ip_version, IpVersion::V4); + assert!(!host.udp_bind_options(0).only_v6); + + let mut tcp_v6 = TcpTransportListener::new( + "tcp://[::1]:0".parse()?, + super::super::plan::unresolved_tcp_listener_options(SocketContext::default()), + None, + host.clone(), + Arc::new(MockDns), + ); + tcp_v6.listen().await?; + assert_eq!( + host.tcp_bind_options(1).bind.context.ip_version, + IpVersion::V6 + ); + assert!(host.tcp_bind_options(1).bind.only_v6); + + let udp_v6_url: Url = "udp://[::1]:0".parse()?; + let mut udp_v6 = UdpTransportListener::::new( + udp_v6_url.clone(), + super::super::plan::unresolved_udp_session_listen_request( + &udp_v6_url, + SocketContext::default(), + ), + UdpSessionAcceptKind::EasyTierMux, + host.clone(), + Arc::new(MockDns), + ); + udp_v6.listen().await?; + assert_eq!(host.udp_bind_options(1).context.ip_version, IpVersion::V6); + assert!(host.udp_bind_options(1).only_v6); + Ok(()) + } + + #[tokio::test] + async fn tcp_listener_applies_protocol_upgrade_limit() -> anyhow::Result<()> { + let host = Arc::new(MockHost::new()); + let mut listener = TcpTransportListener::new( + "tcp://127.0.0.1:0".parse()?, + super::super::plan::unresolved_tcp_listener_options(SocketContext::default()), + Some(std::num::NonZeroUsize::MIN), + host.clone(), + Arc::new(MockDns), + ); + listener.listen().await?; + let socket = host.tcp_listener(0); + socket.accept_from("127.0.0.1:31000".parse()?); + socket.accept_from("127.0.0.1:31001".parse()?); + + let first = listener.accept().await?; + assert!( + crate::foundation::time::timeout(Duration::from_millis(50), listener.accept()) + .await + .is_err() + ); + drop(first); + crate::foundation::time::timeout(Duration::from_secs(1), listener.accept()).await??; + Ok(()) + } + + #[tokio::test] + async fn tcp_upgrade_permit_is_released_before_peer_admission() -> anyhow::Result<()> { + let tunnel_handler = Arc::new(BlockingTunnelHandler { + entered: Notify::new(), + release: Notify::new(), + }); + let protocol = Arc::new(CoreServerProtocolUpgrader::new(Default::default())); + let handler = ProtocolAcceptedTransportHandler::new(&tunnel_handler, protocol); + let upgrade_slots = Arc::new(Semaphore::new(1)); + let permit = upgrade_slots.clone().acquire_owned().await?; + let (stream, _remote) = tokio::io::duplex(64); + + let task = tokio::spawn(async move { + handler + .handle_accepted_socket(AcceptedTransport::Tcp { + socket: MockTcpSocket { + stream, + local_addr: "127.0.0.1:21000".parse().unwrap(), + peer_addr: "127.0.0.1:31000".parse().unwrap(), + }, + local_url: "tcp://127.0.0.1:21000".parse().unwrap(), + upgrade_permit: Some(permit), + }) + .await + }); + + tunnel_handler.entered.notified().await; + let next_permit = + crate::foundation::time::timeout(Duration::from_secs(1), upgrade_slots.acquire_owned()) + .await??; + drop(next_permit); + tunnel_handler.release.notify_one(); + task.await??; + Ok(()) + } + + #[tokio::test] + async fn quic_admission_happens_before_transport_is_returned() -> anyhow::Result<()> { + let host = Arc::new(MockHost::new()); + let mut listener = UdpTransportListener::::new( + "quic://127.0.0.1:0".parse()?, + UdpSessionListenRequest::new(UdpBindOptions::port_bound_listener( + "127.0.0.1:0".parse()?, + )), + UdpSessionAcceptKind::Classified(UdpSessionProtocol::Quic), + host.clone(), + Arc::new(MockDns), + ); + listener.protocol_admission = Some(ServerProtocolAdmissionController::new(1, 1)); + listener.listen().await?; + let socket = host.udp_socket(0); + + socket.receive_from(quic_initial_packet(1), "127.0.0.1:32001".parse()?); + let first = + crate::foundation::time::timeout(Duration::from_secs(1), listener.accept()).await??; + assert!(matches!( + &first, + AcceptedTransport::Udp { + admission: Some(_), + .. + } + )); + + socket.receive_from(quic_initial_packet(2), "127.0.0.1:32002".parse()?); + assert!( + crate::foundation::time::timeout(Duration::from_millis(100), listener.accept()) + .await + .is_err() + ); + + drop(first); + socket.receive_from(quic_initial_packet(3), "127.0.0.1:32003".parse()?); + let third = + crate::foundation::time::timeout(Duration::from_secs(1), listener.accept()).await??; + assert!(matches!( + third, + AcceptedTransport::Udp { + admission: Some(_), + .. + } + )); + Ok(()) + } + + #[tokio::test] + async fn protocol_handler_consumes_multi_tunnel_acceptors_and_udp_upgrades() + -> anyhow::Result<()> { + let tunnel_handler = Arc::new(RecordingTunnelHandler { + calls: AtomicUsize::new(0), + }); + let protocol = Arc::new(RecordingServerProtocolUpgrader::new()); + let handler = ProtocolAcceptedTransportHandler::new(&tunnel_handler, protocol.clone()); + + let (stream, _remote) = tokio::io::duplex(64); + let tcp_result = handler + .handle_accepted_socket(AcceptedTransport::Tcp { + socket: MockTcpSocket { + stream, + local_addr: "127.0.0.1:21000".parse().unwrap(), + peer_addr: "127.0.0.1:31000".parse().unwrap(), + }, + local_url: "tcp://127.0.0.1:21000".parse().unwrap(), + upgrade_permit: None, + }) + .await + .unwrap_err(); + assert_eq!(tcp_result.to_string(), "server tunnel acceptor finished"); + assert_eq!(protocol.tcp_calls.load(Ordering::Relaxed), 1); + assert_eq!(tunnel_handler.calls.load(Ordering::Relaxed), 2); + + let udp_socket = Arc::new(MockUdpSocket::new("127.0.0.1:22000".parse().unwrap())); + let udp_session = UdpSession::identity_standalone( + udp_socket, + "127.0.0.1:32000".parse().unwrap(), + UdpSessionKind::EasyTierMux, + )?; + handler + .handle_accepted_socket(AcceptedTransport::Udp { + session: udp_session, + local_url: "udp://127.0.0.1:22000".parse().unwrap(), + admission: None, + }) + .await?; + assert_eq!(protocol.udp_calls.load(Ordering::Relaxed), 1); + assert_eq!(tunnel_handler.calls.load(Ordering::Relaxed), 3); + + let (stream, _remote) = tokio::io::duplex(64); + handler + .handle_accepted_socket(AcceptedTransport::ByteStream { + socket: MockTcpSocket { + stream, + local_addr: "127.0.0.1:21002".parse()?, + peer_addr: "127.0.0.1:31002".parse()?, + }, + local_url: "external://local".parse()?, + remote_url: Some("external://remote".parse()?), + }) + .await?; + assert_eq!(protocol.byte_stream_calls.load(Ordering::Relaxed), 1); + assert_eq!(tunnel_handler.calls.load(Ordering::Relaxed), 4); + + Ok::<(), anyhow::Error>(()) + } + + #[tokio::test] + async fn core_builtin_server_upgrades_raw_udp_and_byte_streams() -> anyhow::Result<()> { + let tunnel_handler = Arc::new(RecordingTunnelHandler { + calls: AtomicUsize::new(0), + }); + let protocol = Arc::new(CoreServerProtocolUpgrader::new(Default::default())); + let handler = ProtocolAcceptedTransportHandler::new(&tunnel_handler, protocol); + + let udp_socket = Arc::new(MockUdpSocket::new("127.0.0.1:22000".parse()?)); + let udp_session = UdpSession::identity_standalone( + udp_socket, + "127.0.0.1:32000".parse()?, + UdpSessionKind::EasyTierMux, + )?; + let admission_error = handler + .handle_accepted_socket(AcceptedTransport::Udp { + session: udp_session, + local_url: "udp://127.0.0.1:22000".parse()?, + admission: None, + }) + .await + .unwrap_err(); + assert_eq!(admission_error.to_string(), "first admission rejected"); + + let (stream, _remote) = tokio::io::duplex(64); + handler + .handle_accepted_socket(AcceptedTransport::ByteStream { + socket: MockTcpSocket { + stream, + local_addr: "127.0.0.1:21000".parse()?, + peer_addr: "127.0.0.1:31000".parse()?, + }, + local_url: "unix:///tmp/easytier.sock".parse()?, + remote_url: Some("unix://anonymous/remote".parse()?), + }) + .await?; + assert_eq!(tunnel_handler.calls.load(Ordering::Relaxed), 2); + Ok(()) + } + + #[tokio::test] + async fn service_preserves_transport_boundary_and_stops_blocked_handlers() { + let host = Arc::new(MockHost::new()); + let (event_tx, mut event_rx) = mpsc::unbounded_channel(); + let active = Arc::new(AtomicUsize::new(0)); + let handler = Arc::new(RecordingHandler { + events: event_tx, + blocked: Arc::new(Notify::new()), + active: active.clone(), + }); + let service = CoreListenerRuntime::new_with_events( + host.clone(), + Arc::new(MockDns), + Arc::new(RingTunnelRegistry::default()), + vec![ + TransportListenerConfig::Tcp { + url: "tcp://127.0.0.1:0".parse().unwrap(), + options: TcpListenOptions::manual_connect("127.0.0.1:0".parse().unwrap()), + max_pending_upgrades: None, + must_succeed: true, + }, + TransportListenerConfig::Udp { + url: "udp://127.0.0.1:0".parse().unwrap(), + request: UdpSessionListenRequest::new(UdpBindOptions::port_bound_listener( + "127.0.0.1:0".parse().unwrap(), + )), + accept_kind: UdpSessionAcceptKind::EasyTierMux, + must_succeed: true, + }, + TransportListenerConfig::Udp { + url: "wg://127.0.0.1:0".parse().unwrap(), + request: UdpSessionListenRequest::new(UdpBindOptions::port_bound_listener( + "127.0.0.1:0".parse().unwrap(), + )), + accept_kind: UdpSessionAcceptKind::Classified(UdpSessionProtocol::WireGuard), + must_succeed: true, + }, + ], + Vec::new(), + Vec::new(), + handler, + Arc::new(RecordingListenerEvents::default()), + Arc::new(RunningListenerRegistry::default()), + ); + + service.start().await.unwrap(); + host.tcp_listener(0) + .accept_from("127.0.0.1:31000".parse().unwrap()); + host.udp_socket(0).receive_from( + new_syn_packet(1, 2).into_bytes().to_vec(), + "127.0.0.1:32000".parse().unwrap(), + ); + host.udp_socket(1) + .receive_from(wireguard_packet(), "127.0.0.1:32001".parse().unwrap()); + + let mut events = Vec::new(); + for _ in 0..3 { + events.push( + crate::foundation::time::timeout(Duration::from_secs(1), event_rx.recv()) + .await + .expect("accepted transport was not handled") + .expect("accepted transport event channel closed"), + ); + } + assert!(events.contains(&AcceptedEvent::Tcp { port: 21000 })); + assert!(events.contains(&AcceptedEvent::Udp { + port: 22000, + kind: UdpSessionKind::EasyTierMux, + })); + assert!(events.contains(&AcceptedEvent::Udp { + port: 22001, + kind: UdpSessionKind::WireGuard, + })); + assert_eq!(active.load(Ordering::Relaxed), 3); + + crate::foundation::time::timeout(Duration::from_secs(1), service.stop()) + .await + .expect("transport listener service did not stop"); + assert_eq!(active.load(Ordering::Relaxed), 0); + } + + #[tokio::test] + async fn runtime_publishes_plan_failures_before_starting_listeners() { + let events = Arc::new(RecordingListenerEvents::default()); + let service = CoreListenerRuntime::new_with_events( + Arc::new(MockHost::new()), + Arc::new(MockDns), + Arc::new(RingTunnelRegistry::default()), + Vec::new(), + Vec::new(), + vec![ListenerPlanFailure { + url: "unsupported://listener".parse().unwrap(), + message: "unsupported listener".to_owned(), + }], + Arc::new(|_: AcceptedTransport| async { Ok(()) }), + events.clone(), + Arc::new(RunningListenerRegistry::default()), + ); + + service.start().await.unwrap(); + assert!(matches!( + events.events.lock().unwrap().as_slice(), + [CoreEvent::ListenerPlanFailed { url, error }] + if url.as_str() == "unsupported://listener" + && error == "unsupported listener" + )); + service.stop().await; + } +} diff --git a/easytier-core/src/management/full/compiled.rs b/easytier-core/src/management/full/compiled.rs new file mode 100644 index 00000000..6e131c43 --- /dev/null +++ b/easytier-core/src/management/full/compiled.rs @@ -0,0 +1,45 @@ +use std::sync::Arc; + +use easytier_proto::{ + api::{ + config::ConfigRpcServer, + instance::{ + AclManageRpcServer, ConnectorManageRpcServer, CredentialManageRpcServer, + MappedListenerManageRpcServer, PeerManageRpcServer, PortForwardManageRpcServer, + StatsRpcServer, VpnPortalRpcServer, + }, + }, + peer_rpc::PeerCenterRpcServer, +}; + +use super::super::instance_rpc::InstanceManagementRpc; +use crate::{ + instance::{ + CoreInstance, CoreInstanceHost, + manager::{InstanceFactory, InstanceManager}, + }, + rpc::service_registry::ServiceRegistry, +}; + +/// Registers each Instance-targeted management protocol Interface once for +/// the complete process-level Instance collection. +pub fn register_instance_management_rpc( + manager: Arc>, + registry: &ServiceRegistry, +) where + F: InstanceFactory>, + H: CoreInstanceHost, +{ + let rpc = InstanceManagementRpc::::new(manager.clone()); + registry.register(PeerManageRpcServer::new(rpc.clone()), ""); + registry.register(ConnectorManageRpcServer::new(rpc.clone()), ""); + registry.register(MappedListenerManageRpcServer::new(rpc.clone()), ""); + registry.register(VpnPortalRpcServer::new(rpc.clone()), ""); + super::packet_proxy::register(manager.clone(), registry); + registry.register(AclManageRpcServer::new(rpc.clone()), ""); + registry.register(PortForwardManageRpcServer::new(rpc.clone()), ""); + registry.register(StatsRpcServer::new(rpc.clone()), ""); + registry.register(ConfigRpcServer::new(rpc.clone()), ""); + registry.register(CredentialManageRpcServer::new(rpc.clone()), ""); + registry.register(PeerCenterRpcServer::new(rpc), ""); +} diff --git a/easytier-core/src/management/full/config_patch.rs b/easytier-core/src/management/full/config_patch.rs new file mode 100644 index 00000000..6dd29ee1 --- /dev/null +++ b/easytier-core/src/management/full/config_patch.rs @@ -0,0 +1,355 @@ +use std::{fmt::Debug, sync::Arc}; + +use anyhow::Context as _; +use easytier_proto::api::config::{ + self, AclPatch, ConfigPatchAction, ExitNodePatch, InstanceConfigPatch, Patchable, + PortForwardPatch, ProxyNetworkPatch, RoutePatch, UrlPatch, +}; + +use crate::{ + config::{ + peers::AclRuleConfig, + runtime::CoreInstanceRuntimeConfig, + toml::{ConfigLoader as _, TomlConfig}, + }, + instance::{CoreInstance, CoreInstanceConfig, CoreInstanceHost, CoreInstanceState}, +}; + +pub async fn apply_config_patch( + instance: &Arc>, + patch: InstanceConfigPatch, +) -> anyhow::Result<()> +where + H: CoreInstanceHost, +{ + let _operation = instance.operation.lock().await; + if instance.state() != CoreInstanceState::Running { + anyhow::bail!("instance is not ready; config patch rejected"); + } + + let config = instance + .toml_config() + .ok_or_else(|| anyhow::anyhow!("shared TOML configuration is not available"))?; + let candidate = config.detached_snapshot(); + let parsed_prefix = validate_public_ipv6_patch(instance, &config, &patch)?; + let patch_for_host = patch.clone(); + + // Preserve the existing ordered partial-commit contract: earlier valid + // sub-patches remain applied if a later sub-patch fails. + let patch_result: anyhow::Result = async { + let result = patch_port_forwards(&candidate, patch.port_forwards); + validate_and_commit_candidate(instance, &config, &candidate)?; + result?; + + let result = patch_acl(&candidate, patch.acl); + validate_and_commit_candidate(instance, &config, &candidate)?; + result?; + + let result = patch_proxy_networks(&candidate, patch.proxy_networks); + validate_and_commit_candidate(instance, &config, &candidate)?; + result?; + + let result = patch_routes(&candidate, patch.routes); + validate_and_commit_candidate(instance, &config, &candidate)?; + result?; + + let result = patch_exit_nodes_config(&candidate, patch.exit_nodes); + validate_and_commit_candidate(instance, &config, &candidate)?; + instance.update_exit_nodes(result?).await; + + let result = patch_mapped_listeners(&candidate, patch.mapped_listeners); + validate_and_commit_candidate(instance, &config, &candidate)?; + result?; + + patch_connectors(instance, patch.connectors)?; + + let mut provider_config_changed = false; + if let Some(hostname) = patch.hostname { + candidate.set_hostname(Some(hostname)); + } + if let Some(ipv4) = patch.ipv4 + && !candidate.get_dhcp() + { + candidate.set_ipv4(Some(ipv4.into())); + } + if let Some(ipv6) = patch.ipv6 { + candidate.set_ipv6(Some(ipv6.into())); + } + if let Some(disable_relay_data) = patch.disable_relay_data { + let mut flags = candidate.get_flags(); + flags.disable_relay_data = disable_relay_data; + candidate.set_flags(flags); + } + if let Some(enabled) = patch.ipv6_public_addr_provider { + candidate.set_ipv6_public_addr_provider(enabled); + provider_config_changed = true; + } + if let Some(enabled) = patch.ipv6_public_addr_auto { + candidate.set_ipv6_public_addr_auto(enabled); + } + if let Some(prefix) = parsed_prefix { + candidate.set_ipv6_public_addr_prefix(prefix); + provider_config_changed = true; + } + validate_and_commit_candidate(instance, &config, &candidate)?; + instance + .instance_runtime + .synchronize_config(&patch_for_host); + Ok(provider_config_changed) + } + .await; + + instance + .update_runtime_config_under_operation(runtime_config_from_toml(instance, &config)?) + .await?; + let provider_config_changed = patch_result?; + instance + .instance_runtime + .publish_config_patch(patch_for_host); + if provider_config_changed && instance.state() == CoreInstanceState::Running { + instance.reconcile_public_ipv6_provider().await; + } + Ok(()) +} + +fn validate_and_commit_candidate( + instance: &CoreInstance, + shared: &TomlConfig, + candidate: &TomlConfig, +) -> anyhow::Result<()> +where + H: CoreInstanceHost, +{ + let runtime = runtime_config_from_toml(instance, candidate)?; + instance.validate_runtime_config_capabilities(&runtime)?; + shared.replace_from_snapshot(candidate); + Ok(()) +} + +fn runtime_config_from_toml( + instance: &CoreInstance, + config: &TomlConfig, +) -> anyhow::Result +where + H: CoreInstanceHost, +{ + let normalized = CoreInstanceConfig::from_toml_with_host(config, instance.host_config())?; + let current = instance.runtime_config_snapshot(); + let services = normalized.connectivity.runtime; + let mut peer = normalized.peer.snapshot; + peer.runtime.stun_info = current.peer.runtime.stun_info.clone(); + + Ok(CoreInstanceRuntimeConfig { + services, + peer: Arc::new(peer), + }) +} + +fn parse_ipv6_public_addr_prefix_patch( + prefix: Option<&str>, +) -> anyhow::Result>> { + let Some(prefix) = prefix else { + return Ok(None); + }; + let prefix = prefix.trim(); + if prefix.is_empty() { + return Ok(Some(None)); + } + Ok(Some(Some(prefix.parse().with_context(|| { + format!("failed to parse ipv6 public address prefix: {prefix}") + })?))) +} + +fn validate_public_ipv6_patch( + instance: &CoreInstance, + config: &TomlConfig, + patch: &InstanceConfigPatch, +) -> anyhow::Result>> +where + H: CoreInstanceHost, +{ + let parsed_prefix = + parse_ipv6_public_addr_prefix_patch(patch.ipv6_public_addr_prefix.as_deref())?; + let provider_enabled = patch + .ipv6_public_addr_provider + .unwrap_or(config.get_ipv6_public_addr_provider()); + let configured_prefix = parsed_prefix.unwrap_or_else(|| config.get_ipv6_public_addr_prefix()); + let provider_supported = instance + .runtime_config_snapshot() + .services + .public_ipv6_provider + .provider_supported; + crate::config::peers::PublicIpv6ProviderConfig { + provider_enabled, + configured_prefix, + provider_supported, + } + .validate()?; + Ok(parsed_prefix) +} + +fn trace_patchables(patches: &[Patchable]) { + for patch in patches { + match patch.action { + Some(ConfigPatchAction::Add) | Some(ConfigPatchAction::Remove) => { + if let Some(value) = &patch.value { + tracing::info!(?patch.action, ?value, "applying configuration patch"); + } else { + tracing::warn!(?patch.action, "ignored configuration patch without value"); + } + } + Some(ConfigPatchAction::Clear) => { + tracing::info!("clearing configuration collection"); + } + None => tracing::warn!("ignored invalid configuration patch action"), + } + } +} + +fn patch_port_forwards(config: &TomlConfig, patches: Vec) -> anyhow::Result<()> { + if patches.is_empty() { + return Ok(()); + } + let mut current = config.get_port_forwards(); + let patches = patches + .into_iter() + .map(|patch| Patchable { + action: ConfigPatchAction::try_from(patch.action).ok(), + value: patch.cfg.map(Into::into), + }) + .collect::>(); + trace_patchables(&patches); + config::patch_vec(&mut current, patches); + config.set_port_forwards(current); + Ok(()) +} + +fn patch_acl(config: &TomlConfig, patch: Option) -> anyhow::Result<()> { + let Some(patch) = patch else { + return Ok(()); + }; + let mut acl = AclRuleConfig { + acl: config.get_acl(), + tcp_whitelist: config.get_tcp_whitelist(), + udp_whitelist: config.get_udp_whitelist(), + whitelist_priority: None, + }; + if let Some(next) = patch.acl { + acl.acl = Some(next); + } + if !patch.tcp_whitelist.is_empty() { + let patches = patch + .tcp_whitelist + .into_iter() + .map(Into::into) + .collect::>(); + trace_patchables(&patches); + config::patch_vec(&mut acl.tcp_whitelist, patches); + } + if !patch.udp_whitelist.is_empty() { + let patches = patch + .udp_whitelist + .into_iter() + .map(Into::into) + .collect::>(); + trace_patchables(&patches); + config::patch_vec(&mut acl.udp_whitelist, patches); + } + acl.build()?; + config.set_acl(acl.acl); + config.set_tcp_whitelist(acl.tcp_whitelist); + config.set_udp_whitelist(acl.udp_whitelist); + Ok(()) +} + +fn patch_proxy_networks( + config: &TomlConfig, + patches: Vec, +) -> anyhow::Result<()> { + for patch in patches { + match ConfigPatchAction::try_from(patch.action) { + Ok(ConfigPatchAction::Add) => { + let Some(cidr) = patch.cidr.map(Into::into) else { + tracing::warn!("ignored proxy-network add without CIDR"); + continue; + }; + config.add_proxy_cidr(cidr, patch.mapped_cidr.map(Into::into))?; + } + Ok(ConfigPatchAction::Remove) => { + let Some(cidr) = patch.cidr.map(Into::into) else { + tracing::warn!("ignored proxy-network remove without CIDR"); + continue; + }; + config.remove_proxy_cidr(cidr); + } + Ok(ConfigPatchAction::Clear) => config.clear_proxy_cidrs(), + Err(_) => tracing::warn!( + action = patch.action, + "ignored invalid proxy-network action" + ), + } + } + Ok(()) +} + +fn patch_routes(config: &TomlConfig, patches: Vec) -> anyhow::Result<()> { + if patches.is_empty() { + return Ok(()); + } + let mut current = config.get_routes().unwrap_or_default(); + let patches = patches.into_iter().map(Into::into).collect::>(); + trace_patchables(&patches); + config::patch_vec(&mut current, patches); + config.set_routes((!current.is_empty()).then_some(current)); + Ok(()) +} + +fn patch_exit_nodes_config( + config: &TomlConfig, + patches: Vec, +) -> anyhow::Result> { + if patches.is_empty() { + return Ok(config.get_exit_nodes()); + } + let mut current = config.get_exit_nodes(); + let patches = patches.into_iter().map(Into::into).collect::>(); + trace_patchables(&patches); + config::patch_vec(&mut current, patches); + config.set_exit_nodes(current.clone()); + Ok(current) +} + +fn patch_mapped_listeners(config: &TomlConfig, patches: Vec) -> anyhow::Result<()> { + if patches.is_empty() { + return Ok(()); + } + let mut current = config.get_mapped_listeners(); + let patches = patches.into_iter().map(Into::into).collect::>(); + trace_patchables(&patches); + config::patch_vec(&mut current, patches); + config.set_mapped_listeners((!current.is_empty()).then_some(current)); + Ok(()) +} + +fn patch_connectors(instance: &CoreInstance, patches: Vec) -> anyhow::Result<()> +where + H: CoreInstanceHost, +{ + for patch in patches { + let Some(url) = patch.url.map(Into::::into) else { + tracing::warn!("ignored connector patch without URL"); + return Ok(()); + }; + match ConfigPatchAction::try_from(patch.action) { + Ok(ConfigPatchAction::Add) => instance.add_connector(url)?, + Ok(ConfigPatchAction::Remove) => { + if !instance.remove_connector(&url) { + anyhow::bail!("connector not found: {url}"); + } + } + Ok(ConfigPatchAction::Clear) => instance.clear_connectors(), + Err(_) => tracing::warn!(action = patch.action, "ignored invalid connector action"), + } + } + Ok(()) +} diff --git a/easytier-core/src/management/full/instance_info.rs b/easytier-core/src/management/full/instance_info.rs new file mode 100644 index 00000000..f73ab476 --- /dev/null +++ b/easytier-core/src/management/full/instance_info.rs @@ -0,0 +1,79 @@ +use easytier_proto::api::{ + instance::{PeerInfo, list_peer_route_pair}, + manage::{MyNodeInfo, NetworkInstanceRunningInfo}, +}; + +use crate::{ + config::toml::ConfigLoader as _, + instance::{CoreInstance, CoreInstanceHost, CoreInstanceState}, +}; + +/// Builds the process-level running snapshot directly from one core Instance. +pub async fn network_instance_running_info( + instance: &CoreInstance, +) -> anyhow::Result +where + H: CoreInstanceHost, +{ + let running = !matches!( + instance.state(), + CoreInstanceState::Created | CoreInstanceState::Stopped + ); + if !instance.is_ready() { + return Ok(NetworkInstanceRunningInfo { + running, + error_msg: instance.latest_error(), + ..Default::default() + }); + } + + let peers = instance + .peer_snapshots() + .await + .into_iter() + .map(|snapshot| PeerInfo { + peer_id: snapshot.peer_id, + default_conn_id: snapshot.default_conn_id.map(Into::into), + directly_connected_conns: snapshot + .directly_connected_conns + .into_iter() + .map(Into::into) + .collect(), + conns: snapshot.conns.into_iter().map(Into::into).collect(), + }) + .collect::>(); + let node = instance.node_snapshot().await; + let routes = instance + .route_snapshots() + .await + .into_iter() + .map(Into::into) + .collect::>(); + let peer_route_pairs = list_peer_route_pair(peers.clone(), routes.clone()); + let vpn_portal_cfg = Some(instance.vpn_portal_info().await.client_config); + let dev_name = instance + .toml_config() + .map(|config| config.get_flags().dev_name) + .unwrap_or_default(); + + Ok(NetworkInstanceRunningInfo { + dev_name, + my_node_info: Some(MyNodeInfo { + virtual_ipv4: node.ipv4_addr.map(Into::into), + hostname: node.hostname, + version: node.version, + ips: Some(node.ip_list), + stun_info: Some(node.stun_info), + listeners: node.listeners.into_iter().map(Into::into).collect(), + vpn_portal_cfg, + peer_id: node.peer_id, + }), + events: instance.management_events(), + routes, + peers, + peer_route_pairs, + running, + error_msg: instance.latest_error(), + foreign_network_summary: Some(instance.foreign_network_route_summary().await), + }) +} diff --git a/easytier-core/src/management/full/logger_rpc.rs b/easytier-core/src/management/full/logger_rpc.rs new file mode 100644 index 00000000..052e4e74 --- /dev/null +++ b/easytier-core/src/management/full/logger_rpc.rs @@ -0,0 +1,126 @@ +use std::sync::Arc; + +use easytier_proto::{ + api::logger::{ + GetLoggerConfigRequest, GetLoggerConfigResponse, LogLevel, LoggerRpc, + SetLoggerConfigRequest, SetLoggerConfigResponse, + }, + rpc_types::{self, controller::BaseController}, +}; + +pub trait LoggerControl: Send + Sync + 'static { + fn set_level(&self, level: &str) -> anyhow::Result<()>; + + fn level(&self) -> String; +} + +#[derive(Default)] +pub struct UnsupportedLoggerControl; + +impl LoggerControl for UnsupportedLoggerControl { + fn set_level(&self, _level: &str) -> anyhow::Result<()> { + anyhow::bail!("logger control is unsupported by this Host") + } + + fn level(&self) -> String { + "info".to_owned() + } +} + +#[derive(Clone)] +pub struct LoggerManagementRpc { + control: Arc, +} + +impl LoggerManagementRpc { + pub fn new(control: Arc) -> Self { + Self { control } + } +} + +pub fn log_level_name(level: LogLevel) -> &'static str { + match level { + LogLevel::Disabled => "off", + LogLevel::Error => "error", + LogLevel::Warning => "warn", + LogLevel::Info => "info", + LogLevel::Debug => "debug", + LogLevel::Trace => "trace", + } +} + +pub fn parse_log_level(level: &str) -> LogLevel { + match level.to_ascii_lowercase().as_str() { + "off" | "disabled" => LogLevel::Disabled, + "error" => LogLevel::Error, + "warn" | "warning" => LogLevel::Warning, + "debug" => LogLevel::Debug, + "trace" => LogLevel::Trace, + _ => LogLevel::Info, + } +} + +#[async_trait::async_trait] +impl LoggerRpc for LoggerManagementRpc { + type Controller = BaseController; + + async fn set_logger_config( + &self, + _: BaseController, + request: SetLoggerConfigRequest, + ) -> rpc_types::error::Result { + self.control.set_level(log_level_name(request.level()))?; + Ok(SetLoggerConfigResponse {}) + } + + async fn get_logger_config( + &self, + _: BaseController, + _: GetLoggerConfigRequest, + ) -> rpc_types::error::Result { + Ok(GetLoggerConfigResponse { + level: parse_log_level(&self.control.level()).into(), + }) + } +} + +#[cfg(test)] +mod tests { + use std::sync::Mutex; + + use super::*; + + #[derive(Default)] + struct TestLoggerControl(Mutex); + + impl LoggerControl for TestLoggerControl { + fn set_level(&self, level: &str) -> anyhow::Result<()> { + *self.0.lock().unwrap() = level.to_owned(); + Ok(()) + } + + fn level(&self) -> String { + self.0.lock().unwrap().clone() + } + } + + #[tokio::test] + async fn logger_rpc_delegates_process_effects_to_host_control() { + let rpc = LoggerManagementRpc::new(Arc::new(TestLoggerControl::default())); + + rpc.set_logger_config( + BaseController::default(), + SetLoggerConfigRequest { + level: LogLevel::Debug.into(), + }, + ) + .await + .unwrap(); + let response = rpc + .get_logger_config(BaseController::default(), GetLoggerConfigRequest {}) + .await + .unwrap(); + + assert_eq!(response.level(), LogLevel::Debug); + } +} diff --git a/easytier-core/src/management/full/mod.rs b/easytier-core/src/management/full/mod.rs new file mode 100644 index 00000000..837bb3f8 --- /dev/null +++ b/easytier-core/src/management/full/mod.rs @@ -0,0 +1,110 @@ +mod compiled; +mod config_patch; +mod instance_info; +mod logger_rpc; +pub(super) mod packet_proxy; +mod process_rpc; +pub mod remote_client; +mod web_client; + +use std::sync::Arc; + +use easytier_proto::{ + api::{ + logger::{LoggerRpc, LoggerRpcServer}, + manage::WebClientServiceServer, + }, + rpc_types::controller::BaseController, +}; + +use crate::{ + config::toml::ConfigSource, + instance::{ + CoreInstance, CoreInstanceHost, + manager::{InstanceFactory, InstanceManager}, + }, + rpc::service_registry::ServiceRegistry, +}; + +use super::{ + ConfigFileControl, ConfigFilePermission, DaemonGuard, resolve_optional_instance_by_name, +}; + +pub use compiled::register_instance_management_rpc; +pub use config_patch::apply_config_patch; +pub use instance_info::network_instance_running_info; +pub use logger_rpc::{ + LoggerControl, LoggerManagementRpc, UnsupportedLoggerControl, log_level_name, parse_log_level, +}; +pub use process_rpc::{ + ConfigFileStorage, InstanceMutationHooks, InstanceMutationResult, ProcessManagement, + ProcessManagementRpc, UnsupportedConfigFileStorage, +}; +pub use web_client::{ConfigServerEndpoint, WebClient, WebClientConfig}; + +pub use super::instance_rpc::full::call_instance_json_rpc; + +pub fn config_source_from_rpc(source: i32) -> Option { + match easytier_proto::api::manage::ConfigSource::try_from(source).ok() { + Some(easytier_proto::api::manage::ConfigSource::Web) => Some(ConfigSource::Web), + Some(easytier_proto::api::manage::ConfigSource::User) => Some(ConfigSource::User), + _ => None, + } +} + +pub fn config_source_to_rpc(source: ConfigSource) -> i32 { + match source { + ConfigSource::User => easytier_proto::api::manage::ConfigSource::User as i32, + ConfigSource::Web => easytier_proto::api::manage::ConfigSource::Web as i32, + } +} + +/// Registers the complete process-level management surface once. +pub fn register_management_rpc( + instances: Arc>, + registry: &ServiceRegistry, + hooks: Arc, + storage: Arc, + logger: Arc, +) where + F: InstanceFactory, CreateContext = ()>, + F::Error: std::fmt::Debug + std::fmt::Display + Send + Sync + 'static, + H: CoreInstanceHost, +{ + register_instance_management_rpc(instances.clone(), registry); + registry.register(LoggerRpcServer::new(LoggerManagementRpc::new(logger)), ""); + registry.register( + WebClientServiceServer::new(ProcessManagementRpc::::new(instances, hooks, storage)), + "", + ); +} + +pub async fn call_management_json_rpc( + manager: &Arc>, + logger: Arc, + service_name: &str, + method_name: &str, + domain_name: Option<&str>, + payload: serde_json::Value, +) -> crate::proto::rpc_types::error::Result +where + F: InstanceFactory>, + H: CoreInstanceHost, +{ + if service_name == "api.manage.WebClientService" { + return Err(anyhow::anyhow!( + "service {service_name} is not exposed through FFI/JNI generic RPC" + ) + .into()); + } + if service_name == "api.logger.LoggerRpcService" { + return LoggerRpc::json_call_method( + &LoggerManagementRpc::new(logger), + BaseController::default(), + method_name, + payload, + ) + .await; + } + call_instance_json_rpc(manager, service_name, method_name, domain_name, payload).await +} diff --git a/easytier-core/src/management/full/packet_proxy.rs b/easytier-core/src/management/full/packet_proxy.rs new file mode 100644 index 00000000..9c930f24 --- /dev/null +++ b/easytier-core/src/management/full/packet_proxy.rs @@ -0,0 +1,156 @@ +use std::sync::Arc; + +#[cfg(feature = "proxy-packet")] +use easytier_proto::{ + api::instance::{TcpProxyRpc, TcpProxyRpcServer}, + rpc_types::controller::BaseController, +}; + +#[cfg(feature = "proxy-packet")] +use crate::{ + gateway::proxy::wrapped_transport::{WrappedTransportKind, WrappedTransportRole}, + management::instance_rpc::packet_proxy::TcpProxyManagementRpc, +}; +use crate::{ + instance::{ + CoreInstance, CoreInstanceHost, + manager::{InstanceFactory, InstanceManager}, + }, + rpc::service_registry::ServiceRegistry, +}; + +pub(in crate::management) type JsonCall = + Result, serde_json::Value>; + +#[cfg(feature = "proxy-packet")] +pub(in crate::management) fn register( + manager: Arc>, + registry: &ServiceRegistry, +) where + F: InstanceFactory>, + H: CoreInstanceHost, +{ + registry.register( + TcpProxyRpcServer::new(TcpProxyManagementRpc::::tcp(manager.clone())), + "tcp", + ); + for (domain, transport, role) in [ + ( + "kcp_src", + WrappedTransportKind::Kcp, + WrappedTransportRole::Source, + ), + ( + "kcp_dst", + WrappedTransportKind::Kcp, + WrappedTransportRole::Destination, + ), + ( + "quic_src", + WrappedTransportKind::Quic, + WrappedTransportRole::Source, + ), + ( + "quic_dst", + WrappedTransportKind::Quic, + WrappedTransportRole::Destination, + ), + ] { + registry.register( + TcpProxyRpcServer::new(TcpProxyManagementRpc::::wrapped( + manager.clone(), + transport, + role, + )), + domain, + ); + } +} + +#[cfg(not(feature = "proxy-packet"))] +pub(in crate::management) fn register( + _manager: Arc>, + _registry: &ServiceRegistry, +) where + F: InstanceFactory>, + H: CoreInstanceHost, +{ +} + +#[cfg(feature = "proxy-packet")] +pub(in crate::management) async fn call_json( + manager: &Arc>, + service_name: &str, + method_name: &str, + domain_name: Option<&str>, + payload: serde_json::Value, +) -> JsonCall +where + F: InstanceFactory>, + H: CoreInstanceHost, +{ + if service_name != "api.instance.TcpProxyRpcService" { + return Err(payload); + } + + let rpc = match tcp_proxy_json_service(manager.clone(), domain_name) { + Ok(rpc) => rpc, + Err(error) => return Ok(Err(error)), + }; + Ok(rpc + .json_call_method(BaseController::default(), method_name, payload) + .await) +} + +#[cfg(not(feature = "proxy-packet"))] +pub(in crate::management) async fn call_json( + _manager: &Arc>, + _service_name: &str, + _method_name: &str, + _domain_name: Option<&str>, + payload: serde_json::Value, +) -> JsonCall +where + F: InstanceFactory>, + H: CoreInstanceHost, +{ + Err(payload) +} + +#[cfg(feature = "proxy-packet")] +fn tcp_proxy_json_service( + manager: Arc>, + domain_name: Option<&str>, +) -> crate::proto::rpc_types::error::Result> +where + F: InstanceFactory>, + H: CoreInstanceHost, +{ + let rpc = match domain_name { + None | Some("") | Some("tcp") => TcpProxyManagementRpc::tcp(manager), + Some("kcp_src") => TcpProxyManagementRpc::wrapped( + manager, + WrappedTransportKind::Kcp, + WrappedTransportRole::Source, + ), + Some("kcp_dst") => TcpProxyManagementRpc::wrapped( + manager, + WrappedTransportKind::Kcp, + WrappedTransportRole::Destination, + ), + Some("quic_src") => TcpProxyManagementRpc::wrapped( + manager, + WrappedTransportKind::Quic, + WrappedTransportRole::Source, + ), + Some("quic_dst") => TcpProxyManagementRpc::wrapped( + manager, + WrappedTransportKind::Quic, + WrappedTransportRole::Destination, + ), + Some(domain) => { + return Err(anyhow::anyhow!("invalid TcpProxyRpcService domain_name: {domain}").into()); + } + }; + Ok(rpc) +} diff --git a/easytier-core/src/management/full/process_rpc.rs b/easytier-core/src/management/full/process_rpc.rs new file mode 100644 index 00000000..1632d9ff --- /dev/null +++ b/easytier-core/src/management/full/process_rpc.rs @@ -0,0 +1,746 @@ +use std::{ + collections::HashSet, + path::{Path, PathBuf}, + sync::Arc, +}; + +use easytier_proto::{ + api::manage::{ + CollectNetworkInfoRequest, CollectNetworkInfoResponse, DeleteNetworkInstanceRequest, + DeleteNetworkInstanceResponse, GetNetworkInstanceConfigRequest, + GetNetworkInstanceConfigResponse, ListNetworkInstanceMetaRequest, + ListNetworkInstanceMetaResponse, ListNetworkInstanceRequest, ListNetworkInstanceResponse, + NetworkInstanceRunningInfoMap, NetworkMeta, RetainNetworkInstanceRequest, + RetainNetworkInstanceResponse, RunNetworkInstanceRequest, RunNetworkInstanceResponse, + ValidateConfigRequest, ValidateConfigResponse, WebClientService, + }, + rpc_types::{self, controller::BaseController}, +}; + +use crate::{ + config::{ + api::network_config_from_toml, + api_input::NetworkConfigExt as _, + toml::{ConfigLoader as _, ConfigSource, TomlConfig}, + }, + instance::{CoreInstance, CoreInstanceHost, manager::InstanceFactory}, +}; + +use super::{ + ConfigFileControl, ConfigFilePermission, InstanceManager, config_source_from_rpc, + config_source_to_rpc, +}; + +#[async_trait::async_trait] +pub trait InstanceMutationHooks: Send + Sync + 'static { + fn manages_remote_config_instances(&self) -> bool { + false + } + + async fn pre_run_network_instance(&self, _config: &TomlConfig) -> Result<(), String> { + Ok(()) + } + + async fn post_run_network_instance(&self, _instance_id: &uuid::Uuid) -> Result<(), String> { + Ok(()) + } + + async fn post_remove_network_instances( + &self, + _instance_ids: &[uuid::Uuid], + ) -> Result<(), String> { + Ok(()) + } +} + +#[async_trait::async_trait] +impl InstanceMutationHooks for () {} + +/// Host Adapter for configuration-file effects used by process management. +#[async_trait::async_trait] +pub trait ConfigFileStorage: Send + Sync + 'static { + async fn inspect(&self, path: &Path) -> ConfigFileControl; + + async fn read(&self, path: &Path) -> anyhow::Result>>; + + async fn write(&self, path: &Path, contents: &[u8]) -> anyhow::Result<()>; + + async fn remove(&self, path: &Path) -> anyhow::Result<()>; +} + +#[derive(Default)] +pub struct UnsupportedConfigFileStorage; + +#[async_trait::async_trait] +impl ConfigFileStorage for UnsupportedConfigFileStorage { + async fn inspect(&self, path: &Path) -> ConfigFileControl { + ConfigFileControl::new( + Some(path.to_owned()), + ConfigFilePermission::from(ConfigFilePermission::READ_ONLY), + ) + } + + async fn read(&self, _path: &Path) -> anyhow::Result>> { + anyhow::bail!("configuration-file storage is unsupported by this Host") + } + + async fn write(&self, _path: &Path, _contents: &[u8]) -> anyhow::Result<()> { + anyhow::bail!("configuration-file storage is unsupported by this Host") + } + + async fn remove(&self, _path: &Path) -> anyhow::Result<()> { + anyhow::bail!("configuration-file storage is unsupported by this Host") + } +} + +/// Result of one process-level removal transaction. +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct InstanceMutationResult { + pub remaining_instance_ids: Vec, + pub removed_instance_ids: Vec, +} + +/// Transport-independent process-level Instance management. +pub struct ProcessManagement +where + F: InstanceFactory, +{ + instances: Arc>, + hooks: Arc, + storage: Arc, + mutation_lock: Arc>, +} + +impl Clone for ProcessManagement +where + F: InstanceFactory, +{ + fn clone(&self) -> Self { + Self { + instances: self.instances.clone(), + hooks: self.hooks.clone(), + storage: self.storage.clone(), + mutation_lock: self.mutation_lock.clone(), + } + } +} + +impl ProcessManagement +where + F: InstanceFactory, CreateContext = ()>, + F::Error: std::fmt::Debug + std::fmt::Display + Send + Sync + 'static, + H: CoreInstanceHost, +{ + pub fn new( + instances: Arc>, + hooks: Arc, + storage: Arc, + ) -> Self { + let mutation_lock = instances.mutation_lock(); + Self { + instances, + hooks, + storage, + mutation_lock, + } + } + + async fn is_remote_removable(&self, control: &ConfigFileControl) -> bool { + if control.is_read_only() || !control.is_deletable() { + return false; + } + let Some(path) = control.path.as_deref() else { + return true; + }; + !self.storage.inspect(path).await.is_read_only() + } + + async fn ensure_overwritable( + &self, + instance_id: uuid::Uuid, + control: &ConfigFileControl, + require_deletable: bool, + ) -> anyhow::Result<()> { + if control.is_read_only() { + anyhow::bail!("instance {instance_id} is read-only, cannot be overwritten"); + } + if require_deletable && !control.is_deletable() { + anyhow::bail!("instance {instance_id} is no-delete, cannot be overwritten"); + } + if let Some(path) = control.path.as_deref() + && self.storage.inspect(path).await.is_read_only() + { + anyhow::bail!( + "config file {} is read-only, cannot be overwritten", + path.display() + ); + } + Ok(()) + } + + async fn apply_file_cleanup(&self, cleanup: ConfigFileCleanup) { + let result = match cleanup { + ConfigFileCleanup::Remove(path) => self.storage.remove(&path).await, + ConfigFileCleanup::Restore { path, contents } => { + self.storage.write(&path, &contents).await + } + }; + if let Err(error) = result { + tracing::warn!(%error, "failed to roll back configuration file"); + } + } + + async fn rollback_started_instance( + &self, + started_instance_id: Option, + file_cleanup: Option, + restore_instance: Option<(TomlConfig, ConfigFileControl)>, + ) { + if let Some(instance_id) = started_instance_id + && let Err(error) = self.instances.delete_network_instances([instance_id]).await + { + tracing::warn!(%error, "failed to remove rolled-back instance"); + return; + } + if let Some(cleanup) = file_cleanup { + self.apply_file_cleanup(cleanup).await; + } + if let Some((config, control)) = restore_instance + && let Err(error) = self.instances.run_network_instance(config, control) + { + tracing::warn!(%error, "failed to restore overwritten instance"); + } + } + + pub async fn run_network_instance( + &self, + config: TomlConfig, + requested_id: Option, + overwrite: bool, + requested_source: Option, + ) -> anyhow::Result { + let mut instance_id = config.get_id(); + if let Some(requested_id) = requested_id { + instance_id = requested_id; + config.set_id(instance_id); + } + let _mutation = self.mutation_lock.lock().await; + let remote_managed = self.hooks.manages_remote_config_instances(); + + let mut replacing = false; + let mut restore_instance = None; + let mut control = if let Some(control) = self.instances.config_control(instance_id) { + let existing_source = self.instances.config_source(instance_id); + let error_message = self + .instances + .network_info(instance_id) + .await + .and_then(|info| info.error_msg) + .unwrap_or_default(); + if !overwrite && error_message.is_empty() { + return Ok(instance_id); + } + self.ensure_overwritable(instance_id, &control, remote_managed) + .await?; + config.set_network_config_source(requested_source.or(existing_source)); + replacing = true; + restore_instance = self + .instances + .config(instance_id) + .map(|config| (config, control.clone())); + control + } else if let Some(config_dir) = self.instances.config_dir() { + config.set_network_config_source(requested_source); + ConfigFileControl::new( + Some(config_dir.join(format!("{instance_id}.toml"))), + ConfigFilePermission::default(), + ) + } else { + config.set_network_config_source(requested_source); + ConfigFileControl::new(None, ConfigFilePermission::default()) + }; + + self.hooks + .pre_run_network_instance(&config) + .await + .map_err(|error| anyhow::anyhow!("pre-run hook failed: {error}"))?; + + if replacing { + self.ensure_overwritable(instance_id, &control, remote_managed) + .await?; + } + + let mut file_cleanup = None; + if !control.is_read_only() + && let Some(path) = control.path.as_deref() + { + let cleanup = match self.storage.read(path).await { + Ok(Some(contents)) => Some(ConfigFileCleanup::Restore { + path: path.to_owned(), + contents, + }), + Ok(None) => Some(ConfigFileCleanup::Remove(path.to_owned())), + Err(error) => { + return Err(anyhow::anyhow!( + "failed to back up config file {} before overwrite: {error}", + path.display() + )); + } + }; + if let Err(error) = self.storage.write(path, config.dump().as_bytes()).await { + tracing::warn!(%error, path = %path.display(), "failed to write config file"); + control.set_read_only(true); + } else { + file_cleanup = cleanup; + } + } + + if replacing + && let Err(error) = self.instances.delete_network_instances([instance_id]).await + { + self.rollback_started_instance(None, file_cleanup, restore_instance) + .await; + return Err(error); + } + + if let Err(error) = self.instances.run_network_instance(config, control) { + self.rollback_started_instance(None, file_cleanup, restore_instance) + .await; + return Err(error); + } + + if let Err(error) = self.hooks.post_run_network_instance(&instance_id).await { + if remote_managed { + self.rollback_started_instance(Some(instance_id), file_cleanup, restore_instance) + .await; + return Err(anyhow::anyhow!("post-run hook failed: {error}")); + } + tracing::warn!(%error, "post-run hook failed"); + } + Ok(instance_id) + } + + pub async fn retain_network_instances( + &self, + retained: Vec, + ) -> anyhow::Result { + let _mutation = self.mutation_lock.lock().await; + self.retain_network_instances_locked(retained).await + } + + /// Resolves retained names and mutates the collection in one transaction. + pub async fn retain_owned_network_instances_by_name( + &self, + retained_names: Vec, + ) -> anyhow::Result { + let _mutation = self.mutation_lock.lock().await; + let retained = self.resolve_instance_ids_by_name(&retained_names)?; + self.retain_network_instances_locked(retained).await + } + + async fn retain_network_instances_locked( + &self, + retained: Vec, + ) -> anyhow::Result { + let before = self.instances.instance_ids(); + if !self.hooks.manages_remote_config_instances() { + let remaining = self.instances.retain_network_instances(&retained).await?; + let remaining_set = remaining.iter().copied().collect::>(); + let removed = before + .into_iter() + .filter(|id| !remaining_set.contains(id)) + .collect::>(); + self.notify_removed_instances(&removed).await?; + return Ok(InstanceMutationResult { + removed_instance_ids: removed, + remaining_instance_ids: remaining, + }); + } + + let mut retained = retained.into_iter().collect::>(); + let mut removed = Vec::new(); + for instance in self.instances.instances() { + let instance_id = instance.instance_id(); + if retained.contains(&instance_id) { + continue; + } + let Some(control) = self.instances.config_control(instance_id) else { + continue; + }; + if self.is_remote_removable(&control).await { + removed.push(instance_id); + } else { + retained.insert(instance_id); + } + } + let remaining = self + .instances + .retain_network_instances(&retained.into_iter().collect::>()) + .await?; + self.notify_removed_instances(&removed).await?; + Ok(InstanceMutationResult { + remaining_instance_ids: remaining, + removed_instance_ids: removed, + }) + } + + pub async fn delete_network_instances( + &self, + requested: Vec, + ) -> anyhow::Result { + let _mutation = self.mutation_lock.lock().await; + let requested = requested.into_iter().collect::>(); + let remote_managed = self.hooks.manages_remote_config_instances(); + let mut removed = Vec::new(); + let mut files = Vec::new(); + for instance in self.instances.instances() { + let instance_id = instance.instance_id(); + if !requested.contains(&instance_id) { + continue; + } + let Some(control) = self.instances.config_control(instance_id) else { + continue; + }; + let removable = if remote_managed { + self.is_remote_removable(&control).await + } else { + control.is_deletable() + }; + if removable { + removed.push(instance_id); + files.extend(control.path); + } + } + let remaining = self + .instances + .delete_network_instances(removed.clone()) + .await?; + self.notify_removed_instances(&removed).await?; + for path in files { + if remote_managed && self.storage.inspect(&path).await.is_read_only() { + continue; + } + if let Err(error) = self.storage.remove(&path).await { + tracing::warn!(%error, path = %path.display(), "failed to remove config file"); + } + } + Ok(InstanceMutationResult { + remaining_instance_ids: remaining, + removed_instance_ids: removed, + }) + } + + /// Starts one caller-owned Instance while preserving its static control. + pub async fn run_owned_network_instance( + &self, + config: TomlConfig, + control: ConfigFileControl, + ) -> anyhow::Result { + let _mutation = self.mutation_lock.lock().await; + let instance_id = config.get_id(); + if self.instances.instance(instance_id).is_some() { + anyhow::bail!("instance {instance_id} already exists"); + } + let instance_name = config.get_inst_name(); + if super::resolve_optional_instance_by_name(self.instances.as_ref(), &instance_name)? + .is_some() + { + anyhow::bail!("instance name {instance_name} already exists"); + } + self.instances.run_network_instance(config, control) + } + + /// Removes caller-owned Instances without applying remote config-file policy. + pub async fn delete_owned_network_instances( + &self, + requested: Vec, + ) -> anyhow::Result { + self.delete_owned_network_instances_selected_by(|| requested) + .await + } + + /// Selects caller-owned Instances and removes them under one lock. + pub async fn delete_owned_network_instances_selected_by( + &self, + select: impl FnOnce() -> Vec, + ) -> anyhow::Result { + let _mutation = self.mutation_lock.lock().await; + let requested = select(); + self.delete_owned_network_instances_locked(requested).await + } + + /// Resolves requested names and removes them in one transaction. + pub async fn delete_owned_network_instances_by_name( + &self, + requested_names: Vec, + ) -> anyhow::Result { + let _mutation = self.mutation_lock.lock().await; + let requested = self.resolve_instance_ids_by_name(&requested_names)?; + self.delete_owned_network_instances_locked(requested).await + } + + async fn delete_owned_network_instances_locked( + &self, + requested: Vec, + ) -> anyhow::Result { + let before = self.instances.instance_ids(); + let remaining = self.instances.delete_network_instances(requested).await?; + let remaining_set = remaining.iter().copied().collect::>(); + let removed = before + .into_iter() + .filter(|id| !remaining_set.contains(id)) + .collect::>(); + self.notify_removed_instances(&removed).await?; + Ok(InstanceMutationResult { + removed_instance_ids: removed, + remaining_instance_ids: remaining, + }) + } + + fn resolve_instance_ids_by_name(&self, names: &[String]) -> anyhow::Result> { + names + .iter() + .map(|name| { + super::resolve_optional_instance_by_name(self.instances.as_ref(), name) + .map(|instance| instance.map(|instance| instance.instance_id())) + }) + .filter_map(Result::transpose) + .collect() + } + + async fn notify_removed_instances(&self, removed: &[uuid::Uuid]) -> anyhow::Result<()> { + if let Err(error) = self.hooks.post_remove_network_instances(removed).await { + if self.hooks.manages_remote_config_instances() { + anyhow::bail!("post-remove hook failed: {error}"); + } + tracing::warn!(%error, "post-remove hook failed"); + } + Ok(()) + } +} + +enum ConfigFileCleanup { + Remove(PathBuf), + Restore { path: PathBuf, contents: Vec }, +} + +/// Protobuf projection over transport-independent process management. +pub struct ProcessManagementRpc +where + F: InstanceFactory, +{ + management: ProcessManagement, +} + +impl Clone for ProcessManagementRpc +where + F: InstanceFactory, +{ + fn clone(&self) -> Self { + Self { + management: self.management.clone(), + } + } +} + +impl ProcessManagementRpc +where + F: InstanceFactory, CreateContext = ()>, + F::Error: std::fmt::Debug + std::fmt::Display + Send + Sync + 'static, + H: CoreInstanceHost, +{ + pub fn new( + instances: Arc>, + hooks: Arc, + storage: Arc, + ) -> Self { + Self { + management: ProcessManagement::new(instances, hooks, storage), + } + } +} + +#[async_trait::async_trait] +impl WebClientService for ProcessManagementRpc +where + F: InstanceFactory, CreateContext = ()>, + F::Error: std::fmt::Debug + std::fmt::Display + Send + Sync + 'static, + H: CoreInstanceHost, +{ + type Controller = BaseController; + + async fn validate_config( + &self, + _: BaseController, + request: ValidateConfigRequest, + ) -> rpc_types::error::Result { + Ok(ValidateConfigResponse { + toml_config: request.config.unwrap_or_default().gen_config()?.dump(), + }) + } + + async fn run_network_instance( + &self, + _: BaseController, + request: RunNetworkInstanceRequest, + ) -> rpc_types::error::Result { + let config = request + .config + .ok_or_else(|| anyhow::anyhow!("config is required"))? + .gen_config()?; + let requested_id = request.inst_id.map(Into::into); + let requested_source = config_source_from_rpc(request.source); + let management = self.management.clone(); + let instance_id = tokio::spawn(async move { + management + .run_network_instance(config, requested_id, request.overwrite, requested_source) + .await + }) + .await + .map_err(|error| anyhow::anyhow!("instance mutation task failed: {error}"))??; + Ok(RunNetworkInstanceResponse { + inst_id: Some(instance_id.into()), + }) + } + + async fn retain_network_instance( + &self, + _: BaseController, + request: RetainNetworkInstanceRequest, + ) -> rpc_types::error::Result { + let retained = request.inst_ids.into_iter().map(Into::into).collect(); + let management = self.management.clone(); + let result = + tokio::spawn(async move { management.retain_network_instances(retained).await }) + .await + .map_err(|error| anyhow::anyhow!("instance mutation task failed: {error}"))??; + Ok(RetainNetworkInstanceResponse { + remain_inst_ids: result + .remaining_instance_ids + .into_iter() + .map(Into::into) + .collect(), + }) + } + + async fn collect_network_info( + &self, + _: BaseController, + request: CollectNetworkInfoRequest, + ) -> rpc_types::error::Result { + let included = request + .inst_ids + .into_iter() + .map(|id| uuid::Uuid::from(id).to_string()) + .collect::>(); + let map = self + .management + .instances + .collect_network_infos() + .await? + .into_iter() + .map(|(id, info)| (id.to_string(), info)) + .filter(|(id, _)| included.is_empty() || included.contains(id)) + .collect(); + Ok(CollectNetworkInfoResponse { + info: Some(NetworkInstanceRunningInfoMap { map }), + }) + } + + async fn list_network_instance( + &self, + _: BaseController, + _: ListNetworkInstanceRequest, + ) -> rpc_types::error::Result { + Ok(ListNetworkInstanceResponse { + inst_ids: self + .management + .instances + .instance_ids() + .into_iter() + .map(Into::into) + .collect(), + }) + } + + async fn delete_network_instance( + &self, + _: BaseController, + request: DeleteNetworkInstanceRequest, + ) -> rpc_types::error::Result { + let requested = request.inst_ids.into_iter().map(Into::into).collect(); + let management = self.management.clone(); + let result = + tokio::spawn(async move { management.delete_network_instances(requested).await }) + .await + .map_err(|error| anyhow::anyhow!("instance mutation task failed: {error}"))??; + Ok(DeleteNetworkInstanceResponse { + remain_inst_ids: result + .remaining_instance_ids + .into_iter() + .map(Into::into) + .collect(), + }) + } + + async fn get_network_instance_config( + &self, + _: BaseController, + request: GetNetworkInstanceConfigRequest, + ) -> rpc_types::error::Result { + let instance_id = request + .inst_id + .ok_or_else(|| anyhow::anyhow!("instance id is required"))? + .into(); + let control = self + .management + .instances + .config_control(instance_id) + .ok_or_else(|| anyhow::anyhow!("instance config control not found"))?; + if control.is_read_only() { + return Err( + anyhow::anyhow!("configuration for instance {instance_id} is read-only").into(), + ); + } + Ok(GetNetworkInstanceConfigResponse { + config: self + .management + .instances + .config(instance_id) + .map(|config| network_config_from_toml(&config)), + source: config_source_to_rpc( + self.management + .instances + .config_source(instance_id) + .unwrap_or(ConfigSource::User), + ), + }) + } + + async fn list_network_instance_meta( + &self, + _: BaseController, + request: ListNetworkInstanceMetaRequest, + ) -> rpc_types::error::Result { + let mut metas = Vec::with_capacity(request.inst_ids.len()); + for instance_id in request.inst_ids.into_iter().map(uuid::Uuid::from) { + let Some(instance) = self.management.instances.instance(instance_id) else { + continue; + }; + let Some(config) = instance.toml_config() else { + continue; + }; + let Some(control) = self.management.instances.config_control(instance_id) else { + continue; + }; + metas.push(NetworkMeta { + inst_id: Some(instance_id.into()), + network_name: config.get_network_identity().network_name, + config_permission: control.permission.into(), + instance_name: instance.instance_name().to_owned(), + source: config_source_to_rpc(config.get_network_config_source()), + }); + } + Ok(ListNetworkInstanceMetaResponse { metas }) + } +} diff --git a/easytier/src/rpc_service/remote_client.rs b/easytier-core/src/management/full/remote_client.rs similarity index 92% rename from easytier/src/rpc_service/remote_client.rs rename to easytier-core/src/management/full/remote_client.rs index 9efd169d..5a911017 100644 --- a/easytier/src/rpc_service/remote_client.rs +++ b/easytier-core/src/management/full/remote_client.rs @@ -1,19 +1,20 @@ use async_trait::async_trait; use uuid::Uuid; -use crate::{ - common::config::ConfigSource, - proto::{ - api::manage::{ - CollectNetworkInfoRequest, CollectNetworkInfoResponse, DeleteNetworkInstanceRequest, - GetNetworkInstanceConfigRequest, ListNetworkInstanceMetaRequest, - ListNetworkInstanceRequest, NetworkConfig, NetworkMeta, RunNetworkInstanceRequest, - ValidateConfigRequest, ValidateConfigResponse, WebClientService, - }, - rpc_types::controller::BaseController, +use easytier_proto::{ + api::manage::{ + CollectNetworkInfoRequest, CollectNetworkInfoResponse, DeleteNetworkInstanceRequest, + GetNetworkInstanceConfigRequest, ListNetworkInstanceMetaRequest, + ListNetworkInstanceRequest, NetworkConfig, NetworkMeta, RunNetworkInstanceRequest, + ValidateConfigRequest, ValidateConfigResponse, WebClientService, }, + rpc_types::controller::BaseController, }; +use crate::config::toml::ConfigSource; + +use super::{config_source_from_rpc, config_source_to_rpc}; + #[async_trait] pub trait RemoteClientManager where @@ -74,7 +75,7 @@ where inst_id: None, config: Some(config.clone()), overwrite: true, - source: source.to_rpc(), + source: config_source_to_rpc(source), }, ) .await?; @@ -138,7 +139,9 @@ where .await .map_err(RemoteClientError::PersistentError)? .iter() - .map(|x| Into::::into(x.get_network_inst_id().to_string())) + .map(|x| { + Into::::into(x.get_network_inst_id().to_string()) + }) .filter(|id| !ret.inst_ids.contains(id)) .collect::>(); @@ -217,7 +220,7 @@ where inst_id: Some(inst_id.into()), config: Some(cfg), overwrite: true, - source: source.to_rpc(), + source: config_source_to_rpc(source), }, ) .await?; @@ -271,7 +274,7 @@ where network_name: network_name.clone(), config_permission: 0, instance_name: network_name, - source: source.to_rpc(), + source: config_source_to_rpc(source), }, ); } @@ -333,7 +336,7 @@ where .await && let Some(config) = resp.config { - let source = if let Some(source) = ConfigSource::from_rpc(resp.source) { + let source = if let Some(source) = config_source_from_rpc(resp.source) { source } else { self.get_storage() @@ -374,7 +377,7 @@ pub enum RemoteClientError { #[error("Not found: {0}")] NotFound(String), #[error(transparent)] - RpcError(#[from] crate::proto::rpc_types::error::Error), + RpcError(#[from] easytier_proto::rpc_types::error::Error), #[error(transparent)] PersistentError(E), #[error("Other error: {0}")] @@ -389,8 +392,8 @@ pub enum ListNetworkProps { #[derive(Debug, serde::Deserialize, serde::Serialize)] pub struct ListNetworkInstanceIdsJsonResp { - running_inst_ids: Vec, - disabled_inst_ids: Vec, + running_inst_ids: Vec, + disabled_inst_ids: Vec, } #[derive(Debug, serde::Deserialize, serde::Serialize)] diff --git a/easytier-core/src/management/full/web_client.rs b/easytier-core/src/management/full/web_client.rs new file mode 100644 index 00000000..bbdd0baf --- /dev/null +++ b/easytier-core/src/management/full/web_client.rs @@ -0,0 +1,403 @@ +use std::sync::{ + Arc, Weak, + atomic::{AtomicBool, Ordering}, +}; + +use easytier_proto::{ + rpc_types::controller::BaseController, + web::{ + DeviceOsInfo, GetFeatureRequest, GetFeatureResponse, HeartbeatRequest, + WebServerServiceClientFactory, + }, +}; +use tokio::{sync::Mutex, task::JoinSet, time::interval}; +use tokio_util::task::AbortOnDropHandle; +use url::Url; + +use crate::{ + connectivity::protocol::raw::TunnelDialer, + instance::{CoreInstance, CoreInstanceHost, manager::InstanceFactory}, + rpc::{bidirect::BidirectRpcManager, service_registry::ServiceRegistry}, + tunnel::{Tunnel, web_security}, +}; + +use super::{ + ConfigFileStorage, DaemonGuard, InstanceManager, InstanceMutationHooks, LoggerControl, + register_management_rpc, +}; + +const RETRY_INTERVAL: std::time::Duration = std::time::Duration::from_secs(1); +const FEATURE_TIMEOUT: std::time::Duration = std::time::Duration::from_secs(3); + +/// Normalized config-server endpoint and authentication token. +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct ConfigServerEndpoint { + connect_url: Url, + token: String, +} + +impl ConfigServerEndpoint { + pub fn parse(input: &str, supports_scheme: impl FnOnce(&Url) -> bool) -> anyhow::Result { + let endpoint = match Url::parse(input) { + Ok(endpoint) => endpoint, + Err(_) => format!("udp://config-server.easytier.cn:22020/{input}") + .parse() + .map_err(|error| anyhow::anyhow!("failed to parse config server URL: {error}"))?, + }; + if !supports_scheme(&endpoint) { + anyhow::bail!("unsupported config server scheme: {}", endpoint.scheme()); + } + + let token = endpoint + .path_segments() + .and_then(|mut segments| segments.next_back()) + .map(|segment| percent_encoding::percent_decode_str(segment).decode_utf8()) + .transpose() + .map_err(|error| anyhow::anyhow!("failed to decode config server token: {error}"))? + .map(|token| token.to_string()) + .unwrap_or_default(); + if token.is_empty() { + anyhow::bail!("empty token"); + } + + let mut connect_url = endpoint; + if !matches!(connect_url.scheme(), "ws" | "wss") { + connect_url.set_path(""); + } + Ok(Self { connect_url, token }) + } + + pub fn connect_url(&self) -> &Url { + &self.connect_url + } + + pub fn token(&self) -> &str { + &self.token + } +} + +pub struct WebClientConfig { + pub token: String, + pub machine_id: uuid::Uuid, + pub hostname: String, + pub device_os: DeviceOsInfo, + pub easytier_version: String, + pub secure_mode: bool, +} + +struct WebClientController +where + F: InstanceFactory, +{ + config: WebClientConfig, + instances: Arc>, + hooks: Arc, + storage: Arc, + logger: Arc, +} + +impl WebClientController +where + F: InstanceFactory, CreateContext = ()>, + F::Error: std::fmt::Debug + std::fmt::Display + Send + Sync + 'static, + H: CoreInstanceHost, +{ + fn register_management_entry(&self, registry: &ServiceRegistry) { + register_management_rpc( + self.instances.clone(), + registry, + self.hooks.clone(), + self.storage.clone(), + self.logger.clone(), + ); + } +} + +/// Portable config-server client. Hosts only supply identity and adapters. +pub struct WebClient +where + F: InstanceFactory, +{ + _controller: Arc>, + _tasks: AbortOnDropHandle<()>, + _manager_guard: DaemonGuard, + connected: Arc, +} + +impl WebClient +where + F: InstanceFactory, CreateContext = ()>, + F::Error: std::fmt::Debug + std::fmt::Display + Send + Sync + 'static, + H: CoreInstanceHost, +{ + pub fn new( + connector: T, + config: WebClientConfig, + instances: Arc>, + hooks: Arc, + storage: Arc, + logger: Arc, + ) -> Self { + let manager_guard = instances.register_daemon(); + let controller = Arc::new(WebClientController { + config, + instances, + hooks, + storage, + logger, + }); + let connected = Arc::new(AtomicBool::new(false)); + let tasks = AbortOnDropHandle::new(tokio::spawn(Self::routine( + controller.clone(), + connected.clone(), + Box::new(connector), + ))); + + Self { + _controller: controller, + _tasks: tasks, + _manager_guard: manager_guard, + connected, + } + } + + async fn routine( + controller: Arc>, + connected: Arc, + connector: Box, + ) { + loop { + let connection = match connector.connect().await { + Ok(connection) => connection, + Err(error) => { + tracing::warn!(%error, "failed to connect to config server; retrying"); + tokio::time::sleep(RETRY_INTERVAL).await; + continue; + } + }; + + connected.store(true, Ordering::Release); + tracing::info!(?connection, "connected to config server"); + let mut session = WebClientSession::new(connection, controller.clone()); + let support_encryption = + match tokio::time::timeout(FEATURE_TIMEOUT, session.get_feature()).await { + Ok(Ok(feature)) => feature.support_encryption, + Ok(Err(error)) => { + tracing::warn!(%error, "GetFeature RPC failed; using legacy tunnel"); + false + } + Err(_) => { + tracing::warn!("GetFeature RPC timed out; using legacy tunnel"); + false + } + }; + + if support_encryption && web_security::web_secure_tunnel_supported() { + drop(session); + let connection = match connector.connect().await { + Ok(connection) => connection, + Err(error) => { + connected.store(false, Ordering::Release); + tracing::warn!(%error, "failed to reconnect secure config-server tunnel"); + tokio::time::sleep(RETRY_INTERVAL).await; + continue; + } + }; + let connection = match web_security::upgrade_client_tunnel(connection).await { + Ok(connection) => connection, + Err(error) => { + connected.store(false, Ordering::Release); + tracing::warn!(%error, "config-server secure handshake failed"); + tokio::time::sleep(RETRY_INTERVAL).await; + continue; + } + }; + let mut session = WebClientSession::new(connection, controller.clone()); + session.start_heartbeat().await; + session.wait().await; + connected.store(false, Ordering::Release); + continue; + } + + if support_encryption { + if controller.config.secure_mode { + connected.store(false, Ordering::Release); + tracing::warn!( + "secure mode requires web secure-tunnel support in the local build" + ); + tokio::time::sleep(RETRY_INTERVAL).await; + continue; + } + tracing::warn!( + "server supports encryption but the local build is using a legacy tunnel" + ); + } + if controller.config.secure_mode { + connected.store(false, Ordering::Release); + tracing::warn!("secure mode requires config-server encryption support"); + tokio::time::sleep(RETRY_INTERVAL).await; + continue; + } + + session.start_heartbeat().await; + session.wait().await; + connected.store(false, Ordering::Release); + } + } + + pub fn is_connected(&self) -> bool { + self.connected.load(Ordering::Acquire) + } +} + +struct WebClientSession +where + F: InstanceFactory, +{ + rpc: BidirectRpcManager, + controller: Arc>, + heartbeat_started: AtomicBool, + tasks: Mutex>, +} + +impl WebClientSession +where + F: InstanceFactory, CreateContext = ()>, + F::Error: std::fmt::Debug + std::fmt::Display + Send + Sync + 'static, + H: CoreInstanceHost, +{ + fn new(tunnel: Box, controller: Arc>) -> Self { + let rpc = BidirectRpcManager::new(); + rpc.run_with_tunnel(tunnel); + controller.register_management_entry(rpc.rpc_server().registry()); + Self { + rpc, + controller, + heartbeat_started: AtomicBool::new(false), + tasks: Mutex::new(JoinSet::new()), + } + } + + pub async fn start_heartbeat(&self) { + if self.heartbeat_started.swap(true, Ordering::AcqRel) { + return; + } + let mut tasks = self.tasks.lock().await; + Self::heartbeat_routine(&self.rpc, Arc::downgrade(&self.controller), &mut tasks); + } + + fn heartbeat_routine( + rpc: &BidirectRpcManager, + controller: Weak>, + tasks: &mut JoinSet<()>, + ) { + let controller = controller.upgrade().expect("web client controller"); + let machine_id = controller.config.machine_id; + let session_id = uuid::Uuid::new_v4(); + let token = controller.config.token.clone(); + let hostname = controller.config.hostname.clone(); + let device_os = controller.config.device_os.clone(); + let easytier_version = controller.config.easytier_version.clone(); + let controller = Arc::downgrade(&controller); + let client = rpc + .rpc_client() + .scoped_client::>(1, 1, String::new()); + let mut tick = interval(std::time::Duration::from_secs(1)); + + tasks.spawn(async move { + loop { + tick.tick().await; + let Some(controller) = controller.upgrade() else { + break; + }; + let request = HeartbeatRequest { + machine_id: Some(machine_id.into()), + inst_id: Some(session_id.into()), + user_token: token.clone(), + easytier_version: easytier_version.clone(), + hostname: hostname.clone(), + report_time: chrono::Local::now().to_rfc3339(), + device_os: Some(device_os.clone()), + support_config_source: true, + running_network_instances: controller + .instances + .instance_ids() + .into_iter() + .map(Into::into) + .collect(), + }; + + match client.heartbeat(BaseController::default(), request).await { + Ok(response) => { + tracing::debug!(?response, "config-server heartbeat response"); + } + Err(error) => { + tracing::error!(?error, "config-server heartbeat failed"); + break; + } + } + } + }); + } + + async fn wait_routines(&self) { + self.tasks.lock().await.join_next().await; + self.tasks.lock().await.abort_all(); + } + + async fn wait(&mut self) { + tokio::select! { + _ = self.rpc.wait() => {} + _ = self.wait_routines() => {} + } + } + + async fn get_feature( + &self, + ) -> Result { + let client = self + .rpc + .rpc_client() + .scoped_client::>(1, 1, String::new()); + client + .get_feature(BaseController::default(), GetFeatureRequest {}) + .await + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn endpoint_normalizes_shorthand_and_non_websocket_paths() { + let endpoint = ConfigServerEndpoint::parse("team%2Ftoken", |_| true).unwrap(); + assert_eq!(endpoint.token(), "team/token"); + assert_eq!( + endpoint.connect_url().as_str(), + "udp://config-server.easytier.cn:22020" + ); + } + + #[test] + fn endpoint_preserves_websocket_path_and_validates_scheme() { + let endpoint = + ConfigServerEndpoint::parse("wss://example.com/team", |url| url.scheme() == "wss") + .unwrap(); + assert_eq!(endpoint.token(), "team"); + assert_eq!(endpoint.connect_url().as_str(), "wss://example.com/team"); + + let error = + ConfigServerEndpoint::parse("unknown://example.com/team", |_| false).unwrap_err(); + assert!( + error + .to_string() + .contains("unsupported config server scheme") + ); + } + + #[test] + fn endpoint_rejects_an_empty_token() { + assert!(ConfigServerEndpoint::parse("udp://example.com", |_| true).is_err()); + } +} diff --git a/easytier-core/src/management/instance_rpc/full.rs b/easytier-core/src/management/instance_rpc/full.rs new file mode 100644 index 00000000..eb9bc7e7 --- /dev/null +++ b/easytier-core/src/management/instance_rpc/full.rs @@ -0,0 +1,406 @@ +use std::{sync::Arc, time::Duration}; + +use easytier_proto::{ + api::{ + config::{ + ConfigRpc, GetConfigRequest, GetConfigResponse, PatchConfigRequest, PatchConfigResponse, + }, + instance::{ + AclManageRpc, ConnectorManageRpc, CredentialInfo, CredentialManageRpc, + GenerateCredentialRequest, GenerateCredentialResponse, GetAclStatsRequest, + GetAclStatsResponse, GetPrometheusStatsRequest, GetPrometheusStatsResponse, + GetStatsRequest, GetStatsResponse, GetVpnPortalInfoRequest, GetVpnPortalInfoResponse, + GetWhitelistRequest, GetWhitelistResponse, ListCredentialsRequest, + ListCredentialsResponse, ListMappedListenerRequest, ListMappedListenerResponse, + ListPortForwardRequest, ListPortForwardResponse, MappedListener, + MappedListenerManageRpc, MetricSnapshot, PeerManageRpc, PortForwardManageRpc, + RevokeCredentialRequest, RevokeCredentialResponse, StatsRpc, VpnPortalInfo, + VpnPortalRpc, + }, + }, + common::PortForwardConfigPb, + peer_rpc::{ + GetGlobalPeerMapRequest, GetGlobalPeerMapResponse, PeerCenterRpc, ReportPeersRequest, + ReportPeersResponse, + }, + rpc_types::{self, controller::BaseController}, +}; + +use crate::{ + config::{api::network_config_from_toml, toml::ConfigLoader as _}, + instance::{ + CoreInstance, CoreInstanceHost, + manager::{InstanceFactory, InstanceManager}, + }, + peers::credential_manager::{CredentialCreateOptions, CredentialInfo as CoreCredentialInfo}, +}; + +use super::InstanceManagementRpc; +use crate::management::{ + full::{apply_config_patch, packet_proxy}, + resolve_instance, +}; + +/// Dispatches the JSON form of an Instance-targeted management RPC without +/// introducing a second, Host-owned set of service implementations. +pub async fn call_instance_json_rpc( + manager: &Arc>, + service_name: &str, + method_name: &str, + domain_name: Option<&str>, + payload: serde_json::Value, +) -> rpc_types::error::Result +where + F: InstanceFactory>, + H: CoreInstanceHost, +{ + let payload = + match packet_proxy::call_json(manager, service_name, method_name, domain_name, payload) + .await + { + Ok(response) => return response, + Err(payload) => payload, + }; + let ctrl = BaseController::default(); + let rpc = InstanceManagementRpc::::new(manager.clone()); + + match service_name { + "api.instance.PeerManageRpcService" => { + PeerManageRpc::json_call_method(&rpc, ctrl, method_name, payload).await + } + "api.instance.PeerCenterManageRpcService" => { + PeerCenterRpc::json_call_method(&rpc, ctrl, method_name, payload).await + } + "api.instance.ConnectorManageRpcService" => { + ConnectorManageRpc::json_call_method(&rpc, ctrl, method_name, payload).await + } + "api.instance.MappedListenerManageRpcService" => { + MappedListenerManageRpc::json_call_method(&rpc, ctrl, method_name, payload).await + } + "api.instance.VpnPortalRpcService" => { + VpnPortalRpc::json_call_method(&rpc, ctrl, method_name, payload).await + } + "api.instance.AclManageRpcService" => { + AclManageRpc::json_call_method(&rpc, ctrl, method_name, payload).await + } + "api.instance.PortForwardManageRpcService" => { + PortForwardManageRpc::json_call_method(&rpc, ctrl, method_name, payload).await + } + "api.instance.StatsRpcService" => { + StatsRpc::json_call_method(&rpc, ctrl, method_name, payload).await + } + "api.instance.CredentialManageRpcService" => { + CredentialManageRpc::json_call_method(&rpc, ctrl, method_name, payload).await + } + "api.config.ConfigRpcService" => { + ConfigRpc::json_call_method(&rpc, ctrl, method_name, payload).await + } + _ => Err(rpc_types::error::Error::InvalidServiceKey( + service_name.to_owned(), + service_name.to_owned(), + )), + } +} + +fn credential_info_to_api(info: CoreCredentialInfo) -> CredentialInfo { + CredentialInfo { + credential_id: info.credential_id, + groups: info.groups, + allow_relay: info.allow_relay, + expiry_unix: info.expiry_unix, + allowed_proxy_cidrs: info.allowed_proxy_cidrs, + reusable: info.reusable, + } +} + +#[async_trait::async_trait] +impl MappedListenerManageRpc for InstanceManagementRpc +where + F: InstanceFactory>, + H: CoreInstanceHost, +{ + type Controller = BaseController; + + async fn list_mapped_listener( + &self, + _: BaseController, + request: ListMappedListenerRequest, + ) -> rpc_types::error::Result { + let config = self + .instance(request.instance.as_ref())? + .toml_config() + .ok_or_else(|| anyhow::anyhow!("shared TOML configuration is not available"))?; + Ok(ListMappedListenerResponse { + mappedlisteners: config + .get_mapped_listeners() + .into_iter() + .map(|url| MappedListener { + url: Some(url.into()), + }) + .collect(), + }) + } +} + +#[async_trait::async_trait] +impl VpnPortalRpc for InstanceManagementRpc +where + F: InstanceFactory>, + H: CoreInstanceHost, +{ + type Controller = BaseController; + + async fn get_vpn_portal_info( + &self, + _: BaseController, + request: GetVpnPortalInfoRequest, + ) -> rpc_types::error::Result { + let info = self + .instance(request.instance.as_ref())? + .vpn_portal_info() + .await; + Ok(GetVpnPortalInfoResponse { + vpn_portal_info: Some(VpnPortalInfo { + vpn_type: info.vpn_type, + client_config: info.client_config, + connected_clients: info.connected_clients, + }), + }) + } +} + +#[async_trait::async_trait] +impl AclManageRpc for InstanceManagementRpc +where + F: InstanceFactory>, + H: CoreInstanceHost, +{ + type Controller = BaseController; + + async fn get_acl_stats( + &self, + _: BaseController, + request: GetAclStatsRequest, + ) -> rpc_types::error::Result { + Ok(GetAclStatsResponse { + acl_stats: Some(self.instance(request.instance.as_ref())?.acl_stats()), + }) + } + + async fn get_whitelist( + &self, + _: BaseController, + request: GetWhitelistRequest, + ) -> rpc_types::error::Result { + let whitelist = self + .instance(request.instance.as_ref())? + .acl_whitelist_snapshot(); + Ok(GetWhitelistResponse { + tcp_ports: whitelist.tcp_ports, + udp_ports: whitelist.udp_ports, + }) + } +} + +#[async_trait::async_trait] +impl PortForwardManageRpc for InstanceManagementRpc +where + F: InstanceFactory>, + H: CoreInstanceHost, +{ + type Controller = BaseController; + + async fn list_port_forward( + &self, + _: BaseController, + request: ListPortForwardRequest, + ) -> rpc_types::error::Result { + let config = self + .instance(request.instance.as_ref())? + .toml_config() + .ok_or_else(|| anyhow::anyhow!("shared TOML configuration is not available"))?; + Ok(ListPortForwardResponse { + cfgs: config + .get_port_forwards() + .into_iter() + .map(PortForwardConfigPb::from) + .collect(), + }) + } +} + +#[async_trait::async_trait] +impl StatsRpc for InstanceManagementRpc +where + F: InstanceFactory>, + H: CoreInstanceHost, +{ + type Controller = BaseController; + + async fn get_stats( + &self, + _: BaseController, + request: GetStatsRequest, + ) -> rpc_types::error::Result { + Ok(GetStatsResponse { + metrics: self + .instance(request.instance.as_ref())? + .metric_snapshots() + .into_iter() + .map(|snapshot| MetricSnapshot { + name: snapshot.name_str(), + value: snapshot.value, + labels: snapshot + .labels + .labels() + .iter() + .map(|label| (label.key.clone(), label.value.clone())) + .collect(), + }) + .collect(), + }) + } + + async fn get_prometheus_stats( + &self, + _: BaseController, + request: GetPrometheusStatsRequest, + ) -> rpc_types::error::Result { + Ok(GetPrometheusStatsResponse { + prometheus_text: self + .instance(request.instance.as_ref())? + .prometheus_metrics(), + }) + } +} + +#[async_trait::async_trait] +impl CredentialManageRpc for InstanceManagementRpc +where + F: InstanceFactory>, + H: CoreInstanceHost, +{ + type Controller = BaseController; + + async fn generate_credential( + &self, + _: BaseController, + request: GenerateCredentialRequest, + ) -> rpc_types::error::Result { + if request.ttl_seconds <= 0 { + return Err(anyhow::anyhow!("ttl_seconds must be positive").into()); + } + let generated = self + .instance(request.instance.as_ref())? + .generate_credential(CredentialCreateOptions { + groups: request.groups, + allow_relay: request.allow_relay, + allowed_proxy_cidrs: request.allowed_proxy_cidrs, + ttl: Duration::from_secs(request.ttl_seconds as u64), + credential_id: request.credential_id, + reusable: request.reusable.unwrap_or(true), + })?; + Ok(GenerateCredentialResponse { + credential_id: generated.credential_id, + credential_secret: generated.secret, + }) + } + + async fn revoke_credential( + &self, + _: BaseController, + request: RevokeCredentialRequest, + ) -> rpc_types::error::Result { + Ok(RevokeCredentialResponse { + success: self + .instance(request.instance.as_ref())? + .revoke_credential(&request.credential_id)?, + }) + } + + async fn list_credentials( + &self, + _: BaseController, + request: ListCredentialsRequest, + ) -> rpc_types::error::Result { + Ok(ListCredentialsResponse { + credentials: self + .instance(request.instance.as_ref())? + .credential_snapshots() + .into_iter() + .map(credential_info_to_api) + .collect(), + }) + } +} + +#[async_trait::async_trait] +impl PeerCenterRpc for InstanceManagementRpc +where + F: InstanceFactory>, + H: CoreInstanceHost, +{ + type Controller = BaseController; + + async fn get_global_peer_map( + &self, + _: BaseController, + _: GetGlobalPeerMapRequest, + ) -> rpc_types::error::Result { + let instance = resolve_instance(&self.manager, None).map_err(|error| { + if error.to_string().contains("please specify the instance ID") { + anyhow::anyhow!( + "PeerCenter management RPC cannot select an instance automatically when \ + multiple instances are running; please use an API that allows specifying \ + an instance identifier." + ) + } else { + error + } + })?; + Ok(instance.global_peer_map_snapshot()) + } + + async fn report_peers( + &self, + _: BaseController, + _: ReportPeersRequest, + ) -> rpc_types::error::Result { + Err(anyhow::anyhow!("not implemented for management API").into()) + } +} + +#[async_trait::async_trait] +impl ConfigRpc for InstanceManagementRpc +where + F: InstanceFactory>, + H: CoreInstanceHost, +{ + type Controller = BaseController; + + async fn patch_config( + &self, + _: BaseController, + request: PatchConfigRequest, + ) -> rpc_types::error::Result { + let instance = self.instance(request.instance.as_ref())?; + if let Some(patch) = request.patch { + apply_config_patch(&instance, patch).await?; + } + Ok(PatchConfigResponse::default()) + } + + async fn get_config( + &self, + _: BaseController, + request: GetConfigRequest, + ) -> rpc_types::error::Result { + let config = self + .instance(request.instance.as_ref())? + .toml_config() + .ok_or_else(|| anyhow::anyhow!("shared TOML configuration is not available"))?; + Ok(GetConfigResponse { + config: Some(network_config_from_toml(&config)), + }) + } +} diff --git a/easytier-core/src/management/instance_rpc/mod.rs b/easytier-core/src/management/instance_rpc/mod.rs new file mode 100644 index 00000000..b7c99112 --- /dev/null +++ b/easytier-core/src/management/instance_rpc/mod.rs @@ -0,0 +1,327 @@ +use std::sync::Arc; + +use easytier_proto::{ + api::instance::{ + Connector, ConnectorManageRpc, ConnectorStatus, DumpRouteRequest, DumpRouteResponse, + ForeignNetworkEntryPb, GetForeignNetworkSummaryRequest, GetForeignNetworkSummaryResponse, + ListConnectorRequest, ListConnectorResponse, ListForeignNetworkRequest, + ListForeignNetworkResponse, ListGlobalForeignNetworkRequest, + ListGlobalForeignNetworkResponse, ListPeerRequest, ListPeerResponse, + ListPublicIpv6InfoRequest, ListPublicIpv6InfoResponse, ListRouteRequest, ListRouteResponse, + NodeInfo, PeerInfo, PeerManageRpc, ShowNodeInfoRequest, ShowNodeInfoResponse, + TrustedKeyInfoPb, TrustedKeySourcePb, + list_global_foreign_network_response::OneForeignNetwork, + }, + rpc_types::{self, controller::BaseController}, +}; + +use crate::{ + config::{IpPrefix, ProxyNetworkConfig}, + connectivity::manual::{ManualConnectorSnapshot, ManualConnectorStatus}, + instance::{ + CoreInstance, CoreInstanceHost, + manager::{InstanceFactory, InstanceManager}, + }, + peers::{context::TrustedKeySource, foreign_network::ForeignNetworkEntryInfo}, +}; + +use super::resolve_instance; + +#[cfg(feature = "management")] +pub(super) mod full; +#[cfg(all(feature = "management", feature = "proxy-packet"))] +pub(super) mod packet_proxy; +mod projection; + +/// One process-level implementation for Instance-targeted management RPC. +pub struct InstanceManagementRpc +where + F: InstanceFactory, +{ + manager: Arc>, +} + +impl Clone for InstanceManagementRpc +where + F: InstanceFactory, +{ + fn clone(&self) -> Self { + Self { + manager: self.manager.clone(), + } + } +} + +impl InstanceManagementRpc +where + F: InstanceFactory>, + H: CoreInstanceHost, +{ + pub fn new(manager: Arc>) -> Self { + Self { manager } + } + + fn instance( + &self, + identifier: Option<&easytier_proto::api::instance::InstanceIdentifier>, + ) -> rpc_types::error::Result>> { + resolve_instance(&self.manager, identifier).map_err(Into::into) + } +} + +fn foreign_network_info_to_api(info: ForeignNetworkEntryInfo) -> ForeignNetworkEntryPb { + ForeignNetworkEntryPb { + network_secret_digest: info.network_secret_digest, + my_peer_id_for_this_network: info.my_peer_id_for_this_network, + peers: info + .peers + .into_iter() + .map(|peer| PeerInfo { + peer_id: peer.peer_id, + conns: peer.conns.into_iter().map(Into::into).collect(), + ..Default::default() + }) + .collect(), + trusted_keys: info + .trusted_keys + .into_iter() + .map(|key| TrustedKeyInfoPb { + pubkey: key.pubkey, + source: match key.source { + TrustedKeySource::OspfNode => TrustedKeySourcePb::OspfNode.into(), + TrustedKeySource::OspfCredential => TrustedKeySourcePb::OspfCredential.into(), + }, + expiry_unix: key.expiry_unix, + }) + .collect(), + } +} + +fn format_prefix(prefix: &IpPrefix) -> String { + format!("{}/{}", prefix.address, prefix.prefix_len) +} + +fn format_proxy_network(proxy: ProxyNetworkConfig) -> String { + let real = format_prefix(&proxy.real); + match proxy.mapped { + Some(mapped) => format!("{}->{}", real, format_prefix(&mapped)), + None => real, + } +} + +fn connector_snapshots_to_api(snapshots: Vec) -> Vec { + let mut connectors = Vec::with_capacity(snapshots.len()); + for connector in snapshots { + let status = match connector.status { + ManualConnectorStatus::Connected => ConnectorStatus::Connected, + ManualConnectorStatus::Disconnected => ConnectorStatus::Disconnected, + ManualConnectorStatus::Connecting => ConnectorStatus::Connecting, + }; + connectors.insert( + 0, + Connector { + url: Some(connector.url.into()), + status: status.into(), + }, + ); + } + connectors +} + +#[async_trait::async_trait] +impl PeerManageRpc for InstanceManagementRpc +where + F: InstanceFactory>, + H: CoreInstanceHost, +{ + type Controller = BaseController; + + async fn list_peer( + &self, + _: BaseController, + request: ListPeerRequest, + ) -> rpc_types::error::Result { + let peer_infos = self + .instance(request.instance.as_ref())? + .peer_snapshots() + .await + .into_iter() + .map(|snapshot| PeerInfo { + peer_id: snapshot.peer_id, + default_conn_id: snapshot.default_conn_id.map(Into::into), + directly_connected_conns: snapshot + .directly_connected_conns + .into_iter() + .map(Into::into) + .collect(), + conns: snapshot.conns.into_iter().map(Into::into).collect(), + }) + .collect(); + Ok(ListPeerResponse { + peer_infos, + ..Default::default() + }) + } + + async fn list_public_ipv6_info( + &self, + _: BaseController, + request: ListPublicIpv6InfoRequest, + ) -> rpc_types::error::Result { + Ok(self + .instance(request.instance.as_ref())? + .local_public_ipv6_info() + .await + .into()) + } + + async fn list_route( + &self, + _: BaseController, + request: ListRouteRequest, + ) -> rpc_types::error::Result { + Ok(ListRouteResponse { + routes: self + .instance(request.instance.as_ref())? + .route_snapshots() + .await + .into_iter() + .map(Into::into) + .collect(), + }) + } + + async fn dump_route( + &self, + _: BaseController, + request: DumpRouteRequest, + ) -> rpc_types::error::Result { + Ok(DumpRouteResponse { + result: self.instance(request.instance.as_ref())?.dump_route().await, + }) + } + + async fn list_foreign_network( + &self, + _: BaseController, + request: ListForeignNetworkRequest, + ) -> rpc_types::error::Result { + Ok(ListForeignNetworkResponse { + foreign_networks: self + .instance(request.instance.as_ref())? + .foreign_network_snapshots(request.include_trusted_keys) + .await + .into_iter() + .map(|(network_name, info)| (network_name, foreign_network_info_to_api(info))) + .collect(), + }) + } + + async fn list_global_foreign_network( + &self, + _: BaseController, + request: ListGlobalForeignNetworkRequest, + ) -> rpc_types::error::Result { + let mut response = ListGlobalForeignNetworkResponse::default(); + let route_infos = self + .instance(request.instance.as_ref())? + .foreign_network_route_infos() + .await; + for info in &route_infos.infos { + let Some(key) = info.key.as_ref() else { + continue; + }; + let Some(route_info) = info.value.as_ref() else { + continue; + }; + response + .foreign_networks + .entry(key.peer_id) + .or_default() + .foreign_networks + .push(OneForeignNetwork { + network_name: key.network_name.clone(), + peer_ids: route_info.foreign_peer_ids.clone(), + last_updated: match route_info.last_update.as_ref() { + Some(last_update) => projection::format_last_update(last_update)?, + None => String::new(), + }, + version: route_info.version, + }); + } + Ok(response) + } + + async fn get_foreign_network_summary( + &self, + _: BaseController, + request: GetForeignNetworkSummaryRequest, + ) -> rpc_types::error::Result { + Ok(GetForeignNetworkSummaryResponse { + summary: Some( + self.instance(request.instance.as_ref())? + .foreign_network_route_summary() + .await, + ), + }) + } + + async fn show_node_info( + &self, + _: BaseController, + request: ShowNodeInfoRequest, + ) -> rpc_types::error::Result { + let instance = self.instance(request.instance.as_ref())?; + let config = projection::node_config(instance.as_ref())?; + let snapshot = instance.node_snapshot().await; + Ok(ShowNodeInfoResponse { + node_info: Some(NodeInfo { + peer_id: snapshot.peer_id, + ipv4_addr: snapshot + .ipv4_addr + .map(|addr| addr.to_string()) + .unwrap_or_default(), + proxy_cidrs: snapshot + .proxy_networks + .into_iter() + .map(format_proxy_network) + .collect(), + hostname: snapshot.hostname, + stun_info: Some(snapshot.stun_info), + inst_id: snapshot.instance_id.to_string(), + listeners: snapshot + .listeners + .into_iter() + .map(|listener| listener.to_string()) + .collect(), + config, + version: snapshot.version, + feature_flag: Some(snapshot.feature_flags), + ip_list: Some(snapshot.ip_list), + public_ipv6_addr: snapshot.public_ipv6_addr.map(Into::into), + ipv6_public_addr_prefix: snapshot.ipv6_public_addr_prefix.map(Into::into), + }), + }) + } +} + +#[async_trait::async_trait] +impl ConnectorManageRpc for InstanceManagementRpc +where + F: InstanceFactory>, + H: CoreInstanceHost, +{ + type Controller = BaseController; + + async fn list_connector( + &self, + _: BaseController, + request: ListConnectorRequest, + ) -> rpc_types::error::Result { + Ok(ListConnectorResponse { + connectors: connector_snapshots_to_api( + self.instance(request.instance.as_ref())?.list_connectors(), + ), + }) + } +} diff --git a/easytier-core/src/management/instance_rpc/packet_proxy.rs b/easytier-core/src/management/instance_rpc/packet_proxy.rs new file mode 100644 index 00000000..c6cee343 --- /dev/null +++ b/easytier-core/src/management/instance_rpc/packet_proxy.rs @@ -0,0 +1,135 @@ +use std::sync::Arc; + +use easytier_proto::{ + api::instance::{ + ListTcpProxyEntryRequest, ListTcpProxyEntryResponse, TcpProxyEntry, TcpProxyEntryState, + TcpProxyEntryTransportType, TcpProxyRpc, + }, + rpc_types::{self, controller::BaseController}, +}; + +use crate::{ + gateway::proxy::{ + tcp_proxy_engine::{TcpNatEntrySnapshot, TcpNatEntryState as CoreTcpNatEntryState}, + wrapped_transport::{WrappedTransportKind, WrappedTransportRole}, + }, + instance::{ + CoreInstance, CoreInstanceHost, + manager::{InstanceFactory, InstanceManager}, + }, +}; + +use super::super::resolve_instance; + +#[derive(Clone, Copy)] +enum TcpProxySource { + Tcp, + Wrapped(WrappedTransportKind, WrappedTransportRole), +} + +pub(crate) struct TcpProxyManagementRpc +where + F: InstanceFactory, +{ + manager: Arc>, + source: TcpProxySource, +} + +impl Clone for TcpProxyManagementRpc +where + F: InstanceFactory, +{ + fn clone(&self) -> Self { + Self { + manager: self.manager.clone(), + source: self.source, + } + } +} + +impl TcpProxyManagementRpc +where + F: InstanceFactory>, + H: CoreInstanceHost, +{ + pub(crate) fn tcp(manager: Arc>) -> Self { + Self { + manager, + source: TcpProxySource::Tcp, + } + } + + pub(crate) fn wrapped( + manager: Arc>, + transport: WrappedTransportKind, + role: WrappedTransportRole, + ) -> Self { + Self { + manager, + source: TcpProxySource::Wrapped(transport, role), + } + } +} + +fn tcp_entry_snapshot_to_api( + entry: TcpNatEntrySnapshot, + transport_type: TcpProxyEntryTransportType, +) -> TcpProxyEntry { + TcpProxyEntry { + src: Some(entry.src.into()), + dst: Some(entry.dst.into()), + start_time: entry.start_time, + state: match entry.state { + CoreTcpNatEntryState::SynReceived => TcpProxyEntryState::SynReceived, + CoreTcpNatEntryState::ConnectingDst => TcpProxyEntryState::ConnectingDst, + CoreTcpNatEntryState::Connected => TcpProxyEntryState::Connected, + CoreTcpNatEntryState::ClosingSrc => TcpProxyEntryState::ClosingSrc, + CoreTcpNatEntryState::ClosingDst => TcpProxyEntryState::ClosingDst, + CoreTcpNatEntryState::Closed => TcpProxyEntryState::Closed, + } + .into(), + transport_type: transport_type.into(), + } +} + +#[async_trait::async_trait] +impl TcpProxyRpc for TcpProxyManagementRpc +where + F: InstanceFactory>, + H: CoreInstanceHost, +{ + type Controller = BaseController; + + async fn list_tcp_proxy_entry( + &self, + _: BaseController, + request: ListTcpProxyEntryRequest, + ) -> rpc_types::error::Result { + let instance = resolve_instance(&self.manager, request.instance.as_ref())?; + let (snapshots, transport_type) = match self.source { + TcpProxySource::Tcp => ( + instance.tcp_proxy_entry_snapshots(), + TcpProxyEntryTransportType::Tcp, + ), + TcpProxySource::Wrapped(transport, role) => { + if !instance.wrapped_transport_is_started(transport, role) { + return Err(anyhow::anyhow!("wrapped TCP proxy is not available").into()); + } + let transport_type = match transport { + WrappedTransportKind::Kcp => TcpProxyEntryTransportType::Kcp, + WrappedTransportKind::Quic => TcpProxyEntryTransportType::Quic, + }; + ( + instance.wrapped_tcp_proxy_entry_snapshots(transport, role), + transport_type, + ) + } + }; + Ok(ListTcpProxyEntryResponse { + entries: snapshots + .into_iter() + .map(|entry| tcp_entry_snapshot_to_api(entry, transport_type)) + .collect(), + }) + } +} diff --git a/easytier-core/src/management/instance_rpc/projection.rs b/easytier-core/src/management/instance_rpc/projection.rs new file mode 100644 index 00000000..ec28c333 --- /dev/null +++ b/easytier-core/src/management/instance_rpc/projection.rs @@ -0,0 +1,39 @@ +use crate::instance::{CoreInstance, CoreInstanceHost}; + +#[cfg(not(feature = "management"))] +pub(super) fn format_last_update( + last_update: &easytier_proto::common::RuntimeTimestamp, +) -> anyhow::Result { + let last_update = last_update.normalized(); + let date_time = chrono::DateTime::from_timestamp(last_update.seconds, last_update.nanos as u32) + .ok_or_else(|| anyhow::anyhow!("invalid protobuf timestamp"))?; + Ok(format!("\"{date_time:?}\"")) +} + +#[cfg(feature = "management")] +pub(super) fn format_last_update( + last_update: &easytier_proto::common::RuntimeTimestamp, +) -> anyhow::Result { + serde_json::to_string(last_update).map_err(anyhow::Error::from) +} + +#[cfg(not(feature = "management"))] +pub(super) fn node_config(_instance: &CoreInstance) -> anyhow::Result +where + H: CoreInstanceHost, +{ + Ok(String::new()) +} + +#[cfg(feature = "management")] +pub(super) fn node_config(instance: &CoreInstance) -> anyhow::Result +where + H: CoreInstanceHost, +{ + use crate::config::toml::ConfigLoader as _; + + instance + .toml_config() + .ok_or_else(|| anyhow::anyhow!("shared TOML configuration is not available")) + .map(|config| config.dump()) +} diff --git a/easytier-core/src/management/mod.rs b/easytier-core/src/management/mod.rs new file mode 100644 index 00000000..9cad5205 --- /dev/null +++ b/easytier-core/src/management/mod.rs @@ -0,0 +1,52 @@ +//! Process-level management over the canonical Instance collection. + +#[cfg(feature = "management")] +mod full; +mod instance_rpc; +mod rpc_server_hook; +mod selector; +mod server; + +use std::sync::Arc; + +use crate::{ + instance::{CoreInstance, CoreInstanceHost, manager::InstanceFactory}, + rpc::service_registry::ServiceRegistry, +}; +use easytier_proto::api::instance::{ConnectorManageRpcServer, PeerManageRpcServer}; + +pub use crate::instance::manager::{ + ConfigFileControl, ConfigFilePermission, DaemonGuard, InstanceManager, ProcessRuntimeProvider, +}; +#[cfg(feature = "management")] +pub use full::remote_client; +#[cfg(feature = "management")] +pub use full::{ + ConfigFileStorage, ConfigServerEndpoint, InstanceMutationHooks, InstanceMutationResult, + LoggerControl, LoggerManagementRpc, ProcessManagement, ProcessManagementRpc, + UnsupportedConfigFileStorage, UnsupportedLoggerControl, WebClient, WebClientConfig, + apply_config_patch, call_instance_json_rpc, call_management_json_rpc, config_source_from_rpc, + config_source_to_rpc, log_level_name, network_instance_running_info, parse_log_level, + register_instance_management_rpc, register_management_rpc, +}; +pub use instance_rpc::InstanceManagementRpc; +pub use rpc_server_hook::ManagementRpcServerHook; +pub use selector::{ + ManagementInstance, ManagementSelector, resolve_instance, resolve_management_instance, + resolve_optional_instance_by_name, +}; +#[cfg(feature = "management")] +pub use server::ManagementServer; +pub use server::ReadOnlyManagementServer; +/// Registers the read-only status surface used by compact native nodes. +pub fn register_read_only_management_rpc( + manager: Arc>, + registry: &ServiceRegistry, +) where + F: InstanceFactory>, + H: CoreInstanceHost, +{ + let rpc = InstanceManagementRpc::::new(manager); + registry.register(PeerManageRpcServer::new(rpc.clone()), ""); + registry.register(ConnectorManageRpcServer::new(rpc), ""); +} diff --git a/easytier-core/src/management/rpc_server_hook.rs b/easytier-core/src/management/rpc_server_hook.rs new file mode 100644 index 00000000..634763b0 --- /dev/null +++ b/easytier-core/src/management/rpc_server_hook.rs @@ -0,0 +1,59 @@ +use std::{net::IpAddr, sync::Arc}; + +use cidr::IpCidr; +use easytier_proto::common::TunnelInfo; + +use crate::rpc::standalone::RpcServerHook; + +/// Restricts the process-level management endpoint to configured client CIDRs. +pub struct ManagementRpcServerHook { + whitelist: Vec, +} + +impl ManagementRpcServerHook { + pub fn new(whitelist: Option>) -> Self { + Self { + whitelist: whitelist.unwrap_or_else(|| { + vec!["127.0.0.0/8".parse().unwrap(), "::1/128".parse().unwrap()] + }), + } + } +} + +#[async_trait::async_trait] +impl RpcServerHook for ManagementRpcServerHook { + async fn on_new_client( + &self, + tunnel_info: Option, + ) -> Result, anyhow::Error> { + let tunnel_info = tunnel_info.ok_or_else(|| anyhow::anyhow!("tunnel info is None"))?; + let remote_url = tunnel_info + .remote_addr + .as_ref() + .ok_or_else(|| anyhow::anyhow!("remote_addr is None"))?; + let url = url::Url::parse(&remote_url.url) + .map_err(|error| anyhow::anyhow!("failed to parse remote URL: {error}"))?; + let host = url + .host_str() + .ok_or_else(|| anyhow::anyhow!("remote URL has no host"))?; + let ip_addr: IpAddr = host + .parse() + .map_err(|error| anyhow::anyhow!("failed to parse client IP {host}: {error}"))?; + + if self.whitelist.iter().any(|cidr| cidr.contains(&ip_addr)) { + return Ok(Some(tunnel_info)); + } + + Err(anyhow::anyhow!( + "RPC portal client IP {} is not in whitelist {:?}", + ip_addr, + self.whitelist + )) + } +} + +impl From for Arc { + fn from(value: ManagementRpcServerHook) -> Self { + Arc::new(value) + } +} diff --git a/easytier-core/src/management/selector.rs b/easytier-core/src/management/selector.rs new file mode 100644 index 00000000..18abae22 --- /dev/null +++ b/easytier-core/src/management/selector.rs @@ -0,0 +1,227 @@ +use std::sync::Arc; + +use easytier_proto::api::instance::{InstanceIdentifier, instance_identifier::Selector}; + +use crate::instance::{ + CoreInstance, CoreInstanceHost, + manager::{InstanceFactory, InstanceManager, ManagedInstance}, +}; + +/// Transport-independent selector for one managed Instance. +#[derive(Debug, Clone, PartialEq, Eq)] +pub enum ManagementSelector { + Id(uuid::Uuid), + Name(String), + UnambiguousDefault, +} + +/// Instance metadata required by process-level management selection. +pub trait ManagementInstance: ManagedInstance { + fn instance_name(&self) -> &str; +} + +impl ManagementInstance for CoreInstance +where + H: CoreInstanceHost, +{ + fn instance_name(&self) -> &str { + self.instance_name() + } +} + +/// Resolves one stateless management selector through the canonical Manager. +pub fn resolve_instance( + manager: &InstanceManager, + identifier: Option<&InstanceIdentifier>, +) -> anyhow::Result> +where + F: InstanceFactory, + F::Instance: ManagementInstance, +{ + let selector = match identifier.and_then(|identifier| identifier.selector.as_ref()) { + Some(Selector::Id(instance_id)) => ManagementSelector::Id((*instance_id).into()), + Some(Selector::InstanceSelector(selector)) => selector + .name + .clone() + .map(ManagementSelector::Name) + .unwrap_or(ManagementSelector::UnambiguousDefault), + None => ManagementSelector::UnambiguousDefault, + }; + resolve_management_instance(manager, &selector) +} + +/// Resolves one Instance without coupling the caller to an RPC request type. +pub fn resolve_management_instance( + manager: &InstanceManager, + selector: &ManagementSelector, +) -> anyhow::Result> +where + F: InstanceFactory, + F::Instance: ManagementInstance, +{ + if let ManagementSelector::Id(instance_id) = selector { + return manager + .get(*instance_id) + .ok_or_else(|| anyhow::anyhow!("Instance not found")); + } + + let matching = manager + .list() + .into_iter() + .filter(|instance| match selector { + ManagementSelector::Name(name) => instance.instance_name() == name, + ManagementSelector::UnambiguousDefault => true, + ManagementSelector::Id(_) => unreachable!(), + }) + .collect::>(); + + match matching.as_slice() { + [] => anyhow::bail!("No instance matches the selector"), + [instance] => Ok(instance.clone()), + _ => anyhow::bail!( + "{} instances match the selector, please specify the instance ID", + matching.len() + ), + } +} + +/// Resolves a unique name, distinguishing a missing name from ambiguity. +pub fn resolve_optional_instance_by_name( + manager: &InstanceManager, + name: &str, +) -> anyhow::Result>> +where + F: InstanceFactory, + F::Instance: ManagementInstance, +{ + let matching = manager + .list() + .into_iter() + .filter(|instance| instance.instance_name() == name) + .collect::>(); + match matching.as_slice() { + [] => Ok(None), + [instance] => Ok(Some(instance.clone())), + _ => anyhow::bail!( + "{} instances match the selector, please specify the instance ID", + matching.len() + ), + } +} + +#[cfg(test)] +mod tests { + use easytier_proto::{ + api::instance::{ + InstanceIdentifier, instance_identifier::InstanceSelector, + instance_identifier::Selector, + }, + common::Uuid as UuidPb, + }; + use uuid::Uuid; + + use super::*; + use crate::config::toml::{ConfigLoader as _, TomlConfig}; + #[derive(Debug)] + struct TestInstance { + id: Uuid, + name: String, + } + + impl ManagedInstance for TestInstance { + fn instance_id(&self) -> Uuid { + self.id + } + } + + impl ManagementInstance for TestInstance { + fn instance_name(&self) -> &str { + &self.name + } + } + + struct TestFactory; + + impl InstanceFactory for TestFactory { + type Instance = TestInstance; + type CreateContext = String; + type Error = std::convert::Infallible; + + fn create( + &self, + config: TomlConfig, + name: Self::CreateContext, + ) -> Result, Self::Error> { + Ok(Arc::new(TestInstance { + id: config.get_id(), + name, + })) + } + } + + fn add(manager: &InstanceManager, id: Uuid, name: &str) { + let config = TomlConfig::default(); + config.set_id(id); + manager.create(config, name.to_owned()).unwrap(); + } + + fn by_id(id: Uuid) -> InstanceIdentifier { + InstanceIdentifier { + selector: Some(Selector::Id(UuidPb::from(id))), + } + } + + fn by_name(name: &str) -> InstanceIdentifier { + InstanceIdentifier { + selector: Some(Selector::InstanceSelector(InstanceSelector { + name: Some(name.to_owned()), + })), + } + } + + #[test] + fn resolves_uuid_name_and_single_implicit_instance() { + let manager = InstanceManager::new(TestFactory, None); + let id = Uuid::new_v4(); + add(&manager, id, "alpha"); + + assert_eq!(resolve_instance(&manager, Some(&by_id(id))).unwrap().id, id); + assert_eq!( + resolve_instance(&manager, Some(&by_name("alpha"))) + .unwrap() + .id, + id + ); + assert_eq!(resolve_instance(&manager, None).unwrap().id, id); + } + + #[test] + fn rejects_missing_and_ambiguous_selectors() { + let manager = InstanceManager::new(TestFactory, None); + assert_eq!( + resolve_instance(&manager, None).unwrap_err().to_string(), + "No instance matches the selector" + ); + + add(&manager, Uuid::new_v4(), "same"); + add(&manager, Uuid::new_v4(), "same"); + assert!( + resolve_instance(&manager, None) + .unwrap_err() + .to_string() + .contains("2 instances match") + ); + assert!( + resolve_instance(&manager, Some(&by_name("same"))) + .unwrap_err() + .to_string() + .contains("2 instances match") + ); + assert_eq!( + resolve_instance(&manager, Some(&by_name("missing"))) + .unwrap_err() + .to_string(), + "No instance matches the selector" + ); + } +} diff --git a/easytier-core/src/management/server.rs b/easytier-core/src/management/server.rs new file mode 100644 index 00000000..fcc2cc76 --- /dev/null +++ b/easytier-core/src/management/server.rs @@ -0,0 +1,176 @@ +use std::sync::Arc; + +use cidr::IpCidr; + +use crate::{ + instance::{ + CoreInstance, CoreInstanceHost, + manager::{InstanceFactory, ProcessRuntimeProvider}, + }, + process_runtime::{CoreProcessRuntime, ProtectedTcpPortLease}, + proto::rpc_types::error::Error, + rpc::{ + service_registry::ServiceRegistry, + standalone::{RpcServerHook, StandAloneServer}, + }, + socket::SocketListener, + tunnel::Tunnel, +}; + +#[cfg(feature = "management")] +use super::{ConfigFileStorage, InstanceMutationHooks, LoggerControl, register_management_rpc}; +use super::{InstanceManager, ManagementRpcServerHook, register_read_only_management_rpc}; + +struct ManagementListener +where + L: SocketListener> + 'static, +{ + server: StandAloneServer, + process_runtime: Arc, +} + +impl ManagementListener +where + L: SocketListener> + 'static, +{ + fn new(listener: L, process_runtime: Arc) -> Self { + Self { + server: StandAloneServer::new(listener), + process_runtime, + } + } + + fn set_whitelist(&mut self, whitelist: Option>) { + let hook: Arc = Arc::new(ManagementRpcServerHook::new(whitelist)); + self.server.set_hook(hook); + } + + async fn serve(&mut self) -> crate::proto::rpc_types::error::Result<()> { + let process_runtime = self.process_runtime.clone(); + let binding_guard = protect_tcp_port(&process_runtime, &self.server.listener_url())?; + self.server + .serve_with_bound_listener(binding_guard, move |url| { + protect_tcp_port(&process_runtime, url) + }) + .await + } + + fn set_rx_timeout(&mut self, timeout: Option) { + self.server.set_rx_timeout(timeout); + } + + fn registry(&self) -> &ServiceRegistry { + self.server.registry() + } +} + +impl Drop for ManagementListener +where + L: SocketListener> + 'static, +{ + fn drop(&mut self) { + self.server.registry().unregister_all(); + } +} + +fn protect_tcp_port( + process_runtime: &CoreProcessRuntime, + url: &url::Url, +) -> Result, Error> { + match (url.scheme(), url.port()) { + ("tcp", Some(0) | None) => Err(anyhow::anyhow!( + "management TCP listener requires a concrete protected port before binding" + ) + .into()), + ("tcp", Some(port)) => Ok(Some(process_runtime.protect_tcp_port(port))), + _ => Ok(None), + } +} + +#[cfg(feature = "management")] +pub struct ManagementServer +where + L: SocketListener> + 'static, +{ + listener: ManagementListener, +} + +#[cfg(feature = "management")] +impl ManagementServer +where + L: SocketListener> + 'static, +{ + pub fn new( + listener: L, + instances: Arc>, + hooks: Arc, + storage: Arc, + logger: Arc, + ) -> Self + where + F: InstanceFactory, CreateContext = ()> + ProcessRuntimeProvider, + F::Error: std::fmt::Debug + std::fmt::Display + Send + Sync + 'static, + H: CoreInstanceHost, + { + let server = ManagementListener::new(listener, instances.process_runtime()); + register_management_rpc(instances, server.registry(), hooks, storage, logger); + Self { listener: server } + } + + /// Enables CIDR authorization for IP-based management transports. + /// Non-IP local transports intentionally keep the server default. + pub fn set_whitelist(&mut self, whitelist: Option>) { + self.listener.set_whitelist(whitelist); + } + + pub async fn serve(&mut self) -> crate::proto::rpc_types::error::Result<()> { + self.listener.serve().await + } + + pub fn with_rx_timeout(mut self, timeout: Option) -> Self { + self.listener.set_rx_timeout(timeout); + self + } + + pub fn set_rx_timeout(&mut self, timeout: Option) { + self.listener.set_rx_timeout(timeout); + } + + pub fn registry(&self) -> &ServiceRegistry { + self.listener.registry() + } +} + +pub struct ReadOnlyManagementServer +where + L: SocketListener> + 'static, +{ + listener: ManagementListener, +} + +impl ReadOnlyManagementServer +where + L: SocketListener> + 'static, +{ + pub fn new(listener: L, instances: Arc>) -> Self + where + F: InstanceFactory> + ProcessRuntimeProvider, + H: CoreInstanceHost, + { + let server = ManagementListener::new(listener, instances.process_runtime()); + register_read_only_management_rpc(instances, server.registry()); + Self { listener: server } + } + + pub fn set_whitelist(&mut self, whitelist: Option>) { + self.listener.set_whitelist(whitelist); + } + + pub async fn serve(&mut self) -> crate::proto::rpc_types::error::Result<()> { + self.listener.serve().await + } + + pub fn set_rx_timeout(&mut self, timeout: Option) { + self.listener.set_rx_timeout(timeout); + } +} diff --git a/easytier/src/common/compressor.rs b/easytier-core/src/packet/compressor.rs similarity index 71% rename from easytier/src/common/compressor.rs rename to easytier-core/src/packet/compressor.rs index d78972ee..68db8d66 100644 --- a/easytier/src/common/compressor.rs +++ b/easytier-core/src/packet/compressor.rs @@ -1,15 +1,8 @@ -#[cfg(feature = "zstd")] -use anyhow::Context; -#[cfg(feature = "zstd")] -use dashmap::DashMap; -#[cfg(feature = "zstd")] -use std::cell::RefCell; -#[cfg(feature = "zstd")] -use zstd::bulk; - use zerocopy::{AsBytes as _, FromBytes as _}; -use crate::tunnel::packet_def::{COMPRESSOR_TAIL_SIZE, CompressorAlgo, CompressorTail, ZCPacket}; +use super::{COMPRESSOR_TAIL_SIZE, CompressorAlgo, CompressorTail, ZCPacket}; + +mod zstd; type Error = anyhow::Error; @@ -42,17 +35,7 @@ impl DefaultCompressor { compress_algo: CompressorAlgo, ) -> Result, Error> { match compress_algo { - #[cfg(feature = "zstd")] - CompressorAlgo::ZstdDefault => CTX_MAP.with(|map_cell| { - let map = map_cell.borrow(); - let mut ctx_entry = map.entry(compress_algo).or_default(); - ctx_entry.compress(data).with_context(|| { - format!( - "Failed to compress data with algorithm: {:?}", - compress_algo - ) - }) - }), + CompressorAlgo::ZstdDefault => zstd::compress(data, compress_algo), CompressorAlgo::None => Ok(data.to_vec()), } } @@ -63,28 +46,7 @@ impl DefaultCompressor { compress_algo: CompressorAlgo, ) -> Result, Error> { match compress_algo { - #[cfg(feature = "zstd")] - CompressorAlgo::ZstdDefault => DCTX_MAP.with(|map_cell| { - let map = map_cell.borrow(); - let mut ctx_entry = map.entry(compress_algo).or_default(); - for i in 1..=5 { - let mut len = data.len() * 2usize.pow(i); - if i == 5 && len < 64 * 1024 { - len = 64 * 1024; // Ensure a minimum buffer size - } - match ctx_entry.decompress(data, len) { - Ok(buf) => return Ok(buf), - Err(e) if e.to_string().contains("buffer is too small") => { - continue; // Try with a larger buffer - } - Err(e) => return Err(e.into()), - } - } - Err(anyhow::anyhow!( - "Failed to decompress data after multiple attempts with algorithm: {:?}", - compress_algo - )) - }), + CompressorAlgo::ZstdDefault => zstd::decompress(data, compress_algo), CompressorAlgo::None => Ok(data.to_vec()), } } @@ -175,16 +137,15 @@ impl Compressor for DefaultCompressor { } } -#[cfg(feature = "zstd")] -thread_local! { - static CTX_MAP: RefCell>> = RefCell::new(DashMap::new()); - static DCTX_MAP: RefCell>> = RefCell::new(DashMap::new()); +pub(super) fn zstd_available() -> bool { + zstd::AVAILABLE } -#[cfg(all(test, feature = "zstd"))] +#[cfg(test)] pub mod tests { use super::*; + #[cfg(feature = "zstd")] #[tokio::test] async fn test_compress() { let text = b"12345670000000000000000000"; @@ -215,6 +176,7 @@ pub mod tests { assert!(!packet.peer_manager_header().unwrap().is_compressed()); } + #[cfg(feature = "zstd")] #[tokio::test] async fn test_short_text_compress() { let text = b"1234"; @@ -234,4 +196,18 @@ pub mod tests { assert_eq!(packet.payload(), text); assert!(!packet.peer_manager_header().unwrap().is_compressed()); } + + #[cfg(not(feature = "zstd"))] + #[tokio::test] + async fn unavailable_zstd_returns_an_explicit_error() { + let error = DefaultCompressor::new() + .compress_raw(b"payload", CompressorAlgo::ZstdDefault) + .await + .unwrap_err(); + + assert_eq!( + error.to_string(), + "compression algorithm is unavailable in this build: ZstdDefault" + ); + } } diff --git a/easytier-core/src/packet/compressor/zstd.rs b/easytier-core/src/packet/compressor/zstd.rs new file mode 100644 index 00000000..629a234b --- /dev/null +++ b/easytier-core/src/packet/compressor/zstd.rs @@ -0,0 +1,76 @@ +#[cfg(feature = "zstd")] +use std::cell::RefCell; + +#[cfg(feature = "zstd")] +use anyhow::Context as _; +#[cfg(feature = "zstd")] +use dashmap::DashMap; +#[cfg(feature = "zstd")] +use zstd::bulk; + +use super::CompressorAlgo; + +#[cfg(feature = "zstd")] +pub(super) const AVAILABLE: bool = true; +#[cfg(not(feature = "zstd"))] +pub(super) const AVAILABLE: bool = false; + +#[cfg(feature = "zstd")] +thread_local! { + static CTX_MAP: RefCell>> = + RefCell::new(DashMap::new()); + static DCTX_MAP: RefCell>> = + RefCell::new(DashMap::new()); +} + +#[cfg(feature = "zstd")] +pub(super) fn compress(data: &[u8], compress_algo: CompressorAlgo) -> anyhow::Result> { + CTX_MAP.with(|map_cell| { + let map = map_cell.borrow(); + let mut ctx_entry = map.entry(compress_algo).or_default(); + ctx_entry.compress(data).with_context(|| { + format!( + "Failed to compress data with algorithm: {:?}", + compress_algo + ) + }) + }) +} + +#[cfg(not(feature = "zstd"))] +pub(super) fn compress(_data: &[u8], compress_algo: CompressorAlgo) -> anyhow::Result> { + unavailable(compress_algo) +} + +#[cfg(feature = "zstd")] +pub(super) fn decompress(data: &[u8], compress_algo: CompressorAlgo) -> anyhow::Result> { + DCTX_MAP.with(|map_cell| { + let map = map_cell.borrow(); + let mut ctx_entry = map.entry(compress_algo).or_default(); + for i in 1..=5 { + let mut len = data.len() * 2usize.pow(i); + if i == 5 && len < 64 * 1024 { + len = 64 * 1024; + } + match ctx_entry.decompress(data, len) { + Ok(buf) => return Ok(buf), + Err(error) if error.to_string().contains("buffer is too small") => continue, + Err(error) => return Err(error.into()), + } + } + Err(anyhow::anyhow!( + "Failed to decompress data after multiple attempts with algorithm: {:?}", + compress_algo + )) + }) +} + +#[cfg(not(feature = "zstd"))] +pub(super) fn decompress(_data: &[u8], compress_algo: CompressorAlgo) -> anyhow::Result> { + unavailable(compress_algo) +} + +#[cfg(not(feature = "zstd"))] +fn unavailable(compress_algo: CompressorAlgo) -> anyhow::Result> { + Err(super::super::CompressionUnavailableError(compress_algo).into()) +} diff --git a/easytier-core/src/packet/hole_punch.rs b/easytier-core/src/packet/hole_punch.rs new file mode 100644 index 00000000..f430badb --- /dev/null +++ b/easytier-core/src/packet/hole_punch.rs @@ -0,0 +1,79 @@ +use bytes::BytesMut; +use rand::{Rng, SeedableRng}; +use zerocopy::FromBytes as _; + +use super::{UDP_TUNNEL_HEADER_SIZE, UDPTunnelHeader, UdpPacketType, ZCPacket, ZCPacketType}; + +pub(crate) const HOLE_PUNCH_PACKET_BODY_LEN: u16 = 16; + +fn new_udp_packet(f: F, udp_body: &[u8]) -> ZCPacket +where + F: FnOnce(&mut UDPTunnelHeader), +{ + let mut buf = BytesMut::new(); + buf.resize(UDP_TUNNEL_HEADER_SIZE + udp_body.len(), 0); + buf[UDP_TUNNEL_HEADER_SIZE..].copy_from_slice(udp_body); + + let mut ret = ZCPacket::new_from_buf(buf, ZCPacketType::UDP); + let header = ret.mut_udp_tunnel_header().unwrap(); + f(header); + ret +} + +pub(crate) fn new_hole_punch_packet(tid: u32, buf_len: u16) -> ZCPacket { + let mut rng = rand::rngs::StdRng::from_entropy(); + let mut buf = vec![0u8; buf_len as usize]; + rng.fill(&mut buf[..]); + new_udp_packet( + |header| { + header.msg_type = UdpPacketType::HolePunch as u8; + header.conn_id.set(tid); + header.len.set(buf_len); + }, + &buf, + ) +} + +pub(crate) fn hole_punch_packet_tid(data: &[u8], body_len: u16) -> Option { + if data.len() != UDP_TUNNEL_HEADER_SIZE + body_len as usize { + return None; + } + + let header = UDPTunnelHeader::ref_from_prefix(data)?; + let valid = header.msg_type == UdpPacketType::HolePunch as u8 && header.len.get() == body_len; + + valid.then(|| header.conn_id.get()) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn builds_and_parses_hole_punch_packet() { + let tid = 0x1234_5678; + let packet = new_hole_punch_packet(tid, HOLE_PUNCH_PACKET_BODY_LEN); + let bytes = packet.into_bytes(); + + assert_eq!( + bytes.len(), + UDP_TUNNEL_HEADER_SIZE + HOLE_PUNCH_PACKET_BODY_LEN as usize + ); + assert_eq!( + hole_punch_packet_tid(&bytes, HOLE_PUNCH_PACKET_BODY_LEN), + Some(tid) + ); + } + + #[test] + fn rejects_non_matching_hole_punch_packet_length() { + let packet = new_hole_punch_packet(1, HOLE_PUNCH_PACKET_BODY_LEN); + let mut bytes = packet.into_bytes().to_vec(); + bytes.pop(); + + assert_eq!( + hole_punch_packet_tid(&bytes, HOLE_PUNCH_PACKET_BODY_LEN), + None + ); + } +} diff --git a/easytier/src/tunnel/packet_def.rs b/easytier-core/src/packet/mod.rs similarity index 87% rename from easytier/src/tunnel/packet_def.rs rename to easytier-core/src/packet/mod.rs index b659cf4f..6905689c 100644 --- a/easytier/src/tunnel/packet_def.rs +++ b/easytier-core/src/packet/mod.rs @@ -1,6 +1,15 @@ +pub(crate) mod compressor; +mod hole_punch; +pub mod stun; + +pub(crate) use hole_punch::{ + HOLE_PUNCH_PACKET_BODY_LEN, hole_punch_packet_tid, new_hole_punch_packet, +}; + use bytes::Buf; use bytes::Bytes; use bytes::BytesMut; +use easytier_proto::common::CompressionAlgoPb; use zerocopy::AsBytes; use zerocopy::FromBytes; use zerocopy::FromZeroes; @@ -217,20 +226,12 @@ impl PeerManagerHeader { self } - pub fn is_kcp_src_modified(&self) -> bool { - self.packet_type == PacketType::DataWithKcpSrcModified as u8 - } - pub fn mark_quic_src_modified(&mut self) -> &mut Self { assert_eq!(self.packet_type, PacketType::Data as u8); self.packet_type = PacketType::DataWithQuicSrcModified as u8; self } - pub fn is_quic_src_modified(&self) -> bool { - self.packet_type == PacketType::DataWithQuicSrcModified as u8 - } - pub fn set_not_send_to_tun(&mut self, not_send_to_tun: bool) -> &mut Self { let mut flags = PeerManagerHeaderFlags::from_bits(self.flags).unwrap(); if not_send_to_tun { @@ -310,10 +311,55 @@ pub type StandardAeadTail = AeadTail<16, 12>; #[repr(u8)] pub enum CompressorAlgo { None = 0, - #[cfg(feature = "zstd")] ZstdDefault = 1, } +impl CompressorAlgo { + pub fn is_available(self) -> bool { + match self { + Self::None => true, + Self::ZstdDefault => compressor::zstd_available(), + } + } + + pub fn ensure_available(self) -> Result<(), CompressionUnavailableError> { + self.is_available() + .then_some(()) + .ok_or(CompressionUnavailableError(self)) + } +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq, thiserror::Error)] +#[error("invalid compression algorithm: {0:?}")] +pub struct CompressionAlgoError(pub CompressionAlgoPb); + +#[derive(Debug, Clone, Copy, PartialEq, Eq, thiserror::Error)] +#[error("compression algorithm is unavailable in this build: {0:?}")] +pub struct CompressionUnavailableError(pub CompressorAlgo); + +impl TryFrom for CompressorAlgo { + type Error = CompressionAlgoError; + + fn try_from(value: CompressionAlgoPb) -> Result { + match value { + CompressionAlgoPb::Zstd => Ok(CompressorAlgo::ZstdDefault), + CompressionAlgoPb::None => Ok(CompressorAlgo::None), + _ => Err(CompressionAlgoError(value)), + } + } +} + +impl TryFrom for CompressionAlgoPb { + type Error = CompressionAlgoError; + + fn try_from(value: CompressorAlgo) -> Result { + match value { + CompressorAlgo::ZstdDefault => Ok(CompressionAlgoPb::Zstd), + CompressorAlgo::None => Ok(CompressionAlgoPb::None), + } + } +} + #[repr(C, packed)] #[derive(AsBytes, FromBytes, FromZeroes, Clone, Debug, Default)] pub struct CompressorTail { @@ -324,7 +370,6 @@ pub const COMPRESSOR_TAIL_SIZE: usize = std::mem::size_of::(); impl CompressorTail { pub fn get_algo(&self) -> Option { match self.algo { - #[cfg(feature = "zstd")] 1 => Some(CompressorAlgo::ZstdDefault), _ => None, } @@ -601,15 +646,6 @@ impl ZCPacket { PeerManagerHeader::ref_from_prefix(bytes) } - pub fn tcp_tunnel_header(&self) -> Option<&TCPTunnelHeader> { - let offset = self - .packet_type - .get_packet_offsets() - .tcp_tunnel_header_offset; - let bytes = self.bytes_from_offset(offset)?; - TCPTunnelHeader::ref_from_prefix(bytes) - } - pub fn udp_tunnel_header(&self) -> Option<&UDPTunnelHeader> { let offset = self .packet_type @@ -773,6 +809,28 @@ impl ZCPacket { mod tests { use super::*; + #[cfg(feature = "proxy-packet")] + impl PeerManagerHeader { + pub(crate) fn is_kcp_src_modified(&self) -> bool { + self.packet_type == PacketType::DataWithKcpSrcModified as u8 + } + + pub(crate) fn is_quic_src_modified(&self) -> bool { + self.packet_type == PacketType::DataWithQuicSrcModified as u8 + } + } + + impl ZCPacket { + fn tcp_tunnel_header(&self) -> Option<&TCPTunnelHeader> { + let offset = self + .packet_type + .get_packet_offsets() + .tcp_tunnel_header_offset; + let bytes = self.bytes_from_offset(offset)?; + TCPTunnelHeader::ref_from_prefix(bytes) + } + } + #[test] fn test_zc_packet() { let payload = b"hello world"; @@ -817,4 +875,48 @@ mod tests { assert!(packet.mut_wg_tunnel_header().is_none()); } + + #[test] + fn converts_compression_algo_none() { + assert_eq!( + CompressorAlgo::None, + CompressorAlgo::try_from(CompressionAlgoPb::None).unwrap() + ); + assert_eq!( + CompressionAlgoPb::None, + CompressionAlgoPb::try_from(CompressorAlgo::None).unwrap() + ); + } + + #[test] + fn converts_zstd_compression_algo_in_every_profile() { + assert_eq!( + CompressorAlgo::ZstdDefault, + CompressorAlgo::try_from(CompressionAlgoPb::Zstd).unwrap() + ); + assert_eq!( + CompressionAlgoPb::Zstd, + CompressionAlgoPb::try_from(CompressorAlgo::ZstdDefault).unwrap() + ); + assert_eq!( + Some(CompressorAlgo::ZstdDefault), + CompressorTail { algo: 1 }.get_algo() + ); + } + + #[cfg(not(feature = "zstd"))] + #[test] + fn reports_zstd_as_unavailable_without_changing_vocabulary() { + assert!(!CompressorAlgo::ZstdDefault.is_available()); + assert_eq!( + CompressorAlgo::ZstdDefault.ensure_available().unwrap_err(), + CompressionUnavailableError(CompressorAlgo::ZstdDefault) + ); + } + + #[cfg(feature = "zstd")] + #[test] + fn reports_zstd_as_available_when_compiled() { + assert!(CompressorAlgo::ZstdDefault.is_available()); + } } diff --git a/easytier/src/common/stun_codec_ext.rs b/easytier-core/src/packet/stun.rs similarity index 91% rename from easytier/src/common/stun_codec_ext.rs rename to easytier-core/src/packet/stun.rs index 884645af..9ff60cb6 100644 --- a/easytier/src/common/stun_codec_ext.rs +++ b/easytier-core/src/packet/stun.rs @@ -1,18 +1,24 @@ +//! STUN wire codec for the attributes EasyTier NAT traversal uses. +//! +//! EasyTier speaks STUN (RFC 5389/5780) for NAT mapping detection and port +//! mapping. This module holds the wire-level attribute types, their codecs, +//! and the EasyTier transaction-id convention (a `0xdeadbeef` prefix); it +//! performs no I/O. The connectivity layer drives probing and responding on +//! top of it. + use std::net::SocketAddr; use bytecodec::fixnum::{U32beDecoder, U32beEncoder}; +use bytecodec::{ByteCount, Decode, Encode, Eos, Result}; +use bytecodec::{SizedEncode, TryTaggedDecode}; +use stun_codec::macros::track; use stun_codec::net::{SocketAddrDecoder, SocketAddrEncoder, socket_addr_xor}; - use stun_codec::rfc5389::attributes::{ MappedAddress, Software, XorMappedAddress, XorMappedAddress2, }; use stun_codec::rfc5780::attributes::{OtherAddress, ResponseOrigin}; use stun_codec::{AttributeType, Message, TransactionId, define_attribute_enums}; -use bytecodec::{ByteCount, Decode, Encode, Eos, Result, SizedEncode, TryTaggedDecode}; - -use stun_codec::macros::track; - macro_rules! impl_decode { ($decoder:ty, $item:ident, $and_then:expr) => { impl Decode for $decoder { @@ -298,3 +304,16 @@ define_attribute_enums!( ResponseOrigin ] ); + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn easytier_transaction_id_roundtrips_u32() { + let tid = u32_to_tid(0x1122_3344); + + assert_eq!(&tid.as_bytes()[..4], &[0xde, 0xad, 0xbe, 0xef]); + assert_eq!(tid_to_u32(&tid), 0x1122_3344); + } +} diff --git a/easytier/src/peers/acl_filter.rs b/easytier-core/src/peers/acl/filter.rs similarity index 59% rename from easytier/src/peers/acl_filter.rs rename to easytier-core/src/peers/acl/filter.rs index 9b6575ef..14f149f3 100644 --- a/easytier/src/peers/acl_filter.rs +++ b/easytier-core/src/peers/acl/filter.rs @@ -2,26 +2,113 @@ use std::net::{Ipv4Addr, Ipv6Addr}; use std::sync::atomic::Ordering; use std::{ net::IpAddr, - sync::{Arc, atomic::AtomicBool}, + sync::{Arc, Mutex, atomic::AtomicBool}, }; use arc_swap::ArcSwap; use dashmap::DashMap; -use pnet::packet::ipv6::Ipv6Packet; -use pnet::packet::{ - Packet as _, ip::IpNextHeaderProtocols, ipv4::Ipv4Packet, tcp::TcpPacket, udp::UdpPacket, -}; +use easytier_proto::acl::{Acl, AclStats, Action, ChainType, Protocol}; use quanta::Instant; - -use crate::proto::acl::{AclStats, Protocol}; -use crate::tunnel::packet_def::PacketType; -use crate::{ - common::acl_processor::{AclProcessor, AclResult, AclStatKey, AclStatType, PacketInfo}, - proto::acl::{Acl, Action, ChainType}, - tunnel::packet_def::ZCPacket, -}; use tokio_util::task::AbortOnDropHandle; +use crate::{ + packet::{PacketType, ZCPacket}, + peers::acl::processor::{AclProcessor, AclResult, AclStatKey, AclStatType, PacketInfo}, +}; + +const IP_PROTO_ICMP: u8 = 1; +const IP_PROTO_TCP: u8 = 6; +const IP_PROTO_UDP: u8 = 17; +const IP_PROTO_ICMPV6: u8 = 58; + +#[derive(Clone, Copy)] +struct ParsedIpPacket<'a> { + src_ip: IpAddr, + dst_ip: IpAddr, + protocol: u8, + transport_payload: &'a [u8], +} + +fn parse_ip_packet(payload: &[u8]) -> Option> { + let version = payload.first()? >> 4; + match version { + 4 => parse_ipv4_packet(payload), + 6 => parse_ipv6_packet(payload), + _ => None, + } +} + +fn parse_ipv4_packet(payload: &[u8]) -> Option> { + if payload.len() < 20 { + return None; + } + let header_len = usize::from(payload[0] & 0x0f) * 4; + let options_len = header_len.saturating_sub(20); + let payload_offset = 20 + options_len; + let payload_start = payload_offset.min(payload.len()); + let total_length = usize::from(u16::from_be_bytes([payload[2], payload[3]])); + let payload_len = total_length.saturating_sub(header_len); + let payload_end = payload_start.saturating_add(payload_len).min(payload.len()); + + Some(ParsedIpPacket { + src_ip: IpAddr::V4(Ipv4Addr::new( + payload[12], + payload[13], + payload[14], + payload[15], + )), + dst_ip: IpAddr::V4(Ipv4Addr::new( + payload[16], + payload[17], + payload[18], + payload[19], + )), + protocol: payload[9], + transport_payload: &payload[payload_start..payload_end], + }) +} + +fn parse_ipv6_packet(payload: &[u8]) -> Option> { + if payload.len() < 40 { + return None; + } + let payload_len = usize::from(u16::from_be_bytes([payload[4], payload[5]])); + let payload_end = 40usize.saturating_add(payload_len).min(payload.len()); + + Some(ParsedIpPacket { + src_ip: IpAddr::V6(Ipv6Addr::from(<[u8; 16]>::try_from(&payload[8..24]).ok()?)), + dst_ip: IpAddr::V6(Ipv6Addr::from(<[u8; 16]>::try_from(&payload[24..40]).ok()?)), + protocol: payload[6], + transport_payload: &payload[40..payload_end], + }) +} + +fn parse_transport_ports(protocol: u8, payload: &[u8]) -> Option<(Option, Option)> { + let min_len = match protocol { + IP_PROTO_TCP => 20, + IP_PROTO_UDP => 8, + _ => return Some((None, None)), + }; + if payload.len() < min_len { + return None; + } + + Some(( + Some(u16::from_be_bytes([payload[0], payload[1]])), + Some(u16::from_be_bytes([payload[2], payload[3]])), + )) +} + +fn acl_protocol(protocol: u8) -> Protocol { + match protocol { + IP_PROTO_TCP => Protocol::Tcp, + IP_PROTO_UDP => Protocol::Udp, + IP_PROTO_ICMP => Protocol::Icmp, + IP_PROTO_ICMPV6 => Protocol::IcmPv6, + _ => Protocol::Unspecified, + } +} + #[derive(Debug, Eq, PartialEq, Hash)] struct OutboundAllowRecord { src_ip: IpAddr, @@ -63,7 +150,8 @@ pub struct AclFilter { // Track allowed outbound packets and automatically allow their corresponding inbound response // packets, even if they would normally be dropped by ACL rules outbound_allow_records: Arc>, - clean_task: AbortOnDropHandle<()>, + #[allow(dead_code)] + clean_task: Mutex>>, } impl Default for AclFilter { @@ -80,13 +168,21 @@ impl AclFilter { acl_processor: ArcSwap::from(Arc::new(AclProcessor::new(Acl::default()))), acl_enabled: Arc::new(AtomicBool::new(false)), outbound_allow_records, - clean_task: AbortOnDropHandle::new(tokio::spawn(async move { + clean_task: Mutex::new(Some(AbortOnDropHandle::new(tokio::spawn(async move { let max_life = std::time::Duration::from_secs(30); loop { record_clone.retain(|_, v| v.elapsed() < max_life); - tokio::time::sleep(std::time::Duration::from_secs(30)).await; + crate::foundation::time::sleep(std::time::Duration::from_secs(30)).await; } - })), + })))), + } + } + + pub(crate) async fn stop_cleanup_task(&self) { + let task = self.clean_task.lock().unwrap().take(); + if let Some(task) = task { + task.abort(); + let _ = task.await; } } @@ -142,73 +238,14 @@ impl AclFilter { fn extract_packet_info( &self, packet: &ZCPacket, - route: &(dyn super::route_trait::Route + Send + Sync + 'static), + route: &(dyn crate::peers::route::Route + Send + Sync + 'static), ) -> Option { let payload = packet.payload(); - let src_ip; - let dst_ip; - let src_port; - let dst_port; - let protocol; - - let ipv4_packet = Ipv4Packet::new(payload)?; - if ipv4_packet.get_version() == 4 { - src_ip = IpAddr::V4(ipv4_packet.get_source()); - dst_ip = IpAddr::V4(ipv4_packet.get_destination()); - protocol = ipv4_packet.get_next_level_protocol(); - - (src_port, dst_port) = match protocol { - IpNextHeaderProtocols::Tcp => { - let tcp_packet = TcpPacket::new(ipv4_packet.payload())?; - ( - Some(tcp_packet.get_source()), - Some(tcp_packet.get_destination()), - ) - } - IpNextHeaderProtocols::Udp => { - let udp_packet = UdpPacket::new(ipv4_packet.payload())?; - ( - Some(udp_packet.get_source()), - Some(udp_packet.get_destination()), - ) - } - _ => (None, None), - }; - } else if ipv4_packet.get_version() == 6 { - let ipv6_packet = Ipv6Packet::new(payload)?; - src_ip = IpAddr::V6(ipv6_packet.get_source()); - dst_ip = IpAddr::V6(ipv6_packet.get_destination()); - protocol = ipv6_packet.get_next_header(); - - (src_port, dst_port) = match protocol { - IpNextHeaderProtocols::Tcp => { - let tcp_packet = TcpPacket::new(ipv6_packet.payload())?; - ( - Some(tcp_packet.get_source()), - Some(tcp_packet.get_destination()), - ) - } - IpNextHeaderProtocols::Udp => { - let udp_packet = UdpPacket::new(ipv6_packet.payload())?; - ( - Some(udp_packet.get_source()), - Some(udp_packet.get_destination()), - ) - } - _ => (None, None), - }; - } else { - return None; - } - - let acl_protocol = match protocol { - IpNextHeaderProtocols::Tcp => Protocol::Tcp, - IpNextHeaderProtocols::Udp => Protocol::Udp, - IpNextHeaderProtocols::Icmp => Protocol::Icmp, - IpNextHeaderProtocols::Icmpv6 => Protocol::IcmPv6, - _ => Protocol::Unspecified, - }; + let parsed = parse_ip_packet(payload)?; + let (src_port, dst_port) = + parse_transport_ports(parsed.protocol, parsed.transport_payload)?; + let acl_protocol = acl_protocol(parsed.protocol); let src_groups = packet .get_src_peer_id() @@ -220,8 +257,8 @@ impl AclFilter { .unwrap_or_else(|| Arc::new(Vec::new())); Some(PacketInfo { - src_ip, - dst_ip, + src_ip: parsed.src_ip, + dst_ip: parsed.dst_ip, src_port, dst_port, protocol: acl_protocol, @@ -321,7 +358,7 @@ impl AclFilter { is_in: bool, my_ipv4: Option, is_local_ipv6: impl Fn(Ipv6Addr) -> bool, - route: &(dyn super::route_trait::Route + Send + Sync + 'static), + route: &(dyn crate::peers::route::Route + Send + Sync + 'static), ) -> bool { if !self.acl_enabled.load(Ordering::Relaxed) { return true; @@ -406,12 +443,20 @@ mod tests { use quanta::Instant; - use crate::{ - common::acl_processor::PacketInfo, - proto::acl::{Acl, ChainType, Protocol}, + use easytier_proto::acl::{Acl, ChainType, Protocol}; + + use crate::peers::acl::processor::PacketInfo; + + use super::{ + AclFilter, IP_PROTO_ICMP, IP_PROTO_TCP, IP_PROTO_UDP, OutboundAllowRecord, acl_protocol, + parse_ip_packet, parse_transport_ports, }; - use super::{AclFilter, OutboundAllowRecord}; + impl AclFilter { + pub(crate) fn cleanup_task_is_stopped(&self) -> bool { + self.clean_task.lock().unwrap().is_none() + } + } fn packet_info(dst_ip: IpAddr) -> PacketInfo { PacketInfo { @@ -426,6 +471,123 @@ mod tests { } } + #[test] + fn parse_ipv4_tcp_packet_extracts_addrs_and_ports() { + let mut packet = vec![0u8; 40]; + packet[0] = 0x45; + packet[2..4].copy_from_slice(&40u16.to_be_bytes()); + packet[9] = IP_PROTO_TCP; + packet[12..16].copy_from_slice(&[10, 0, 0, 1]); + packet[16..20].copy_from_slice(&[10, 0, 0, 2]); + packet[20..22].copy_from_slice(&1234u16.to_be_bytes()); + packet[22..24].copy_from_slice(&80u16.to_be_bytes()); + + let parsed = parse_ip_packet(&packet).unwrap(); + let (src_port, dst_port) = + parse_transport_ports(parsed.protocol, parsed.transport_payload).unwrap(); + + assert_eq!(parsed.src_ip, IpAddr::V4(Ipv4Addr::new(10, 0, 0, 1))); + assert_eq!(parsed.dst_ip, IpAddr::V4(Ipv4Addr::new(10, 0, 0, 2))); + assert_eq!(acl_protocol(parsed.protocol), Protocol::Tcp); + assert_eq!(src_port, Some(1234)); + assert_eq!(dst_port, Some(80)); + } + + #[test] + fn parse_ipv6_udp_packet_extracts_addrs_and_ports() { + let src: Ipv6Addr = "2001:db8::1".parse().unwrap(); + let dst: Ipv6Addr = "2001:db8::2".parse().unwrap(); + let mut packet = vec![0u8; 48]; + packet[0] = 0x60; + packet[4..6].copy_from_slice(&8u16.to_be_bytes()); + packet[6] = IP_PROTO_UDP; + packet[8..24].copy_from_slice(&src.octets()); + packet[24..40].copy_from_slice(&dst.octets()); + packet[40..42].copy_from_slice(&5353u16.to_be_bytes()); + packet[42..44].copy_from_slice(&53u16.to_be_bytes()); + + let parsed = parse_ip_packet(&packet).unwrap(); + let (src_port, dst_port) = + parse_transport_ports(parsed.protocol, parsed.transport_payload).unwrap(); + + assert_eq!(parsed.src_ip, IpAddr::V6(src)); + assert_eq!(parsed.dst_ip, IpAddr::V6(dst)); + assert_eq!(acl_protocol(parsed.protocol), Protocol::Udp); + assert_eq!(src_port, Some(5353)); + assert_eq!(dst_port, Some(53)); + } + + #[test] + fn parse_ipv4_uses_declared_total_length() { + let mut packet = vec![0u8; 40]; + packet[0] = 0x45; + packet[2..4].copy_from_slice(&24u16.to_be_bytes()); + packet[9] = IP_PROTO_TCP; + packet[12..16].copy_from_slice(&[10, 0, 0, 1]); + packet[16..20].copy_from_slice(&[10, 0, 0, 2]); + packet[20..22].copy_from_slice(&1234u16.to_be_bytes()); + packet[22..24].copy_from_slice(&80u16.to_be_bytes()); + + let parsed = parse_ip_packet(&packet).unwrap(); + + assert_eq!(parsed.transport_payload.len(), 4); + assert!(parse_transport_ports(parsed.protocol, parsed.transport_payload).is_none()); + } + + #[test] + fn parse_ipv6_uses_declared_payload_length() { + let mut packet = vec![0u8; 48]; + packet[0] = 0x60; + packet[4..6].copy_from_slice(&4u16.to_be_bytes()); + packet[6] = IP_PROTO_UDP; + packet[40..42].copy_from_slice(&5353u16.to_be_bytes()); + packet[42..44].copy_from_slice(&53u16.to_be_bytes()); + + let parsed = parse_ip_packet(&packet).unwrap(); + + assert_eq!(parsed.transport_payload.len(), 4); + assert!(parse_transport_ports(parsed.protocol, parsed.transport_payload).is_none()); + } + + #[test] + fn parse_ipv4_keeps_pnet_ihl_less_than_five_behavior() { + let mut packet = vec![0u8; 40]; + packet[0] = 0x44; + packet[2..4].copy_from_slice(&40u16.to_be_bytes()); + packet[9] = IP_PROTO_TCP; + packet[20..22].copy_from_slice(&1234u16.to_be_bytes()); + packet[22..24].copy_from_slice(&80u16.to_be_bytes()); + + let parsed = parse_ip_packet(&packet).unwrap(); + let (src_port, dst_port) = + parse_transport_ports(parsed.protocol, parsed.transport_payload).unwrap(); + + assert_eq!(parsed.transport_payload.len(), 20); + assert_eq!(src_port, Some(1234)); + assert_eq!(dst_port, Some(80)); + } + + #[test] + fn parse_ipv4_keeps_pnet_truncated_options_behavior() { + let mut packet = vec![0u8; 20]; + packet[0] = 0x4f; + packet[2..4].copy_from_slice(&60u16.to_be_bytes()); + packet[9] = IP_PROTO_ICMP; + packet[12..16].copy_from_slice(&[10, 0, 0, 1]); + packet[16..20].copy_from_slice(&[10, 0, 0, 2]); + + let parsed = parse_ip_packet(&packet).unwrap(); + let (src_port, dst_port) = + parse_transport_ports(parsed.protocol, parsed.transport_payload).unwrap(); + + assert_eq!(parsed.src_ip, IpAddr::V4(Ipv4Addr::new(10, 0, 0, 1))); + assert_eq!(parsed.dst_ip, IpAddr::V4(Ipv4Addr::new(10, 0, 0, 2))); + assert_eq!(acl_protocol(parsed.protocol), Protocol::Icmp); + assert!(parsed.transport_payload.is_empty()); + assert_eq!(src_port, None); + assert_eq!(dst_port, None); + } + #[test] fn classify_chain_type_treats_public_ipv6_lease_as_inbound() { let leased_ipv6 = Ipv6Addr::new(0x2001, 0xdb8, 0x100, 0, 0, 0, 0, 0x123); diff --git a/easytier-core/src/peers/acl/mod.rs b/easytier-core/src/peers/acl/mod.rs new file mode 100644 index 00000000..139e2542 --- /dev/null +++ b/easytier-core/src/peers/acl/mod.rs @@ -0,0 +1,7 @@ +//! Access-control list packet filtering: the per-rule processor and the +//! filter wiring it into the peer/NIC packet pipelines. + +pub(crate) mod filter; +pub(crate) mod processor; + +pub(crate) use filter::AclFilter; diff --git a/easytier/src/common/acl_processor.rs b/easytier-core/src/peers/acl/processor.rs similarity index 86% rename from easytier/src/common/acl_processor.rs rename to easytier-core/src/peers/acl/processor.rs index 4ae88693..649f112a 100644 --- a/easytier/src/common/acl_processor.rs +++ b/easytier-core/src/peers/acl/processor.rs @@ -8,10 +8,9 @@ use std::{ use quanta::Instant; -use crate::common::{config::ConfigLoader, global_ctx::ArcGlobalCtx, token_bucket::TokenBucket}; -use crate::proto::acl::*; -use anyhow::Context as _; +use crate::foundation::token_bucket::TokenBucket; use dashmap::DashMap; +use easytier_proto::acl::*; use tokio::task::JoinSet; // Performance-optimized key for rate limiting to avoid string allocations @@ -110,13 +109,11 @@ impl AclCacheKey { // Cache entry with timestamp for LRU cleanup #[derive(Debug, Clone)] pub(crate) struct AclCacheEntry { - pub action: Action, pub matched_rule: RuleId, pub last_access: Instant, // New fields to track rule characteristics for proper cache behavior pub conn_track_key: Option, pub rate_limit_keys: Vec, - pub chain_type: ChainType, pub acl_result: Option, pub rule_stats_vec: Vec>, } @@ -144,11 +141,6 @@ pub struct AclResult { } impl AclResult { - /// Get matched rule as string (lazy evaluation) - pub fn matched_rule_string(&self) -> Option { - self.matched_rule.as_ref().map(|r| r.to_string_cached()) - } - /// Get matched rule as string reference for logging (compatibility method) pub fn matched_rule_str(&self) -> Option { self.matched_rule.as_ref().map(|r| r.as_str()) @@ -376,7 +368,7 @@ impl AclProcessor { let cleanup_interval = self.cache_cleanup_interval; self.tasks.spawn(async move { - let mut interval = tokio::time::interval(cleanup_interval); + let mut interval = crate::foundation::time::interval(cleanup_interval); loop { interval.tick().await; Self::cleanup_cache(&rule_cache, cache_max_size); @@ -389,7 +381,7 @@ impl AclProcessor { let conn_track = self.conn_track.clone(); self.tasks.spawn(async move { - let mut interval = tokio::time::interval(cleanup_interval); + let mut interval = crate::foundation::time::interval(cleanup_interval); loop { interval.tick().await; Self::cleanup_expired_connections(conn_track.clone(), 60); @@ -514,12 +506,10 @@ impl AclProcessor { }; let mut cache_entry = AclCacheEntry { - action: Action::Allow, matched_rule: RuleId::Default, last_access: Instant::now(), conn_track_key: None, rate_limit_keys: vec![], - chain_type, acl_result: None, rule_stats_vec: vec![], }; @@ -927,42 +917,8 @@ impl AclProcessor { conn_track.remove(&key); } } - - /// Get cache hit rate - pub fn get_cache_hit_rate(&self) -> f64 { - let cache_hits = self - .stats - .get(&AclStatKey::CacheHits) - .map(|v| *v.value()) - .unwrap_or(0); - let total_requests = cache_hits - + self - .stats - .get(&AclStatKey::RuleMatches) - .map(|v| *v.value()) - .unwrap_or(0); - - if total_requests == 0 { - 0.0 - } else { - cache_hits as f64 / total_requests as f64 - } - } } -// 新增辅助函数 -fn parse_port_start(port_strs: &[String]) -> Option { - port_strs - .iter() - .filter_map(|s| parse_port_range(s).map(|(start, _)| start)) - .min() -} -fn parse_port_end(port_strs: &[String]) -> Option { - port_strs - .iter() - .filter_map(|s| parse_port_range(s).map(|(_, end)| end)) - .max() -} fn parse_port_range(s: &str) -> Option<(u16, u16)> { if let Some((start, end)) = s.split_once('-') { let start = start.trim().parse().ok()?; @@ -1043,190 +999,6 @@ impl AclStatKey { } } -pub struct AclRuleBuilder { - pub acl: Option, - pub tcp_whitelist: Vec, - pub udp_whitelist: Vec, - pub whitelist_priority: Option, -} - -impl AclRuleBuilder { - fn parse_port_list(port_list: &[String]) -> anyhow::Result> { - let mut ports = Vec::new(); - - for port_spec in port_list { - if port_spec.contains('-') { - // Handle port range like "8000-9000" - let parts: Vec<&str> = port_spec.split('-').collect(); - if parts.len() != 2 { - return Err(anyhow::anyhow!("Invalid port range format: {}", port_spec)); - } - - let start: u16 = parts[0] - .parse() - .with_context(|| format!("Invalid start port in range: {}", port_spec))?; - let end: u16 = parts[1] - .parse() - .with_context(|| format!("Invalid end port in range: {}", port_spec))?; - - if start > end { - return Err(anyhow::anyhow!( - "Start port must be <= end port in range: {}", - port_spec - )); - } - - // acl can handle port range - ports.push(port_spec.clone()); - } else { - // Handle single port - let port: u16 = port_spec - .parse() - .with_context(|| format!("Invalid port number: {}", port_spec))?; - ports.push(port.to_string()); - } - } - - Ok(ports) - } - - fn generate_acl_from_whitelists(&mut self) -> anyhow::Result<()> { - if self.tcp_whitelist.is_empty() && self.udp_whitelist.is_empty() { - return Ok(()); - } - - // Create inbound chain for whitelist rules - let mut inbound_chain = Chain { - name: "inbound_whitelist".to_string(), - chain_type: ChainType::Inbound as i32, - description: "Auto-generated inbound whitelist from CLI".to_string(), - enabled: true, - rules: vec![], - default_action: Action::Allow as i32, - }; - - let mut rule_priority = self.whitelist_priority.unwrap_or(1000u32); - - // Add TCP whitelist rules - if !self.tcp_whitelist.is_empty() { - let tcp_ports = Self::parse_port_list(&self.tcp_whitelist)?; - let tcp_rule = Rule { - name: "tcp_whitelist".to_string(), - description: "Auto-generated TCP whitelist rule".to_string(), - priority: rule_priority, - enabled: true, - protocol: Protocol::Tcp as i32, - ports: tcp_ports, - source_ips: vec![], - destination_ips: vec![], - source_ports: vec![], - action: Action::Allow as i32, - rate_limit: 0, - burst_limit: 0, - stateful: true, - source_groups: vec![], - destination_groups: vec![], - }; - let tcp_rule_deny_other = Rule { - name: "tcp_whitelist_deny_other".to_string(), - description: "Auto-generated TCP whitelist rule to deny other ports".to_string(), - priority: 0, - enabled: true, - protocol: Protocol::Tcp as i32, - ports: vec!["0-65535".to_string()], - source_ips: vec![], - destination_ips: vec![], - source_ports: vec![], - action: Action::Drop as i32, - rate_limit: 0, - burst_limit: 0, - stateful: false, - source_groups: vec![], - destination_groups: vec![], - }; - inbound_chain.rules.push(tcp_rule); - inbound_chain.rules.push(tcp_rule_deny_other); - rule_priority -= 1; - } - - // Add UDP whitelist rules - if !self.udp_whitelist.is_empty() { - let udp_ports = Self::parse_port_list(&self.udp_whitelist)?; - let udp_rule = Rule { - name: "udp_whitelist".to_string(), - description: "Auto-generated UDP whitelist rule".to_string(), - priority: rule_priority, - enabled: true, - protocol: Protocol::Udp as i32, - ports: udp_ports, - source_ips: vec![], - destination_ips: vec![], - source_ports: vec![], - action: Action::Allow as i32, - rate_limit: 0, - burst_limit: 0, - stateful: false, - source_groups: vec![], - destination_groups: vec![], - }; - let udp_rule_deny_other = Rule { - name: "udp_whitelist_deny_other".to_string(), - description: "Auto-generated UDP whitelist rule to deny other ports".to_string(), - priority: 0, - enabled: true, - protocol: Protocol::Udp as i32, - ports: vec!["0-65535".to_string()], - source_ips: vec![], - destination_ips: vec![], - source_ports: vec![], - action: Action::Drop as i32, - rate_limit: 0, - burst_limit: 0, - stateful: false, - source_groups: vec![], - destination_groups: vec![], - }; - inbound_chain.rules.push(udp_rule); - inbound_chain.rules.push(udp_rule_deny_other); - } - - if self.acl.is_none() { - self.acl = Some(Acl::default()); - } - - let acl = self.acl.as_mut().unwrap(); - - if let Some(ref mut acl_v1) = acl.acl_v1 { - acl_v1.chains.push(inbound_chain); - } else { - acl.acl_v1 = Some(AclV1 { - chains: vec![inbound_chain], - group: Some(GroupInfo { - declares: vec![], - members: vec![], - }), - }); - } - - Ok(()) - } - - fn do_build(mut self) -> anyhow::Result> { - self.generate_acl_from_whitelists()?; - Ok(self.acl.clone()) - } - - pub fn build(global_ctx: &ArcGlobalCtx) -> anyhow::Result> { - let builder = AclRuleBuilder { - acl: global_ctx.config.get_acl(), - tcp_whitelist: global_ctx.config.get_tcp_whitelist(), - udp_whitelist: global_ctx.config.get_udp_whitelist(), - whitelist_priority: None, - }; - builder.do_build() - } -} - #[derive(Debug, Clone, Copy)] pub enum AclStatType { Total, @@ -1241,6 +1013,28 @@ mod tests { use std::hash::{Hash, Hasher}; use std::net::{IpAddr, Ipv4Addr}; + impl AclProcessor { + fn get_cache_hit_rate(&self) -> f64 { + let cache_hits = self + .stats + .get(&AclStatKey::CacheHits) + .map(|v| *v.value()) + .unwrap_or(0); + let total_requests = cache_hits + + self + .stats + .get(&AclStatKey::RuleMatches) + .map(|v| *v.value()) + .unwrap_or(0); + + if total_requests == 0 { + 0.0 + } else { + cache_hits as f64 / total_requests as f64 + } + } + } + #[tokio::test] async fn test_group_based_acl_rules() { let mut acl_config = Acl::default(); diff --git a/easytier-core/src/peers/admission.rs b/easytier-core/src/peers/admission.rs new file mode 100644 index 00000000..66c690d6 --- /dev/null +++ b/easytier-core/src/peers/admission.rs @@ -0,0 +1,131 @@ +use std::sync::{Arc, Weak}; + +use async_trait::async_trait; + +use crate::{ + connectivity::protocol::raw, + events::{CoreEvent, CoreEventSink}, + listener::{ + AcceptedSocketHandler, + transport::{AcceptedTransport, AcceptedTunnelHandler}, + }, + tunnel::Tunnel, +}; + +use super::peer_manager::PeerManagerCore; + +pub(crate) struct PeerAcceptedTunnelHandler { + peer_manager: Weak, + events: Arc, +} + +impl PeerAcceptedTunnelHandler { + pub(crate) fn new( + peer_manager: &Arc, + events: Arc, + ) -> Arc { + Arc::new(Self { + peer_manager: Arc::downgrade(peer_manager), + events, + }) + } +} + +#[async_trait] +impl AcceptedTunnelHandler for PeerAcceptedTunnelHandler { + async fn handle_tunnel(&self, tunnel: Box) -> anyhow::Result<()> { + let tunnel_info = tunnel + .info() + .ok_or_else(|| anyhow::anyhow!("accepted tunnel has no tunnel info"))?; + let local_url = tunnel_info + .local_addr + .clone() + .unwrap_or_default() + .to_string(); + let remote_url = tunnel_info + .remote_addr + .clone() + .unwrap_or_default() + .to_string(); + self.events.emit(CoreEvent::TunnelAccepted { + local_url: local_url.clone(), + remote_url: remote_url.clone(), + }); + tracing::info!(ret = ?tunnel, "conn accepted"); + + let Some(peer_manager) = self.peer_manager.upgrade() else { + let error = "peer manager is gone, cannot handle tunnel".to_owned(); + self.events.emit(CoreEvent::TunnelAdmissionFailed { + local_url, + remote_url, + error: error.clone(), + }); + tracing::error!(error = %error, "handle conn error"); + return Err(anyhow::anyhow!(error)); + }; + if let Err(error) = peer_manager.add_tunnel_as_server(tunnel, true).await { + self.events.emit(CoreEvent::TunnelAdmissionFailed { + local_url, + remote_url, + error: error.to_string(), + }); + tracing::error!(?error, "handle conn error"); + return Err(error.into()); + } + Ok(()) + } +} + +pub(crate) struct RawAcceptedTransportHandler { + peer_manager: Weak, +} + +impl RawAcceptedTransportHandler { + pub(crate) fn new(peer_manager: &Arc) -> Self { + Self { + peer_manager: Arc::downgrade(peer_manager), + } + } +} + +#[async_trait] +impl AcceptedSocketHandler> for RawAcceptedTransportHandler +where + TcpSocket: crate::socket::tcp::VirtualTcpSocket, +{ + async fn handle_accepted_socket( + &self, + accepted: AcceptedTransport, + ) -> anyhow::Result<()> { + let peer_manager = self + .peer_manager + .upgrade() + .ok_or_else(|| anyhow::anyhow!("peer manager is gone"))?; + let tunnel = match accepted { + AcceptedTransport::Tunnel { tunnel, .. } => tunnel, + AcceptedTransport::Tcp { + socket, local_url, .. + } => { + if local_url.scheme() != "tcp" { + anyhow::bail!("unsupported raw TCP listener protocol: {local_url}"); + } + raw::upgrade_accepted_tcp_with_local_url(socket, local_url)? + } + AcceptedTransport::Udp { + session, local_url, .. + } => { + if local_url.scheme() != "udp" { + anyhow::bail!("unsupported raw UDP listener protocol: {local_url}"); + } + raw::upgrade_accepted_udp_with_local_url(session, local_url)? + } + AcceptedTransport::ByteStream { + socket, + local_url, + remote_url, + } => raw::upgrade_accepted_byte_stream(socket, local_url, remote_url)?, + }; + peer_manager.add_tunnel_as_server(tunnel, true).await?; + Ok(()) + } +} diff --git a/easytier-core/src/peers/conn/mod.rs b/easytier-core/src/peers/conn/mod.rs new file mode 100644 index 00000000..5a407ba9 --- /dev/null +++ b/easytier-core/src/peers/conn/mod.rs @@ -0,0 +1,8 @@ +//! Peer connection primitives: noise sessions, individual peer connections, +//! and the peer map that multiplexes them. + +pub(crate) mod peer; +pub(crate) mod peer_conn; +pub(crate) mod peer_conn_ping; +pub(crate) mod peer_map; +pub(crate) mod peer_session; diff --git a/easytier-core/src/peers/conn/peer.rs b/easytier-core/src/peers/conn/peer.rs new file mode 100644 index 00000000..04daa012 --- /dev/null +++ b/easytier-core/src/peers/conn/peer.rs @@ -0,0 +1,290 @@ +use std::sync::Arc; + +use crossbeam::atomic::AtomicCell; +use dashmap::{DashMap, DashSet}; +use parking_lot::RwLock; + +use tokio::{select, sync::mpsc}; + +use tracing::Instrument; + +use super::peer_conn::{PeerConn, PeerConnId}; +use crate::peers::{ + PacketRecvChan, + context::{ArcPeerContext, PeerEvent}, + util::shrink_dashmap, +}; +use crate::{ + config::PeerId, + packet::ZCPacket, + peers::error::Error, + proto::{core_peer::peer::PeerConnInfo, peer_rpc::PeerIdentityType}, +}; +use tokio_util::task::AbortOnDropHandle; + +type ArcPeerConn = Arc; +type ConnMap = Arc>; + +pub struct Peer { + pub peer_node_id: PeerId, + conns: ConnMap, + context: ArcPeerContext, + + packet_recv_chan: PacketRecvChan, + + close_event_sender: mpsc::Sender, + #[allow(dead_code)] + close_event_listener: AbortOnDropHandle<()>, + + shutdown_notifier: Arc, + + default_conn_id: Arc>, + peer_identity_type: Arc>>, + peer_public_key: Arc>>>, + #[allow(dead_code)] + default_conn_id_clear_task: AbortOnDropHandle<()>, +} + +impl Peer { + pub(crate) fn new( + peer_node_id: PeerId, + packet_recv_chan: PacketRecvChan, + context: ArcPeerContext, + ) -> Self { + let conns: ConnMap = Arc::new(DashMap::new()); + let (close_event_sender, mut close_event_receiver) = mpsc::channel(10); + let shutdown_notifier = Arc::new(tokio::sync::Notify::new()); + let peer_identity_type = Arc::new(AtomicCell::new(None)); + let peer_identity_type_copy = peer_identity_type.clone(); + let peer_public_key = Arc::new(RwLock::new(None)); + let peer_public_key_copy = peer_public_key.clone(); + + let conns_copy = conns.clone(); + let shutdown_notifier_copy = shutdown_notifier.clone(); + let context_copy = context.clone(); + let close_event_listener = AbortOnDropHandle::new(tokio::spawn( + async move { + loop { + select! { + ret = close_event_receiver.recv() => { + if ret.is_none() { + break; + } + let ret = ret.unwrap(); + tracing::warn!( + ?peer_node_id, + ?ret, + "notified that peer conn is closed", + ); + + if let Some((_, conn)) = conns_copy.remove(&ret) { + context_copy.issue_event(PeerEvent::PeerConnRemoved( + conn.get_conn_info(), + )); + shrink_dashmap(&conns_copy, Some(4)); + if conns_copy.is_empty() { + peer_identity_type_copy.store(None); + *peer_public_key_copy.write() = None; + } + } + } + + _ = shutdown_notifier_copy.notified() => { + close_event_receiver.close(); + tracing::warn!(?peer_node_id, "peer close event listener notified"); + } + } + } + tracing::info!("peer {} close event listener exit", peer_node_id); + } + .instrument(tracing::info_span!( + "peer_close_event_listener", + ?peer_node_id, + )), + )); + + let default_conn_id = Arc::new(AtomicCell::new(PeerConnId::default())); + + let conns_copy = conns.clone(); + let default_conn_id_copy = default_conn_id.clone(); + let default_conn_id_clear_task = AbortOnDropHandle::new(tokio::spawn(async move { + loop { + crate::foundation::time::sleep(std::time::Duration::from_secs(5)).await; + if conns_copy.len() > 1 { + default_conn_id_copy.store(PeerConnId::default()); + } + } + })); + + Peer { + peer_node_id, + conns, + packet_recv_chan, + context, + + close_event_sender, + close_event_listener, + + shutdown_notifier, + default_conn_id, + peer_identity_type, + peer_public_key, + default_conn_id_clear_task, + } + } + + pub async fn add_peer_conn(&self, mut conn: PeerConn) -> Result<(), Error> { + let conn_identity_type = conn.get_peer_identity_type(); + let peer_identity_type = self.peer_identity_type.load(); + if let Some(peer_identity_type) = peer_identity_type { + if peer_identity_type != conn_identity_type { + return Err(Error::SecretKeyError(format!( + "peer identity type mismatch. peer: {:?}, conn: {:?}", + peer_identity_type, conn_identity_type + ))); + } + } else { + self.peer_identity_type.store(Some(conn_identity_type)); + } + + let close_notifier = conn.get_close_notifier(); + let conn_info = conn.get_conn_info(); + let conn_pubkey = conn_info.noise_remote_static_pubkey.clone(); + { + let mut peer_pubkey = self.peer_public_key.write(); + if let Some(existing_pubkey) = peer_pubkey.as_ref() { + if existing_pubkey != &conn_pubkey { + return Err(Error::SecretKeyError(format!( + "peer public key mismatch. peer_id: {}, existing_len: {}, new_len: {}", + self.peer_node_id, + existing_pubkey.len(), + conn_pubkey.len() + ))); + } + } else { + *peer_pubkey = Some(conn_pubkey); + } + } + + conn.start_recv_loop(self.packet_recv_chan.clone()).await; + conn.start_pingpong(); + self.conns.insert(conn.get_conn_id(), Arc::new(conn)); + + let close_event_sender = self.close_event_sender.clone(); + tokio::spawn(async move { + let conn_id = close_notifier.get_conn_id(); + if let Some(mut waiter) = close_notifier.get_waiter().await { + let _ = waiter.recv().await; + } + if let Err(e) = close_event_sender.send(conn_id).await { + tracing::warn!(?conn_id, "failed to send close event: {}", e); + } + }); + + self.context + .issue_event(PeerEvent::PeerConnAdded(conn_info)); + Ok(()) + } + + async fn select_conn(&self) -> Option { + let default_conn_id = self.default_conn_id.load(); + if let Some(conn) = self.conns.get(&default_conn_id) { + return Some(conn.clone()); + } + + // find a conn with the smallest latency + let mut min_latency = u64::MAX; + for conn in self.conns.iter() { + let latency = conn.value().get_stats().latency_us; + if latency < min_latency { + min_latency = latency; + self.default_conn_id.store(conn.get_conn_id()); + } + } + + self.conns + .get(&self.default_conn_id.load()) + .map(|conn| conn.clone()) + } + + pub async fn send_msg(&self, msg: ZCPacket) -> Result<(), Error> { + let Some(conn) = self.select_conn().await else { + return Err(Error::PeerNoConnectionError(self.peer_node_id)); + }; + conn.send_msg(msg).await?; + + Ok(()) + } + + pub async fn close_peer_conn(&self, conn_id: &PeerConnId) -> Result<(), Error> { + let has_key = self.conns.contains_key(conn_id); + if !has_key { + return Err(Error::NotFound); + } + self.close_event_sender.send(*conn_id).await.unwrap(); + Ok(()) + } + + pub async fn list_peer_conns(&self) -> Vec { + let mut conns = vec![]; + for conn in self.conns.iter() { + // do not lock here, otherwise it will cause dashmap deadlock + conns.push(conn.clone()); + } + + let mut ret = Vec::new(); + for conn in conns { + let info = conn.get_conn_info(); + if !info.is_closed { + ret.push(info); + } else { + let conn_id = info.conn_id.parse().unwrap(); + let _ = self.close_peer_conn(&conn_id).await; + } + } + ret + } + + pub fn has_live_conns(&self) -> bool { + self.conns.iter().any(|entry| !entry.value().is_closed()) + } + + pub fn has_directly_connected_conn(&self) -> bool { + self.conns + .iter() + .any(|entry| !entry.value().is_closed() && !entry.value().is_hole_punched()) + } + + pub fn get_directly_connections(&self) -> DashSet { + self.conns + .iter() + .filter(|entry| !(entry.value()).is_hole_punched()) + .map(|entry| (entry.value()).get_conn_id()) + .collect() + } + + pub fn get_default_conn_id(&self) -> PeerConnId { + self.default_conn_id.load() + } + + pub fn get_peer_identity_type(&self) -> Option { + self.peer_identity_type.load() + } + + pub fn get_peer_public_key(&self) -> Option> { + self.peer_public_key.read().clone() + } +} + +// pritn on drop +impl Drop for Peer { + fn drop(&mut self) { + self.conns.retain(|_, conn| { + self.context + .issue_event(PeerEvent::PeerConnRemoved(conn.get_conn_info())); + false + }); + self.shutdown_notifier.notify_one(); + tracing::info!("peer {} drop", self.peer_node_id); + } +} diff --git a/easytier/src/peers/peer_conn.rs b/easytier-core/src/peers/conn/peer_conn.rs similarity index 54% rename from easytier/src/peers/peer_conn.rs rename to easytier-core/src/peers/conn/peer_conn.rs index 8b9b6fe6..867d3da5 100644 --- a/easytier/src/peers/peer_conn.rs +++ b/easytier-core/src/peers/conn/peer_conn.rs @@ -19,34 +19,31 @@ use guarden::guard; use hmac::Mac; use prost::Message; -use tokio::{ - sync::broadcast, - task::JoinSet, - time::{Duration, timeout}, -}; +use tokio::{sync::broadcast, task::JoinSet}; use tracing::Instrument; use zerocopy::AsBytes; use snow::{HandshakeState, params::NoiseParams}; +use crate::foundation::time::{Duration, timeout}; + use super::{ - PacketRecvChan, peer_conn_ping::PeerConnPinger, peer_session::{PeerSession, PeerSessionAction}, - traffic_metrics::AggregateTrafficMetrics, +}; +use crate::peers::{ + PacketRecvChan, + context::{ArcPeerContext, NetworkIdentity, NetworkSecretDigest}, }; use crate::{ - common::{ - PeerId, - config::{NetworkIdentity, NetworkSecretDigest}, - error::Error, - global_ctx::ArcGlobalCtx, - }, - peers::peer_session::{PeerSessionStore, SessionKey, UpsertResponderSessionReturn}, + config::PeerId, + packet::{PacketType, ZCPacket}, + peers::conn::peer_session::{PeerSessionStore, SessionKey, UpsertResponderSessionReturn}, + peers::error::Error, proto::{ - api::instance::{PeerConnInfo, PeerConnStats}, - common::{LimiterConfig, SecureModeConfig, TunnelInfo}, + common::{SecureModeConfig, TunnelInfo}, + core_peer::peer::{PeerConnInfo, PeerConnStats}, peer_rpc::{ HandshakeRequest, PeerConnNoiseMsg1Pb, PeerConnNoiseMsg2Pb, PeerConnNoiseMsg3Pb, PeerConnSessionActionPb, PeerIdentityType, SecureAuthLevel, @@ -56,10 +53,8 @@ use crate::{ Tunnel, TunnelError, ZCPacketStream, filter::{StatsRecorderTunnelFilter, TunnelFilter, TunnelFilterChain, TunnelWithFilter}, mpsc::{MpscTunnel, MpscTunnelSender}, - packet_def::{PacketType, ZCPacket}, stats::{Throughput, WindowLatency}, }, - use_global_var, }; pub type PeerConnId = uuid::Uuid; @@ -76,12 +71,12 @@ struct SecretProof { /// The result of noise handshake. #[derive(Debug)] +#[allow(dead_code)] struct NoiseHandshakeResult { peer_id: PeerId, session: Arc, local_static_pubkey: Vec, remote_static_pubkey: Vec, - handshake_hash: Vec, secure_auth_level: SecureAuthLevel, peer_identity_type: PeerIdentityType, remote_network_name: String, @@ -91,9 +86,6 @@ struct NoiseHandshakeResult { // foreign network manager use this to verify peer. // the challenge will be sent to authorized peer and compare the proof against it. client_secret_proof: Option, - - my_encrypt_algo: String, - remote_encrypt_algo: String, } #[derive(Clone)] @@ -105,15 +97,6 @@ struct PeerSessionTunnelFilter { } impl PeerSessionTunnelFilter { - fn new(enabled: bool) -> Self { - Self { - enabled, - my_peer_id: Arc::new(AtomicCell::new(PeerId::default())), - peer_id: Arc::new(AtomicCell::new(None)), - session: Arc::new(ArcSwapOption::empty()), - } - } - fn new_with_peer(my_peer_id: PeerId, enabled: bool) -> Self { Self { enabled, @@ -135,7 +118,7 @@ impl PeerSessionTunnelFilter { self.session.store(Some(session)); } - fn should_skip_encrypt(&self, hdr: &crate::tunnel::packet_def::PeerManagerHeader) -> bool { + fn should_skip_encrypt(&self, hdr: &crate::packet::PeerManagerHeader) -> bool { hdr.packet_type == PacketType::NoiseHandshakeMsg1 as u8 || hdr.packet_type == PacketType::NoiseHandshakeMsg2 as u8 || hdr.packet_type == PacketType::NoiseHandshakeMsg3 as u8 @@ -288,12 +271,13 @@ pub struct PeerConn { my_peer_id: PeerId, peer_id_hint: Option, - global_ctx: ArcGlobalCtx, + context: ArcPeerContext, secure_mode_cfg: Option, session_filter: PeerSessionTunnelFilter, noise_handshake_result: Option, + #[allow(dead_code)] tunnel: Arc>>, sink: MpscTunnelSender, recv: Mutex>>>, @@ -330,27 +314,27 @@ impl Debug for PeerConn { } impl PeerConn { - pub fn new( + pub(crate) fn new( my_peer_id: PeerId, - global_ctx: ArcGlobalCtx, + context: ArcPeerContext, tunnel: Box, peer_session_store: Arc, ) -> Self { - Self::new_with_peer_id_hint(my_peer_id, global_ctx, tunnel, None, peer_session_store) + Self::new_with_peer_id_hint(my_peer_id, context, tunnel, None, peer_session_store) } - pub fn new_with_peer_id_hint( + pub(crate) fn new_with_peer_id_hint( my_peer_id: PeerId, - global_ctx: ArcGlobalCtx, + context: ArcPeerContext, tunnel: Box, peer_id_hint: Option, peer_session_store: Arc, ) -> Self { - let flags = global_ctx.get_flags(); + let flags = context.flags(); let tunnel_info = tunnel.info(); let (ctrl_sender, _ctrl_receiver) = broadcast::channel(8); - let secure_mode_cfg = global_ctx.config.get_secure_mode(); + let secure_mode_cfg = context.secure_mode(); let session_filter = PeerSessionTunnelFilter::new_with_peer( my_peer_id, secure_mode_cfg @@ -375,7 +359,7 @@ impl PeerConn { my_peer_id, peer_id_hint, - global_ctx, + context, secure_mode_cfg, session_filter, @@ -524,7 +508,7 @@ impl PeerConn { send_secret_digest: bool, metric_network_name: &str, ) -> Result<(), Error> { - let network = self.global_ctx.get_network_identity(); + let network = self.context.network_identity(); let mut req = HandshakeRequest { magic: MAGIC, my_peer_id: self.my_peer_id, @@ -537,7 +521,7 @@ impl PeerConn { // only send network secret digest if the network is the same if send_secret_digest { req.network_secret_digest - .extend_from_slice(&network.network_secret_digest.unwrap_or_default()); + .extend_from_slice(&network.secret_digest().unwrap_or_default()); } else { // fill zero req.network_secret_digest @@ -641,19 +625,8 @@ impl PeerConn { } fn get_pinned_remote_static_pubkey_b64(&self) -> Option { - let remote_url_str = self - .tunnel_info - .as_ref() - .and_then(|t| t.remote_addr.as_ref()) - .map(|u| u.url.as_str())?; - let remote_url: url::Url = remote_url_str.parse().ok()?; - - self.global_ctx - .config - .get_peers() - .into_iter() - .find(|p| p.uri == remote_url) - .and_then(|p| p.peer_public_key) + self.context + .pinned_remote_static_pubkey(self.tunnel_info.as_ref()) } async fn send_noise_msg( @@ -718,7 +691,7 @@ impl PeerConn { ) -> Result { // 1. Verify proof if let Some(proof) = proof - && let Some(mac) = self.global_ctx.get_secret_proof(handshake_hash) + && let Some(mac) = self.context.secret_proof(handshake_hash) && mac.verify_slice(proof).is_ok() { return Ok(SecureAuthLevel::NetworkSecretConfirmed); @@ -734,7 +707,7 @@ impl PeerConn { // If no network_secret, pinned key must be in trusted list if !has_network_secret && !self - .global_ctx + .context .is_pubkey_trusted(remote_pubkey, remote_network_name) { return Err(Error::WaitRespError( @@ -746,7 +719,7 @@ impl PeerConn { // 3. Check if pubkey is in trusted list if self - .global_ctx + .context .is_pubkey_trusted(remote_pubkey, remote_network_name) { return Ok(SecureAuthLevel::PeerVerified); @@ -771,9 +744,7 @@ impl PeerConn { remote_sent_secret_proof: bool, is_client: bool, ) -> PeerIdentityType { - if !remote_role_hint_is_same_network - || remote_network_name != self.global_ctx.get_network_name() - { + if !remote_role_hint_is_same_network || remote_network_name != self.context.network_name() { if is_client { PeerIdentityType::SharedNode } else if remote_sent_secret_proof { @@ -807,7 +778,7 @@ impl PeerConn { let builder = snow::Builder::new(params); let (local_private_key, local_static_pubkey) = self.get_keypair()?; - let network = self.global_ctx.get_network_identity(); + let network = self.context.network_identity(); let a_session_generation = self .peer_id_hint .and_then(|peer_id| { @@ -876,18 +847,11 @@ impl PeerConn { let handshake_hash_for_proof = hs.get_handshake_hash().to_vec(); let secret_proof_32 = self - .global_ctx - .get_secret_proof(&handshake_hash_for_proof) + .context + .secret_proof(&handshake_hash_for_proof) .map(|mac| mac.finalize().into_bytes().to_vec()); - let secret_digest = if use_global_var!(HMAC_SECRET_DIGEST) { - self.global_ctx - .get_secret_proof("digest".as_bytes()) - .map(|mac| mac.finalize().into_bytes().to_vec()) - .unwrap_or_default() - } else { - network.network_secret_digest.unwrap_or_default().to_vec() - }; + let secret_digest = self.context.secret_digest(&network); let msg3_pb = PeerConnNoiseMsg3Pb { a_conn_id_echo: Some(a_conn_id.into()), @@ -938,9 +902,7 @@ impl PeerConn { true, ); - let handshake_hash = hs.get_handshake_hash().to_vec(); - - let algo = self.global_ctx.get_flags().encryption_algorithm.clone(); + let algo = self.context.flags().encryption_algorithm.clone(); let root_key = msg2_pb .root_key_32 .as_deref() @@ -971,16 +933,12 @@ impl PeerConn { session, local_static_pubkey: local_static_pubkey.to_vec(), remote_static_pubkey: remote_static, - handshake_hash, secure_auth_level, peer_identity_type, remote_network_name, // we have authorized the peer with noise handshake, so just set secret digest same as us even remote is a shared node. secret_digest, client_secret_proof: None, - - my_encrypt_algo: self.my_encrypt_algo.clone(), - remote_encrypt_algo: msg2_pb.server_encryption_algorithm.clone(), }) } @@ -1026,22 +984,6 @@ impl PeerConn { Ok(msg) } - async fn read_next_message_with_timeout( - &mut self, - read_timeout: Duration, - ) -> Result { - timeout(read_timeout, async { - let mut locked = self.recv.lock().await; - let recv = locked.as_mut().unwrap(); - Ok(recv - .next() - .await - .ok_or(Error::WaitRespError("read next message failed".to_owned()))??) - }) - .await - .map_err(|e| Error::WaitRespError(format!("read next message timeout: {e:?}")))? - } - async fn do_noise_handshake_as_server( &mut self, first_msg1: ZCPacket, @@ -1080,19 +1022,19 @@ impl PeerConn { // this may update my peer id handshake_recved(self, &remote_network_name)?; - let server_network_name = self.global_ctx.get_network_name(); + let server_network_name = self.context.network_name(); let (role_hint, secret_proof_32) = if msg1_pb.a_network_name == server_network_name { ( 1, - self.global_ctx - .get_secret_proof(hs.get_handshake_hash()) + self.context + .secret_proof(hs.get_handshake_hash()) .map(|m| m.finalize().into_bytes().to_vec()), ) } else { (2, None) }; - let algo = self.global_ctx.get_flags().encryption_algorithm.clone(); + let algo = self.context.flags().encryption_algorithm.clone(); let UpsertResponderSessionReturn { session, action, @@ -1179,10 +1121,7 @@ impl PeerConn { &handshake_hash_for_proof, &remote_static, None, // Server doesn't have pinned_remote_pubkey - self.global_ctx - .get_network_identity() - .network_secret - .is_some(), + self.context.network_identity().network_secret.is_some(), false, // is_initiator &remote_network_name, )? @@ -1197,14 +1136,11 @@ impl PeerConn { false, ); - let handshake_hash = hs.get_handshake_hash().to_vec(); - Ok(NoiseHandshakeResult { peer_id: remote_peer_id, session, local_static_pubkey: local_static_pubkey.to_vec(), remote_static_pubkey: remote_static, - handshake_hash, secure_auth_level, peer_identity_type, remote_network_name, @@ -1213,9 +1149,6 @@ impl PeerConn { challenge: handshake_hash_for_proof, proof: p.clone(), }), - - my_encrypt_algo: self.my_encrypt_algo.clone(), - remote_encrypt_algo: msg1_pb.client_encryption_algorithm.clone(), }) } @@ -1272,7 +1205,7 @@ impl PeerConn { self.info = Some(rsp); self.is_client = Some(false); - let send_digest = self.get_network_identity() == self.global_ctx.get_network_identity(); + let send_digest = self.get_network_identity() == self.context.network_identity(); self.send_handshake(send_digest, &self.get_network_identity().network_name) .await?; } else { @@ -1289,11 +1222,6 @@ impl PeerConn { } } - #[tracing::instrument] - pub async fn do_handshake_as_server(&mut self) -> Result<(), Error> { - self.do_handshake_as_server_ext(|_, _| Ok(())).await - } - #[tracing::instrument] pub async fn do_handshake_as_client(&mut self) -> Result<(), Error> { if self.is_secure_mode_enabled() { @@ -1306,7 +1234,7 @@ impl PeerConn { self.info = Some(handshake_rsp); self.is_client = Some(true); } else { - let network = self.global_ctx.get_network_identity(); + let network = self.context.network_identity(); self.send_handshake(true, &network.network_name).await?; tracing::info!("waiting for handshake request from server"); let rsp = self.wait_handshake_loop().await?; @@ -1324,23 +1252,12 @@ impl PeerConn { } } - pub fn handshake_done(&self) -> bool { - self.info.is_some() - } - - fn control_metrics(&self, network_name: &str) -> AggregateTrafficMetrics { - AggregateTrafficMetrics::control( - self.global_ctx.stats_manager().clone(), - network_name.to_string(), - ) - } - fn record_control_tx(&self, network_name: &str, bytes: u64) { - self.control_metrics(network_name).record_tx(bytes); + self.context.record_control_tx(network_name, bytes); } fn record_control_rx(&self, network_name: &str, bytes: u64) { - self.control_metrics(network_name).record_rx(bytes); + self.context.record_control_rx(network_name, bytes); } pub async fn start_recv_loop(&mut self, packet_recv_chan: PacketRecvChan) { @@ -1350,37 +1267,14 @@ impl PeerConn { let close_event_notifier = self.close_event_notifier.clone(); let ctrl_sender = self.ctrl_resp_sender.clone(); let conn_info_for_instrument = self.get_conn_info(); - let control_metrics = self.control_metrics(&conn_info_for_instrument.network_name); + let context = self.context.clone(); + let control_network_name = conn_info_for_instrument.network_name.clone(); - let is_foreign_network = conn_info_for_instrument.network_name - != self.global_ctx.get_network_identity().network_name; - let recv_limiter = if is_foreign_network - && self.global_ctx.get_flags().foreign_relay_bps_limit != u64::MAX - { - let relay_network_bps_limit = self.global_ctx.get_flags().foreign_relay_bps_limit; - let limiter_config = LimiterConfig { - burst_rate: None, - bps: Some(relay_network_bps_limit), - fill_duration_ms: None, - }; - Some(self.global_ctx.token_bucket_manager().get_or_create( - &format!("{}:recv", conn_info_for_instrument.network_name), - limiter_config.into(), - )) - } else if self.global_ctx.get_flags().instance_recv_bps_limit != u64::MAX { - let limiter_config = LimiterConfig { - burst_rate: None, - bps: Some(self.global_ctx.get_flags().instance_recv_bps_limit), - fill_duration_ms: None, - }; - Some( - self.global_ctx - .token_bucket_manager() - .get_or_create("instance:recv", limiter_config.into()), - ) - } else { - None - }; + let is_foreign_network = + conn_info_for_instrument.network_name != self.context.network_identity().network_name; + let recv_limiter = self + .context + .recv_limiter(&conn_info_for_instrument.network_name, is_foreign_network); self.tasks.spawn( async move { @@ -1404,15 +1298,15 @@ impl PeerConn { }; if peer_mgr_hdr.packet_type == PacketType::Ping as u8 { - control_metrics.record_rx(buf_len); + context.record_control_rx(&control_network_name, buf_len); peer_mgr_hdr.packet_type = PacketType::Pong as u8; if let Err(e) = sink.send(zc_packet).await { tracing::error!(?e, "peer conn send req error"); } else { - control_metrics.record_tx(buf_len); + context.record_control_tx(&control_network_name, buf_len); } } else if peer_mgr_hdr.packet_type == PacketType::Pong as u8 { - control_metrics.record_rx(buf_len); + context.record_control_rx(&control_network_name, buf_len); if let Err(e) = ctrl_sender.send(zc_packet) { tracing::error!(?e, "peer conn send ctrl resp error"); } @@ -1447,7 +1341,8 @@ impl PeerConn { self.latency_stats.clone(), self.loss_rate_stats.clone(), self.throughput.clone(), - self.control_metrics(&self.get_conn_info().network_name), + self.context.clone(), + self.get_conn_info().network_name, ); let close_event_notifier = self.close_event_notifier.clone(); @@ -1487,7 +1382,7 @@ impl PeerConn { fn network_secret_digest_is_empty(network: &NetworkIdentity) -> bool { network - .network_secret_digest + .secret_digest() .as_ref() .is_none_or(|digest| digest.iter().all(|byte| *byte == 0)) } @@ -1501,22 +1396,22 @@ impl PeerConn { return false; }; - self.global_ctx - .get_secret_proof(&secret_proof.challenge) + self.context + .secret_proof(&secret_proof.challenge) .is_some_and(|mac| mac.verify_slice(&secret_proof.proof).is_ok()) } - pub(crate) fn matches_local_network_secret(&self) -> bool { + pub fn matches_local_network_secret(&self) -> bool { if self.matches_local_secret_proof() { return true; } - let my_identity = self.global_ctx.get_network_identity(); + let my_identity = self.context.network_identity(); let peer_identity = self.get_network_identity(); !Self::network_secret_digest_is_empty(&my_identity) && !Self::network_secret_digest_is_empty(&peer_identity) - && my_identity.network_secret_digest == peer_identity.network_secret_digest + && my_identity.secret_digest() == peer_identity.secret_digest() } pub fn get_close_notifier(&self) -> Arc { @@ -1597,970 +1492,3 @@ impl Drop for PeerConn { self.close_event_notifier.notify_close(); } } - -#[cfg(test)] -pub mod tests { - use std::{sync::Arc, time::Duration}; - - use rand::rngs::OsRng; - - use super::*; - use crate::common::config::PeerConfig; - use crate::common::global_ctx::GlobalCtx; - use crate::common::global_ctx::tests::get_mock_global_ctx; - use crate::common::new_peer_id; - use crate::common::stats_manager::{LabelSet, LabelType, MetricName}; - use crate::peers::create_packet_recv_chan; - use crate::peers::recv_packet_from_chan; - use crate::tunnel::common::tests::wait_for_condition; - use crate::tunnel::filter::PacketRecorderTunnelFilter; - use crate::tunnel::filter::tests::DropSendTunnelFilter; - use crate::tunnel::ring::create_ring_tunnel_pair; - use tokio_util::task::AbortOnDropHandle; - - pub fn set_secure_mode_cfg(global_ctx: &GlobalCtx, enabled: bool) { - if !enabled { - global_ctx.config.set_secure_mode(None); - } else { - // generate x25519 key pair - let private = x25519_dalek::StaticSecret::random_from_rng(OsRng); - let public = x25519_dalek::PublicKey::from(&private); - - global_ctx.config.set_secure_mode(Some(SecureModeConfig { - enabled: true, - local_private_key: Some(BASE64_STANDARD.encode(private.as_bytes())), - local_public_key: Some(BASE64_STANDARD.encode(public.as_bytes())), - })); - } - } - - fn metric_value(global_ctx: &GlobalCtx, metric: MetricName, network_name: &str) -> u64 { - global_ctx - .stats_manager() - .get_metric( - metric, - &LabelSet::new().with_label_type(LabelType::NetworkName(network_name.to_string())), - ) - .map(|metric| metric.value) - .unwrap_or(0) - } - - #[test] - fn peer_session_filter_skips_relay_packet_for_next_hop() { - let my_peer_id = 10; - let next_hop_peer_id = 20; - let dst_peer_id = 30; - let filter = PeerSessionTunnelFilter::new_with_peer(my_peer_id, true); - filter.set_peer_id(next_hop_peer_id); - - let session = Arc::new(PeerSession::new( - next_hop_peer_id, - PeerSession::new_root_key(), - 1, - 0, - "aes-gcm".to_string(), - "aes-gcm".to_string(), - None, - )); - session.invalidate(); - filter.set_session(session); - - let mut packet = ZCPacket::new_with_payload(b"relay payload"); - packet.fill_peer_manager_hdr(my_peer_id, dst_peer_id, PacketType::Data as u8); - packet - .mut_peer_manager_header() - .unwrap() - .set_encrypted(true); - let original_len = packet.buf_len(); - - let packet = filter - .before_send(packet) - .expect("relay packet should bypass next-hop session"); - - let hdr = packet.peer_manager_header().unwrap(); - assert_eq!(hdr.from_peer_id.get(), my_peer_id); - assert_eq!(hdr.to_peer_id.get(), dst_peer_id); - assert!(hdr.is_encrypted()); - assert_eq!(packet.buf_len(), original_len); - } - - #[tokio::test] - async fn peer_conn_handshake_same_id() { - let ps = Arc::new(PeerSessionStore::new()); - let (c, s) = create_ring_tunnel_pair(); - let c_peer_id = new_peer_id(); - let s_peer_id = c_peer_id; - - let mut c_peer = PeerConn::new(c_peer_id, get_mock_global_ctx(), Box::new(c), ps.clone()); - let mut s_peer = PeerConn::new(s_peer_id, get_mock_global_ctx(), Box::new(s), ps.clone()); - - let (c_ret, s_ret) = tokio::join!( - c_peer.do_handshake_as_client(), - s_peer.do_handshake_as_server() - ); - - assert!(c_ret.is_err()); - assert!(s_ret.is_err()); - } - - #[tokio::test] - async fn peer_conn_handshake() { - let (c, s) = create_ring_tunnel_pair(); - - let c_recorder = Arc::new(PacketRecorderTunnelFilter::new()); - let s_recorder = Arc::new(PacketRecorderTunnelFilter::new()); - - let c = TunnelWithFilter::new(c, c_recorder.clone()); - let s = TunnelWithFilter::new(s, s_recorder.clone()); - - let c_peer_id = new_peer_id(); - let s_peer_id = new_peer_id(); - - let ps = Arc::new(PeerSessionStore::new()); - let c_ctx = get_mock_global_ctx(); - let s_ctx = get_mock_global_ctx(); - - let mut c_peer = PeerConn::new(c_peer_id, c_ctx.clone(), Box::new(c), ps.clone()); - - let mut s_peer = PeerConn::new(s_peer_id, s_ctx.clone(), Box::new(s), ps.clone()); - - let (c_ret, s_ret) = tokio::join!( - c_peer.do_handshake_as_client(), - s_peer.do_handshake_as_server() - ); - - c_ret.unwrap(); - s_ret.unwrap(); - - assert_eq!(c_recorder.sent.lock().unwrap().len(), 1); - assert_eq!(c_recorder.received.lock().unwrap().len(), 1); - - assert_eq!(s_recorder.sent.lock().unwrap().len(), 1); - assert_eq!(s_recorder.received.lock().unwrap().len(), 1); - - assert_eq!( - metric_value(&c_ctx, MetricName::TrafficControlBytesTx, "default"), - c_recorder - .sent - .lock() - .unwrap() - .iter() - .map(|pkt| pkt.buf_len() as u64) - .sum::() - ); - assert_eq!( - metric_value(&c_ctx, MetricName::TrafficControlBytesRx, "default"), - c_recorder - .received - .lock() - .unwrap() - .iter() - .map(|pkt| pkt.buf_len() as u64) - .sum::() - ); - assert_eq!( - metric_value(&s_ctx, MetricName::TrafficControlBytesTx, "default"), - s_recorder - .sent - .lock() - .unwrap() - .iter() - .map(|pkt| pkt.buf_len() as u64) - .sum::() - ); - assert_eq!( - metric_value(&s_ctx, MetricName::TrafficControlBytesRx, "default"), - s_recorder - .received - .lock() - .unwrap() - .iter() - .map(|pkt| pkt.buf_len() as u64) - .sum::() - ); - - assert_eq!(c_peer.get_peer_id(), s_peer_id); - assert_eq!(s_peer.get_peer_id(), c_peer_id); - assert_eq!(c_peer.get_network_identity(), s_peer.get_network_identity()); - assert_eq!( - c_peer.get_network_identity().network_name, - NetworkIdentity::default().network_name - ); - assert_eq!(c_peer.get_network_identity().network_secret, None); - assert_eq!( - c_peer.get_network_identity().network_secret_digest, - NetworkIdentity::default().network_secret_digest - ); - } - - #[tokio::test] - async fn peer_conn_secure_mode_pubkey_and_encryption() { - let (c, s) = create_ring_tunnel_pair(); - - let c_recorder = Arc::new(PacketRecorderTunnelFilter::new()); - let s_recorder = Arc::new(PacketRecorderTunnelFilter::new()); - - let c = TunnelWithFilter::new(c, c_recorder.clone()); - let s = TunnelWithFilter::new(s, s_recorder.clone()); - - let c_peer_id = new_peer_id(); - let s_peer_id = new_peer_id(); - - let c_ctx = get_mock_global_ctx(); - let s_ctx = get_mock_global_ctx(); - set_secure_mode_cfg(&c_ctx, true); - set_secure_mode_cfg(&s_ctx, true); - - let ps = Arc::new(PeerSessionStore::new()); - let mut c_peer = PeerConn::new(c_peer_id, c_ctx.clone(), Box::new(c), ps.clone()); - let mut s_peer = PeerConn::new(s_peer_id, s_ctx.clone(), Box::new(s), ps.clone()); - - let (c_ret, s_ret) = tokio::join!( - c_peer.do_handshake_as_client(), - s_peer.do_handshake_as_server() - ); - - c_ret.unwrap(); - s_ret.unwrap(); - - assert_eq!( - metric_value(&c_ctx, MetricName::TrafficControlBytesTx, "default"), - c_recorder - .sent - .lock() - .unwrap() - .iter() - .map(|pkt| pkt.buf_len() as u64) - .sum::() - ); - assert_eq!( - metric_value(&c_ctx, MetricName::TrafficControlBytesRx, "default"), - c_recorder - .received - .lock() - .unwrap() - .iter() - .map(|pkt| pkt.buf_len() as u64) - .sum::() - ); - assert_eq!( - metric_value(&s_ctx, MetricName::TrafficControlBytesTx, "default"), - s_recorder - .sent - .lock() - .unwrap() - .iter() - .map(|pkt| pkt.buf_len() as u64) - .sum::() - ); - assert_eq!( - metric_value(&s_ctx, MetricName::TrafficControlBytesRx, "default"), - s_recorder - .received - .lock() - .unwrap() - .iter() - .map(|pkt| pkt.buf_len() as u64) - .sum::() - ); - - let c_info = c_peer.get_conn_info(); - let s_info = s_peer.get_conn_info(); - - assert_eq!(c_info.noise_local_static_pubkey.len(), 32); - assert_eq!(c_info.noise_remote_static_pubkey.len(), 32); - assert_eq!(s_info.noise_local_static_pubkey.len(), 32); - assert_eq!(s_info.noise_remote_static_pubkey.len(), 32); - - assert_eq!( - c_info.noise_remote_static_pubkey, - s_info.noise_local_static_pubkey - ); - assert_eq!( - s_info.noise_remote_static_pubkey, - c_info.noise_local_static_pubkey - ); - - let network = s_ctx.get_network_identity(); - let mut expected = HandshakeRequest { - magic: MAGIC, - my_peer_id: s_peer_id, - version: VERSION, - features: Vec::new(), - network_name: network.network_name.clone(), - ..Default::default() - }; - expected - .network_secret_digest - .extend_from_slice(&network.network_secret_digest.unwrap_or_default()); - let expected_payload = expected.encode_to_vec(); - - println!("sent: {:?}", c_recorder.sent.lock().unwrap()); - - let wire_hs = c_recorder - .sent - .lock() - .unwrap() - .iter() - .find(|p| { - p.peer_manager_header() - .is_some_and(|h| h.packet_type == PacketType::NoiseHandshakeMsg3 as u8) - }) - .unwrap() - .clone(); - assert_ne!(wire_hs.payload(), expected_payload.as_slice()); - } - - #[tokio::test] - async fn peer_conn_secure_mode_server_accept_legacy_client() { - let (c, s) = create_ring_tunnel_pair(); - - let c_peer_id = new_peer_id(); - let s_peer_id = new_peer_id(); - - let c_ctx = get_mock_global_ctx(); - let s_ctx = get_mock_global_ctx(); - - c_ctx - .config - .set_network_identity(NetworkIdentity::new("user".to_string(), "sec1".to_string())); - s_ctx.config.set_network_identity(NetworkIdentity { - network_name: "shared".to_string(), - network_secret: None, - network_secret_digest: None, - }); - set_secure_mode_cfg(&s_ctx, true); - - let ps = Arc::new(PeerSessionStore::new()); - let mut c_peer = PeerConn::new(c_peer_id, c_ctx, Box::new(c), ps.clone()); - let mut s_peer = PeerConn::new(s_peer_id, s_ctx, Box::new(s), ps.clone()); - - let (c_ret, s_ret) = tokio::join!( - c_peer.do_handshake_as_client(), - s_peer.do_handshake_as_server() - ); - - c_ret.unwrap(); - s_ret.unwrap(); - - assert_eq!( - c_peer.get_conn_info().secure_auth_level, - SecureAuthLevel::None as i32, - ); - assert_eq!( - s_peer.get_conn_info().secure_auth_level, - SecureAuthLevel::None as i32, - ); - - assert_eq!(c_peer.get_conn_info().network_name, "shared".to_string()); - assert_eq!(s_peer.get_conn_info().network_name, "user".to_string()); - } - - #[tokio::test] - async fn peer_conn_secure_mode_different_network_name_ok() { - let (c, s) = create_ring_tunnel_pair(); - - let c_peer_id = new_peer_id(); - let s_peer_id = new_peer_id(); - - let c_ctx = get_mock_global_ctx(); - let s_ctx = get_mock_global_ctx(); - - c_ctx - .config - .set_network_identity(NetworkIdentity::new("user".to_string(), "sec1".to_string())); - s_ctx.config.set_network_identity(NetworkIdentity::new( - "shared".to_string(), - "sec2".to_string(), - )); - - set_secure_mode_cfg(&c_ctx, true); - set_secure_mode_cfg(&s_ctx, true); - - let ps = Arc::new(PeerSessionStore::new()); - let mut c_peer = PeerConn::new(c_peer_id, c_ctx, Box::new(c), ps.clone()); - let mut s_peer = PeerConn::new(s_peer_id, s_ctx, Box::new(s), ps.clone()); - - let (c_ret, s_ret) = tokio::join!( - c_peer.do_handshake_as_client(), - s_peer.do_handshake_as_server() - ); - c_ret.unwrap(); - s_ret.unwrap(); - - assert_eq!( - c_peer.get_conn_info().secure_auth_level, - SecureAuthLevel::EncryptedUnauthenticated as i32, - ); - assert_eq!( - s_peer.get_conn_info().secure_auth_level, - SecureAuthLevel::EncryptedUnauthenticated as i32, - ); - - assert_eq!(c_peer.get_conn_info().network_name, "shared".to_string()); - assert_eq!(s_peer.get_conn_info().network_name, "user".to_string()); - } - - #[tokio::test] - async fn peer_conn_secure_mode_data_roundtrip() { - let (c, s) = create_ring_tunnel_pair(); - - let c_peer_id = new_peer_id(); - let s_peer_id = new_peer_id(); - - let c_ctx = get_mock_global_ctx(); - let s_ctx = get_mock_global_ctx(); - set_secure_mode_cfg(&c_ctx, true); - set_secure_mode_cfg(&s_ctx, true); - - let ps = Arc::new(PeerSessionStore::new()); - let mut c_peer = PeerConn::new(c_peer_id, c_ctx, Box::new(c), ps.clone()); - let mut s_peer = PeerConn::new(s_peer_id, s_ctx, Box::new(s), ps.clone()); - - let (c_ret, s_ret) = tokio::join!( - c_peer.do_handshake_as_client(), - s_peer.do_handshake_as_server() - ); - c_ret.unwrap(); - s_ret.unwrap(); - - let (packet_send, mut packet_recv) = create_packet_recv_chan(); - s_peer.start_recv_loop(packet_send).await; - - let payload = b"secure-data-123"; - let mut pkt = ZCPacket::new_with_payload(payload); - pkt.fill_peer_manager_hdr(c_peer_id, s_peer_id, PacketType::Data as u8); - c_peer.send_msg(pkt).await.unwrap(); - - let got = timeout(Duration::from_secs(2), async move { - recv_packet_from_chan(&mut packet_recv).await - }) - .await - .unwrap() - .unwrap(); - - assert_eq!(got.payload(), payload); - assert_eq!( - got.peer_manager_header().unwrap().packet_type, - PacketType::Data as u8 - ); - } - - #[tokio::test] - async fn peer_conn_secure_mode_network_secret_confirmed() { - let (c, s) = create_ring_tunnel_pair(); - - let c_peer_id = new_peer_id(); - let s_peer_id = new_peer_id(); - - let c_ctx = get_mock_global_ctx(); - let s_ctx = get_mock_global_ctx(); - - c_ctx - .config - .set_network_identity(NetworkIdentity::new("net1".to_string(), "sec1".to_string())); - s_ctx - .config - .set_network_identity(NetworkIdentity::new("net1".to_string(), "sec1".to_string())); - - set_secure_mode_cfg(&c_ctx, true); - set_secure_mode_cfg(&s_ctx, true); - - let ps = Arc::new(PeerSessionStore::new()); - let mut c_peer = PeerConn::new(c_peer_id, c_ctx, Box::new(c), ps.clone()); - let mut s_peer = PeerConn::new(s_peer_id, s_ctx, Box::new(s), ps.clone()); - - let (c_ret, s_ret) = tokio::join!( - c_peer.do_handshake_as_client(), - s_peer.do_handshake_as_server() - ); - c_ret.unwrap(); - s_ret.unwrap(); - - assert_eq!( - c_peer.get_conn_info().secure_auth_level, - SecureAuthLevel::NetworkSecretConfirmed as i32, - ); - assert_eq!( - s_peer.get_conn_info().secure_auth_level, - SecureAuthLevel::NetworkSecretConfirmed as i32, - ); - assert_eq!( - c_peer.get_conn_info().peer_identity_type, - PeerIdentityType::Admin as i32, - ); - assert_eq!( - s_peer.get_conn_info().peer_identity_type, - PeerIdentityType::Admin as i32, - ); - } - - #[tokio::test] - async fn peer_conn_secure_mode_shared_node_pubkey_verified() { - let (c, s) = create_ring_tunnel_pair(); - - let c_peer_id = new_peer_id(); - let s_peer_id = new_peer_id(); - - let c_ctx = get_mock_global_ctx(); - let s_ctx = get_mock_global_ctx(); - - c_ctx - .config - .set_network_identity(NetworkIdentity::new("net1".to_string(), "sec2".to_string())); - s_ctx.config.set_network_identity(NetworkIdentity { - network_name: "net2".to_string(), - network_secret: None, - network_secret_digest: None, - }); - - let remote_url: url::Url = c.info().unwrap().remote_addr.unwrap().url.parse().unwrap(); - - set_secure_mode_cfg(&c_ctx, true); - set_secure_mode_cfg(&s_ctx, true); - - c_ctx.config.set_peers(vec![PeerConfig { - uri: remote_url, - peer_public_key: Some( - s_ctx - .config - .get_secure_mode() - .unwrap() - .local_public_key - .unwrap(), - ), - }]); - - let ps = Arc::new(PeerSessionStore::new()); - let mut c_peer = PeerConn::new(c_peer_id, c_ctx, Box::new(c), ps.clone()); - let mut s_peer = PeerConn::new(s_peer_id, s_ctx, Box::new(s), ps.clone()); - - let (c_ret, s_ret) = tokio::join!( - c_peer.do_handshake_as_client(), - s_peer.do_handshake_as_server() - ); - c_ret.unwrap(); - s_ret.unwrap(); - - assert_eq!( - c_peer.get_conn_info().secure_auth_level, - SecureAuthLevel::PeerVerified as i32, - ); - assert_eq!( - c_peer.get_conn_info().peer_identity_type, - PeerIdentityType::SharedNode as i32, - ); - assert_eq!( - s_peer.get_conn_info().peer_identity_type, - PeerIdentityType::Admin as i32, - ); - } - - #[tokio::test] - async fn peer_conn_secure_mode_shared_node_without_pin_is_unauthenticated() { - let (c, s) = create_ring_tunnel_pair(); - - let c_peer_id = new_peer_id(); - let s_peer_id = new_peer_id(); - - let c_ctx = get_mock_global_ctx(); - let s_ctx = get_mock_global_ctx(); - - c_ctx - .config - .set_network_identity(NetworkIdentity::new("net1".to_string(), "sec2".to_string())); - s_ctx.config.set_network_identity(NetworkIdentity { - network_name: "net2".to_string(), - network_secret: None, - network_secret_digest: None, - }); - - set_secure_mode_cfg(&c_ctx, true); - set_secure_mode_cfg(&s_ctx, true); - - let ps = Arc::new(PeerSessionStore::new()); - let mut c_peer = PeerConn::new(c_peer_id, c_ctx, Box::new(c), ps.clone()); - let mut s_peer = PeerConn::new(s_peer_id, s_ctx, Box::new(s), ps.clone()); - - let (c_ret, s_ret) = tokio::join!( - c_peer.do_handshake_as_client(), - s_peer.do_handshake_as_server() - ); - c_ret.unwrap(); - s_ret.unwrap(); - - assert_eq!( - c_peer.get_conn_info().secure_auth_level, - SecureAuthLevel::EncryptedUnauthenticated as i32, - ); - assert_eq!( - s_peer.get_conn_info().secure_auth_level, - SecureAuthLevel::EncryptedUnauthenticated as i32, - ); - assert_eq!( - c_peer.get_conn_info().peer_identity_type, - PeerIdentityType::SharedNode as i32, - ); - assert_eq!( - s_peer.get_conn_info().peer_identity_type, - PeerIdentityType::Admin as i32, - ); - } - - async fn peer_conn_pingpong_test_common( - drop_start: u32, - drop_end: u32, - conn_closed: bool, - drop_both: bool, - ) { - let (c, s) = create_ring_tunnel_pair(); - - // drop 1-3 packets should not affect pingpong - let c_recorder = Arc::new(DropSendTunnelFilter::new(drop_start, drop_end)); - let c = TunnelWithFilter::new(c, c_recorder.clone()); - - let c_peer_id = new_peer_id(); - let s_peer_id = new_peer_id(); - - let ps = Arc::new(PeerSessionStore::new()); - let mut c_peer = PeerConn::new(c_peer_id, get_mock_global_ctx(), Box::new(c), ps.clone()); - let mut s_peer = PeerConn::new(s_peer_id, get_mock_global_ctx(), Box::new(s), ps.clone()); - - let (c_ret, s_ret) = tokio::join!( - c_peer.do_handshake_as_client(), - s_peer.do_handshake_as_server() - ); - - s_peer.start_recv_loop(create_packet_recv_chan().0).await; - // do not start ping for s, s only reponde to ping from c - - assert!(c_ret.is_ok()); - assert!(s_ret.is_ok()); - - let close_notifier = c_peer.get_close_notifier(); - c_peer.start_pingpong(); - c_peer.start_recv_loop(create_packet_recv_chan().0).await; - - let throughput = c_peer.throughput.clone(); - let _t = AbortOnDropHandle::new(tokio::spawn(async move { - // if not drop both, we mock some rx traffic for client peer to test pinger - if drop_both { - return; - } - loop { - tokio::time::sleep(Duration::from_millis(100)).await; - throughput.record_rx_bytes(3); - } - })); - - tokio::time::sleep(Duration::from_secs(15)).await; - - if conn_closed { - assert!(close_notifier.is_closed()); - } else { - assert!(!close_notifier.is_closed()); - } - } - - #[tokio::test] - async fn peer_conn_pingpong_records_control_metrics() { - let (c, s) = create_ring_tunnel_pair(); - - let c_peer_id = new_peer_id(); - let s_peer_id = new_peer_id(); - - let c_ctx = get_mock_global_ctx(); - let s_ctx = get_mock_global_ctx(); - let ps = Arc::new(PeerSessionStore::new()); - let mut c_peer = PeerConn::new(c_peer_id, c_ctx.clone(), Box::new(c), ps.clone()); - let mut s_peer = PeerConn::new(s_peer_id, s_ctx.clone(), Box::new(s), ps.clone()); - - let (c_ret, s_ret) = tokio::join!( - c_peer.do_handshake_as_client(), - s_peer.do_handshake_as_server() - ); - - assert!(c_ret.is_ok()); - assert!(s_ret.is_ok()); - - s_peer.start_recv_loop(create_packet_recv_chan().0).await; - c_peer.start_pingpong(); - c_peer.start_recv_loop(create_packet_recv_chan().0).await; - - wait_for_condition( - || { - let c_ctx = c_ctx.clone(); - let s_ctx = s_ctx.clone(); - async move { - metric_value(&c_ctx, MetricName::TrafficControlBytesTx, "default") > 0 - && metric_value(&c_ctx, MetricName::TrafficControlBytesRx, "default") > 0 - && metric_value(&s_ctx, MetricName::TrafficControlBytesTx, "default") > 0 - && metric_value(&s_ctx, MetricName::TrafficControlBytesRx, "default") > 0 - } - }, - Duration::from_secs(5), - ) - .await; - } - - #[tokio::test] - async fn peer_conn_pingpong_timeout_not_close() { - peer_conn_pingpong_test_common(3, 5, false, false).await; - } - - #[tokio::test] - async fn peer_conn_pingpong_oneside_timeout() { - peer_conn_pingpong_test_common(4, 12, false, false).await; - } - - #[tokio::test] - async fn peer_conn_pingpong_bothside_timeout() { - peer_conn_pingpong_test_common(3, 14, true, true).await; - } - - #[tokio::test] - async fn close_tunnel_during_handshake() { - let ps = Arc::new(PeerSessionStore::new()); - let (c, s) = create_ring_tunnel_pair(); - let mut c_peer = PeerConn::new( - new_peer_id(), - get_mock_global_ctx(), - Box::new(c), - ps.clone(), - ); - let j = tokio::spawn(async move { - tokio::time::sleep(Duration::from_secs(1)).await; - drop(s); - }); - timeout(Duration::from_millis(1500), c_peer.do_handshake_as_client()) - .await - .unwrap() - .unwrap_err(); - let _ = tokio::join!(j); - } - - /// Helper: set up a credential node's GlobalCtx with a specific private key - /// (no network_secret, secure mode enabled with the given keypair) - fn set_credential_mode_cfg( - global_ctx: &GlobalCtx, - network_name: &str, - private_key: &x25519_dalek::StaticSecret, - ) { - use crate::common::config::NetworkIdentity; - let public = x25519_dalek::PublicKey::from(private_key); - global_ctx - .config - .set_network_identity(NetworkIdentity::new_credential(network_name.to_string())); - global_ctx.config.set_secure_mode(Some(SecureModeConfig { - enabled: true, - local_private_key: Some(BASE64_STANDARD.encode(private_key.as_bytes())), - local_public_key: Some(BASE64_STANDARD.encode(public.as_bytes())), - })); - } - - /// Test: credential node connects to admin node, admin has credential in trusted list. - /// Handshake should succeed with PeerVerified auth level on server side. - #[tokio::test] - async fn peer_conn_credential_node_connects_to_admin() { - let (c, s) = create_ring_tunnel_pair(); - - let c_peer_id = new_peer_id(); - let s_peer_id = new_peer_id(); - - // Admin node (server) has network_secret - let s_ctx = get_mock_global_ctx(); - s_ctx.config.set_network_identity(NetworkIdentity::new( - "net1".to_string(), - "secret".to_string(), - )); - set_secure_mode_cfg(&s_ctx, true); - - // Generate a credential on admin and get the private key for the client - let (cred_id, cred_secret) = s_ctx.get_credential_manager().generate_credential( - vec!["guest".to_string()], - false, - vec![], - std::time::Duration::from_secs(3600), - ); - - // Credential node (client) uses credential private key - let c_ctx = get_mock_global_ctx(); - let privkey_bytes: [u8; 32] = BASE64_STANDARD - .decode(&cred_secret) - .unwrap() - .try_into() - .unwrap(); - let private = x25519_dalek::StaticSecret::from(privkey_bytes); - set_credential_mode_cfg(&c_ctx, "net1", &private); - - let ps = Arc::new(PeerSessionStore::new()); - let mut c_peer = PeerConn::new(c_peer_id, c_ctx, Box::new(c), ps.clone()); - let mut s_peer = PeerConn::new(s_peer_id, s_ctx, Box::new(s), ps.clone()); - - let (c_ret, s_ret) = tokio::join!( - c_peer.do_handshake_as_client(), - s_peer.do_handshake_as_server() - ); - - c_ret.unwrap(); - s_ret.unwrap(); - - // Server should see credential node as PeerVerified - assert_eq!( - s_peer.get_conn_info().secure_auth_level, - SecureAuthLevel::PeerVerified as i32, - ); - assert_eq!( - s_peer.get_conn_info().peer_identity_type, - PeerIdentityType::Credential as i32, - ); - - // Client (credential node) keeps encrypted unauthenticated level - assert_eq!( - c_peer.get_conn_info().secure_auth_level, - SecureAuthLevel::EncryptedUnauthenticated as i32, - ); - assert_eq!( - c_peer.get_conn_info().peer_identity_type, - PeerIdentityType::Admin as i32, - ); - - // Verify credential ID matches - let _ = cred_id; // just to use it - } - - /// Test: unknown credential node (not in trusted list) is rejected by admin. - #[tokio::test] - async fn peer_conn_unknown_credential_rejected() { - let (c, s) = create_ring_tunnel_pair(); - - let c_peer_id = new_peer_id(); - let s_peer_id = new_peer_id(); - - // Admin node (server) with no credentials generated - let s_ctx = get_mock_global_ctx(); - s_ctx.config.set_network_identity(NetworkIdentity::new( - "net1".to_string(), - "secret".to_string(), - )); - set_secure_mode_cfg(&s_ctx, true); - - // Unknown credential node (client) with random key, not in admin's trusted list - let c_ctx = get_mock_global_ctx(); - let random_private = x25519_dalek::StaticSecret::random_from_rng(OsRng); - set_credential_mode_cfg(&c_ctx, "net1", &random_private); - - let ps = Arc::new(PeerSessionStore::new()); - let mut c_peer = PeerConn::new(c_peer_id, c_ctx, Box::new(c), ps.clone()); - let mut s_peer = PeerConn::new(s_peer_id, s_ctx, Box::new(s), ps.clone()); - - let (c_ret, s_ret) = tokio::join!( - c_peer.do_handshake_as_client(), - s_peer.do_handshake_as_server() - ); - - // Server should reject the unknown credential - assert!(s_ret.is_err(), "server should reject unknown credential"); - // Client may also fail due to connection being closed - let _ = c_ret; - } - - /// Test: two admin nodes with same network_secret still get NetworkSecretConfirmed. - /// (Regression test: credential system should not break normal admin-to-admin auth) - #[tokio::test] - async fn peer_conn_admin_to_admin_still_works() { - let (c, s) = create_ring_tunnel_pair(); - - let c_peer_id = new_peer_id(); - let s_peer_id = new_peer_id(); - - let c_ctx = get_mock_global_ctx(); - let s_ctx = get_mock_global_ctx(); - - c_ctx.config.set_network_identity(NetworkIdentity::new( - "net1".to_string(), - "secret".to_string(), - )); - s_ctx.config.set_network_identity(NetworkIdentity::new( - "net1".to_string(), - "secret".to_string(), - )); - - set_secure_mode_cfg(&c_ctx, true); - set_secure_mode_cfg(&s_ctx, true); - - let ps = Arc::new(PeerSessionStore::new()); - let mut c_peer = PeerConn::new(c_peer_id, c_ctx, Box::new(c), ps.clone()); - let mut s_peer = PeerConn::new(s_peer_id, s_ctx, Box::new(s), ps.clone()); - - let (c_ret, s_ret) = tokio::join!( - c_peer.do_handshake_as_client(), - s_peer.do_handshake_as_server() - ); - - c_ret.unwrap(); - s_ret.unwrap(); - - assert_eq!( - c_peer.get_conn_info().secure_auth_level, - SecureAuthLevel::NetworkSecretConfirmed as i32, - ); - assert_eq!( - s_peer.get_conn_info().secure_auth_level, - SecureAuthLevel::NetworkSecretConfirmed as i32, - ); - } - - /// Test: revoked credential is rejected on new connection attempt. - #[tokio::test] - async fn peer_conn_revoked_credential_rejected() { - // Admin generates credential, then revokes it - let admin_ctx = get_mock_global_ctx(); - admin_ctx.config.set_network_identity(NetworkIdentity::new( - "net1".to_string(), - "secret".to_string(), - )); - set_secure_mode_cfg(&admin_ctx, true); - - let (cred_id, cred_secret) = admin_ctx.get_credential_manager().generate_credential( - vec![], - false, - vec![], - std::time::Duration::from_secs(3600), - ); - - // Revoke the credential - assert!( - admin_ctx - .get_credential_manager() - .revoke_credential(&cred_id) - ); - - // Now try to connect with the revoked credential - let (c, s) = create_ring_tunnel_pair(); - let c_peer_id = new_peer_id(); - let s_peer_id = new_peer_id(); - - let c_ctx = get_mock_global_ctx(); - let privkey_bytes: [u8; 32] = BASE64_STANDARD - .decode(&cred_secret) - .unwrap() - .try_into() - .unwrap(); - let private = x25519_dalek::StaticSecret::from(privkey_bytes); - set_credential_mode_cfg(&c_ctx, "net1", &private); - - let ps = Arc::new(PeerSessionStore::new()); - let mut c_peer = PeerConn::new(c_peer_id, c_ctx, Box::new(c), ps.clone()); - let mut s_peer = PeerConn::new(s_peer_id, admin_ctx, Box::new(s), ps.clone()); - - let (c_ret, s_ret) = tokio::join!( - c_peer.do_handshake_as_client(), - s_peer.do_handshake_as_server() - ); - - // Server should reject the revoked credential - assert!(s_ret.is_err(), "server should reject revoked credential"); - let _ = c_ret; - } -} diff --git a/easytier/src/peers/peer_conn_ping.rs b/easytier-core/src/peers/conn/peer_conn_ping.rs similarity index 92% rename from easytier/src/peers/peer_conn_ping.rs rename to easytier-core/src/peers/conn/peer_conn_ping.rs index c72401de..05b1f533 100644 --- a/easytier/src/peers/peer_conn_ping.rs +++ b/easytier-core/src/peers/conn/peer_conn_ping.rs @@ -3,25 +3,21 @@ use std::{ Arc, atomic::{AtomicU32, Ordering}, }, - time::Duration, + time::{Duration, Instant}, }; -use quanta::Instant; use rand::{Rng, thread_rng}; -use tokio::{ - sync::broadcast, - task::JoinSet, - time::{Interval, timeout}, -}; +use tokio::{sync::broadcast, task::JoinSet}; use tracing::Instrument; use crate::{ - common::{PeerId, error::Error}, - peers::traffic_metrics::AggregateTrafficMetrics, + config::PeerId, + foundation::time::{Interval, interval, timeout}, + packet::{PacketType, ZCPacket}, + peers::{context::ArcPeerContext, error::Error}, tunnel::{ TunnelError, mpsc::MpscTunnelSender, - packet_def::{PacketType, ZCPacket}, stats::{Throughput, WindowLatency}, }, }; @@ -62,7 +58,7 @@ impl PingIntervalController { Self { throughput, loss_counter, - interval: tokio::time::interval(Duration::from_secs(1)), + interval: interval(Duration::from_secs(1)), logic_time: 0, last_send_logic_time: 0, @@ -120,7 +116,8 @@ pub struct PeerConnPinger { latency_stats: Arc, loss_rate_stats: Arc, throughput_stats: Arc, - control_metrics: AggregateTrafficMetrics, + context: ArcPeerContext, + network_name: String, tasks: JoinSet>, } @@ -143,7 +140,8 @@ impl PeerConnPinger { latency_stats: Arc, loss_rate_stats: Arc, throughput_stats: Arc, - control_metrics: AggregateTrafficMetrics, + context: ArcPeerContext, + network_name: String, ) -> Self { Self { my_peer_id, @@ -154,7 +152,8 @@ impl PeerConnPinger { ctrl_sender, loss_rate_stats, throughput_stats, - control_metrics, + context, + network_name, } } @@ -168,7 +167,8 @@ impl PeerConnPinger { my_node_id: PeerId, peer_id: PeerId, sink: &MpscTunnelSender, - control_metrics: &AggregateTrafficMetrics, + context: &ArcPeerContext, + network_name: &str, receiver: &mut broadcast::Receiver, seq: u32, ) -> Result { @@ -176,7 +176,7 @@ impl PeerConnPinger { let req = Self::new_ping_packet(my_node_id, peer_id, seq); let req_len = req.buf_len() as u64; sink.send(req).await?; - control_metrics.record_tx(req_len); + context.record_control_tx(network_name, req_len); let now = Instant::now(); // wait until we get a pong packet in ctrl_resp_receiver @@ -223,7 +223,8 @@ impl PeerConnPinger { pub async fn pingpong(&mut self) { let sink = self.sink.clone(); - let control_metrics = self.control_metrics.clone(); + let context = self.context.clone(); + let network_name = self.network_name.clone(); let my_node_id = self.my_peer_id; let peer_id = self.peer_id; let latency_stats = self.latency_stats.clone(); @@ -269,7 +270,8 @@ impl PeerConnPinger { ); let sink = sink.clone(); - let control_metrics = control_metrics.clone(); + let context = context.clone(); + let network_name = network_name.clone(); let receiver = ctrl_resp_sender.subscribe(); let ping_res_sender = ping_res_sender.clone(); pingpong_tasks.spawn(async move { @@ -278,7 +280,8 @@ impl PeerConnPinger { my_node_id, peer_id, &sink, - &control_metrics, + &context, + &network_name, &mut receiver, req_seq, ) diff --git a/easytier/src/peers/peer_map.rs b/easytier-core/src/peers/conn/peer_map.rs similarity index 79% rename from easytier/src/peers/peer_map.rs rename to easytier-core/src/peers/conn/peer_map.rs index 46e55aaf..4a89496c 100644 --- a/easytier/src/peers/peer_map.rs +++ b/easytier-core/src/peers/conn/peer_map.rs @@ -1,4 +1,5 @@ use std::{ + collections::{BTreeSet, HashMap, HashSet}, net::{Ipv4Addr, Ipv6Addr}, sync::Arc, }; @@ -9,52 +10,60 @@ use parking_lot::Mutex; use tokio::sync::RwLock; use crate::{ - common::{ - PeerId, + config::PeerId, + packet::ZCPacket, + peers::{ + context::{ArcPeerContext, NetworkIdentity, PeerEvent}, error::Error, - global_ctx::{ArcGlobalCtx, GlobalCtxEvent, NetworkIdentity}, - shrink_dashmap, + util::shrink_dashmap, }, proto::{ - api::instance::{self, PeerConnInfo}, - peer_rpc::{PeerIdentityType, RoutePeerInfo}, + core_peer::peer::{PeerConnInfo, Route as CoreRoute}, + peer_rpc::{ + DirectConnectedPeerInfo, PeerIdentityType, PeerInfoForGlobalMap, RoutePeerInfo, + }, }, - tunnel::{TunnelError, packet_def::ZCPacket}, + tunnel::TunnelError, }; use super::{ - PacketRecvChan, peer::Peer, peer_conn::{PeerConn, PeerConnId}, - route_trait::{ArcRoute, NextHopPolicy}, +}; +use crate::peers::{ + PacketRecvChan, + route::{ArcRoute, NextHopPolicy}, }; pub struct PeerMap { - global_ctx: ArcGlobalCtx, + context: ArcPeerContext, my_peer_id: PeerId, peer_map: DashMap>, packet_send: PacketRecvChan, routes: RwLock>, - alive_client_urls: Arc>>, + alive_client_urls: Arc>>>, } impl PeerMap { - pub fn new(packet_send: PacketRecvChan, global_ctx: ArcGlobalCtx, my_peer_id: PeerId) -> Self { + pub(crate) fn new( + packet_send: PacketRecvChan, + context: ArcPeerContext, + my_peer_id: PeerId, + ) -> Self { PeerMap { - global_ctx, + context, my_peer_id, peer_map: DashMap::new(), packet_send, routes: RwLock::new(Vec::new()), - alive_client_urls: Arc::new(Mutex::new(multimap::MultiMap::new())), + alive_client_urls: Arc::new(Mutex::new(HashMap::new())), } } async fn add_new_peer(&self, peer: Peer) { let peer_id = peer.peer_node_id; self.peer_map.insert(peer_id, Arc::new(peer)); - self.global_ctx - .issue_event(GlobalCtxEvent::PeerAdded(peer_id)); + self.context.issue_event(PeerEvent::PeerAdded(peer_id)); } pub async fn add_new_peer_conn(&self, peer_conn: PeerConn) -> Result<(), Error> { @@ -62,7 +71,7 @@ impl PeerMap { let peer_id = peer_conn.get_peer_id(); let no_entry = self.peer_map.get(&peer_id).is_none(); if no_entry { - let new_peer = Peer::new(peer_id, self.packet_send.clone(), self.global_ctx.clone()); + let new_peer = Peer::new(peer_id, self.packet_send.clone(), self.context.clone()); new_peer.add_peer_conn(peer_conn).await?; self.add_new_peer(new_peer).await; } else { @@ -84,7 +93,9 @@ impl PeerMap { let alive_client_url: url::Url = conn_info.tunnel?.remote_addr?.into(); self.alive_client_urls .lock() - .insert(alive_client_url.clone(), conn_id); + .entry(alive_client_url.clone()) + .or_default() + .insert(conn_id); tokio::spawn(async move { if let Some(mut waiter) = close_notifier.get_waiter().await { @@ -94,12 +105,12 @@ impl PeerMap { return; }; let mut guard = alive_conns.lock(); - if let Some(mut conn_ids) = guard.remove(&alive_client_url) { + if let Some(conn_ids) = guard.get_mut(&alive_client_url) { conn_ids.retain(|id| id != &conn_id); - if !conn_ids.is_empty() { - guard.insert_many(alive_client_url, conn_ids); + if conn_ids.is_empty() { + guard.remove(&alive_client_url); } - }; + } let alive_conn_count = guard.len(); drop(guard); tracing::debug!( @@ -326,8 +337,7 @@ impl PeerMap { let remove_ret = self.peer_map.remove(&peer_id); shrink_dashmap(&self.peer_map, None); - self.global_ctx - .issue_event(GlobalCtxEvent::PeerRemoved(peer_id)); + self.context.issue_event(PeerEvent::PeerRemoved(peer_id)); tracing::info!( ?peer_id, has_old_value = ?remove_ret.is_some(), @@ -342,6 +352,20 @@ impl PeerMap { routes.insert(0, route); } + pub(crate) async fn clear_resources(&self) { + for peer_id in self.list_peers() { + let _ = self.close_peer(peer_id).await; + } + let routes = { + let mut routes = self.routes.write().await; + std::mem::take(&mut *routes) + }; + for route in routes { + route.close().await; + } + self.alive_client_urls.lock().clear(); + } + pub async fn clean_peer_without_conn(&self) { let mut to_remove = vec![]; @@ -367,7 +391,7 @@ impl PeerMap { route_map } - pub async fn list_route_infos(&self) -> Vec { + pub async fn list_route_infos(&self) -> Vec { if let Some(route) = self.routes.read().await.iter().next() { return route.list_routes().await; } @@ -390,18 +414,50 @@ impl PeerMap { pub fn my_peer_id(&self) -> PeerId { self.my_peer_id } - - pub fn get_global_ctx(&self) -> ArcGlobalCtx { - self.global_ctx.clone() - } } impl Drop for PeerMap { fn drop(&mut self) { tracing::debug!( self.my_peer_id, - network = ?self.global_ctx.get_network_identity(), + network = ?self.context.network_identity(), "PeerMap is dropped" ); } } + +/// Aggregates the directly-connected peers of several peer maps into the +/// peer-center reporting format, keeping the lowest observed latency per peer. +pub(crate) async fn direct_peer_info(peer_maps: &[Arc]) -> PeerInfoForGlobalMap { + let mut peers = BTreeSet::new(); + for peer_map in peer_maps { + peers.extend(peer_map.list_peers()); + } + + let mut ret = PeerInfoForGlobalMap::default(); + for peer in peers { + let mut conns = None; + for peer_map in peer_maps { + if let Some(found) = peer_map.list_peer_conns(peer).await { + conns = Some(found); + break; + } + } + let Some(min_lat) = conns + .into_iter() + .flatten() + .map(|conn| conn.stats.as_ref().unwrap().latency_us) + .min() + else { + continue; + }; + + ret.direct_peers.insert( + peer, + DirectConnectedPeerInfo { + latency_ms: std::cmp::max(1, (min_lat as u32 / 1000) as i32), + }, + ); + } + ret +} diff --git a/easytier/src/peers/peer_session.rs b/easytier-core/src/peers/conn/peer_session.rs similarity index 91% rename from easytier/src/peers/peer_session.rs rename to easytier-core/src/peers/conn/peer_session.rs index 0e3e7f38..93dba786 100644 --- a/easytier/src/peers/peer_session.rs +++ b/easytier-core/src/peers/conn/peer_session.rs @@ -2,18 +2,15 @@ use std::sync::{ Arc, RwLock, atomic::{AtomicBool, Ordering}, }; -use std::time::Duration; +use std::time::{Duration, Instant}; use anyhow::anyhow; use crossbeam::atomic::AtomicCell; use dashmap::DashMap; -use quanta::Instant; -use super::secure_datagram::{SecureDatagramDirection, SecureDatagramSession}; -use crate::{ - common::{PeerId, shrink_dashmap}, - tunnel::packet_def::ZCPacket, -}; +use crate::peers::util::shrink_dashmap; +use crate::tunnel::secure_datagram::{SecureDatagramDirection, SecureDatagramSession}; +use crate::{config::PeerId, packet::ZCPacket}; const SESSION_IDLE_TIMEOUT: Duration = Duration::from_secs(60); @@ -101,10 +98,6 @@ impl PeerSessionStore { self.sessions.remove(key); } - pub fn insert_session(&self, key: SessionKey, session: Arc) { - self.sessions.insert(key, PeerSessionEntry::new(session)); - } - pub fn evict_unused_sessions(&self) { self.evict_unused_sessions_idle(SESSION_IDLE_TIMEOUT); } @@ -269,8 +262,6 @@ impl std::fmt::Debug for PeerSession { } impl PeerSession { - const SYNC_RX_GRACE_AFTER_MS: u64 = SecureDatagramSession::SYNC_RX_GRACE_AFTER_MS; - pub fn new( peer_id: PeerId, root_key: [u8; 32], @@ -405,11 +396,32 @@ impl PeerSession { } } +#[cfg(any(test, feature = "test-utils"))] +mod test_utils { + use super::*; + + impl PeerSessionStore { + #[doc(hidden)] + pub(crate) fn contains_valid(&self, key: &SessionKey) -> bool { + self.sessions + .get(key) + .is_some_and(|entry| entry.session.is_valid()) + } + } +} + #[cfg(test)] mod tests { use super::*; + impl PeerSessionStore { + fn insert_session(&self, key: SessionKey, session: Arc) { + self.sessions.insert(key, PeerSessionEntry::new(session)); + } + } + #[test] + #[cfg(all(feature = "aes-gcm", feature = "chacha20"))] fn peer_session_supports_asymmetric_algorithms() { let a: PeerId = 10; let b: PeerId = 20; @@ -451,14 +463,6 @@ mod tests { assert_eq!(pkt2.payload(), plaintext2); } - #[test] - fn sync_root_key_preserves_generic_grace_window_constant() { - assert_eq!( - PeerSession::SYNC_RX_GRACE_AFTER_MS, - SecureDatagramSession::SYNC_RX_GRACE_AFTER_MS - ); - } - #[test] fn peer_session_store_keeps_recent_session_without_external_refs() { let store = PeerSessionStore::new(); @@ -531,4 +535,29 @@ mod tests { "invalid sessions should not be kept by recent activity" ); } + + #[cfg(feature = "test-utils")] + #[test] + fn contains_valid_does_not_refresh_session_activity() { + let store = PeerSessionStore::new(); + let key = SessionKey::new("net".to_string(), 20); + let session = Arc::new(PeerSession::new( + 20, + PeerSession::new_root_key(), + 1, + 0, + "aes-gcm".to_string(), + "aes-gcm".to_string(), + None, + )); + store.insert_session(key.clone(), session); + let last_used_at = store.sessions.get(&key).unwrap().last_used_at.load(); + + assert!(store.contains_valid(&key)); + + assert_eq!( + store.sessions.get(&key).unwrap().last_used_at.load(), + last_used_at + ); + } } diff --git a/easytier-core/src/peers/context.rs b/easytier-core/src/peers/context.rs new file mode 100644 index 00000000..db891259 --- /dev/null +++ b/easytier-core/src/peers/context.rs @@ -0,0 +1,1775 @@ +use std::{ + collections::HashMap, + net::IpAddr, + sync::{ + Arc, Mutex, + atomic::{AtomicBool, Ordering}, + }, + time::{SystemTime, UNIX_EPOCH}, +}; + +use arc_swap::ArcSwap; +use cidr::{Ipv4Cidr, Ipv4Inet, Ipv6Cidr, Ipv6Inet}; +use dashmap::DashMap; +use easytier_proto::{ + acl::Acl, + common::{ + FlagsInConfig, LimiterConfig, PeerFeatureFlag, SecureModeConfig, StunInfo, TunnelInfo, + }, + peer_rpc::{PeerGroupInfo, TrustedCredentialPubkeyProof}, +}; +use hmac::{Hmac, Mac}; +use sha2::Sha256; + +pub use crate::config::{NetworkIdentity, NetworkSecretDigest}; +use crate::{ + config::peers::{HostRoutingPolicy, PeerGroupIdentity, PeerRuntimeConfig, PeerRuntimeSnapshot}, + config::runtime::CoreRuntimeConfigStore, + config::{ + CoreConfig, IpPrefix, NodeConfig, PeerId, PeerPolicyConfig, ProxyNetworkConfig, + RouteConfig, TrafficConfig, + }, + events::{CoreEvent, CoreEventSink}, + foundation::stats::{LabelSet, LabelType, MetricName, StatsManager}, + foundation::token_bucket::{ArcByteLimiter, TokenBucketManager}, + peers::{ + credential_manager::{CredentialManager, CredentialStorage}, + util::shrink_dashmap, + whitelist::check_network_in_relay_whitelist, + }, +}; + +pub(crate) const SECRET_PROOF_PREFIX: &[u8] = b"easytier secret proof"; +const PEER_EVENT_CAPACITY: usize = 100; + +#[derive(Debug, Clone)] +#[allow(clippy::enum_variant_names)] +pub(crate) enum PeerEvent { + PeerAdded(PeerId), + PeerRemoved(PeerId), + PeerConnAdded(easytier_proto::core_peer::peer::PeerConnInfo), + PeerConnRemoved(easytier_proto::core_peer::peer::PeerConnInfo), +} + +#[derive(Debug, Clone, PartialEq, Eq)] +#[allow(clippy::enum_variant_names)] +pub(crate) enum PeerContextEvent { + PeerAdded(PeerId), + PeerRemoved(PeerId), + PeerConnAdded, + PeerConnRemoved, +} + +pub(crate) type PeerContextEventSubscriber = tokio::sync::broadcast::Receiver; + +/// Normalized product and host inputs used to derive one peer runtime version. +/// +/// The host owns platform-specific normalization of node, route, identity, and +/// capability values. Peer policy remains derived in core from the submitted +/// flags and ACL. +#[derive(Debug, Clone)] +pub struct PeerRuntimeSnapshotInput { + pub node: NodeConfig, + pub routes: RouteConfig, + pub network_identity: NetworkIdentity, + pub stun_info: StunInfo, + pub flags: FlagsInConfig, + pub secure_mode: Option, + pub host_routing: HostRoutingPolicy, + pub acl: Option, + pub easytier_version: String, + pub vpn_portal_cidr: Option, + pub pinned_peers: Vec<(url::Url, Option)>, + pub ospf_update_my_foreign_network_interval_sec: u64, + pub max_direct_conns_per_peer_in_foreign_network: usize, + pub hmac_secret_digest: bool, +} + +#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)] +struct PeerTrafficLimits { + instance_recv_bps: Option, + foreign_relay_bps: Option, +} + +impl PeerTrafficLimits { + fn from_portable(runtime: &PeerRuntimeConfig, flags: &FlagsInConfig) -> Self { + let traffic = &runtime.core.traffic; + Self { + instance_recv_bps: Self::normalize( + traffic.instance_recv_bps_limit, + flags.instance_recv_bps_limit, + ), + foreign_relay_bps: Self::normalize( + traffic.foreign_relay_bps_limit, + flags.foreign_relay_bps_limit, + ), + } + } + + fn normalize(explicit: Option, legacy: u64) -> Option { + match explicit { + Some(u64::MAX) => None, + Some(limit) => Some(limit), + None if !matches!(legacy, 0 | u64::MAX) => Some(legacy), + None => None, + } + } +} + +impl PeerRuntimeSnapshot { + pub fn from_host_input(input: PeerRuntimeSnapshotInput) -> Self { + let PeerRuntimeSnapshotInput { + node, + routes, + network_identity, + stun_info, + flags, + secure_mode, + host_routing, + acl, + easytier_version, + vpn_portal_cidr, + pinned_peers, + ospf_update_my_foreign_network_interval_sec, + max_direct_conns_per_peer_in_foreign_network, + hmac_secret_digest, + } = input; + let feature_flags = PeerFeatureFlag { + kcp_input: !flags.disable_kcp_input, + no_relay_kcp: flags.disable_relay_kcp, + support_conn_list_sync: true, + quic_input: !flags.disable_quic_input, + no_relay_quic: flags.disable_relay_quic, + need_p2p: flags.need_p2p, + disable_p2p: flags.disable_p2p, + avoid_relay_data: flags.disable_relay_data, + ..Default::default() + }; + let peer_policy = PeerPolicyConfig { + p2p_enabled: !flags.disable_p2p, + relay_peer_rpc: flags.relay_all_peer_rpc, + relay_data: !flags.disable_relay_data, + latency_first: flags.latency_first, + encryption_required: flags.enable_encryption, + }; + let traffic = TrafficConfig { + mtu: u16::try_from(flags.mtu) + .ok() + .filter(|configured| *configured != 0), + instance_recv_bps_limit: (flags.instance_recv_bps_limit != u64::MAX) + .then_some(flags.instance_recv_bps_limit), + foreign_relay_bps_limit: (flags.foreign_relay_bps_limit != u64::MAX) + .then_some(flags.foreign_relay_bps_limit), + }; + let avoid_relay_data_preference = check_network_in_relay_whitelist( + &flags.relay_network_whitelist, + &network_identity.network_name, + ) + .is_err(); + let (acl_group_declarations, peer_group_memberships) = peer_acl_groups(acl.as_ref()); + + Self { + runtime: PeerRuntimeConfig { + core: CoreConfig { + node, + routes, + peer_policy, + traffic, + }, + network_identity, + stun_info, + feature_flags, + secure_mode, + host_routing, + }, + easytier_version, + avoid_relay_data_preference, + flags, + vpn_portal_cidr, + pinned_peers, + peer_group_memberships, + acl_group_declarations, + ospf_update_my_foreign_network_interval_sec, + max_direct_conns_per_peer_in_foreign_network, + hmac_secret_digest, + } + } + + fn traffic_limits(&self) -> PeerTrafficLimits { + PeerTrafficLimits::from_portable(&self.runtime, &self.flags) + } +} + +fn peer_acl_groups(acl: Option<&Acl>) -> (Vec, Vec) { + let group = acl + .and_then(|acl| acl.acl_v1.as_ref()) + .and_then(|acl| acl.group.as_ref()); + let declarations = group.map_or_else(Vec::new, |group| { + group + .declares + .iter() + .map(|identity| PeerGroupIdentity { + group_name: identity.group_name.clone(), + group_secret: identity.group_secret.clone(), + }) + .collect() + }); + let memberships = group.map_or_else(Vec::new, |group| { + group + .declares + .iter() + .filter(|identity| group.members.contains(&identity.group_name)) + .map(|identity| PeerGroupIdentity { + group_name: identity.group_name.clone(), + group_secret: identity.group_secret.clone(), + }) + .collect() + }); + (declarations, memberships) +} + +/// Supplies the instance's current STUN observation. +pub(crate) trait PeerStunInfoSource: Send + Sync { + fn stun_info(&self) -> StunInfo { + StunInfo::default() + } +} + +impl PeerStunInfoSource for () {} + +/// Supplies public-IPv6 state observed or leased by the host. +pub(crate) trait PeerPublicIpv6State: Send + Sync { + fn public_ipv6_lease_contains(&self, _ip: &std::net::Ipv6Addr) -> bool { + false + } + + fn public_ipv6_provider_enabled(&self) -> bool { + false + } + + fn advertised_ipv6_public_addr_prefix(&self) -> Option { + None + } +} + +impl PeerPublicIpv6State for () {} + +/// Host adapters used to assemble the core-owned peer context. Each field stays +/// narrow so peer modules cannot reach unrelated host state after construction. +pub(crate) struct CorePeerContextAdapters { + pub stun_info_source: Option>, + pub events: Arc, + pub credential_storage: Option>, +} + +/// Peer context backed by one core-owned submitted snapshot and its instance +/// runtime resources. +pub(crate) struct CorePeerContext { + config: CoreRuntimeConfigStore, + avoid_relay_data_preference: AtomicBool, + stun_info_source: Option>, + public_ipv6_state: Arc, + limiter_state: Mutex, + stats_manager: Arc, + credentials: Arc, + trusted_keys: Arc, + peer_events: tokio::sync::broadcast::Sender, + events: Arc, +} + +impl CorePeerContext { + pub(crate) fn new( + config: CoreRuntimeConfigStore, + public_ipv6_state: Arc, + adapters: CorePeerContextAdapters, + ) -> Self { + Self::new_with_stats_manager( + config, + public_ipv6_state, + adapters, + Arc::new(StatsManager::new()), + ) + } + + /// Builds a foreign-network context that contributes to the same + /// instance-level metrics registry while retaining independent identity, + /// credential, trusted-key, event, and limiter state. + pub fn new_foreign( + config: CoreRuntimeConfigStore, + adapters: CorePeerContextAdapters, + parent: &CorePeerContext, + ) -> Self { + Self::new_with_stats_manager(config, Arc::new(()), adapters, parent.stats_manager()) + } + + fn new_with_stats_manager( + config: CoreRuntimeConfigStore, + public_ipv6_state: Arc, + adapters: CorePeerContextAdapters, + stats_manager: Arc, + ) -> Self { + let avoid_relay_data_preference = + AtomicBool::new(config.snapshot().peer.avoid_relay_data_preference); + let credentials = Arc::new( + adapters + .credential_storage + .map_or_else(CredentialManager::new, CredentialManager::from_storage), + ); + Self { + config, + avoid_relay_data_preference, + stun_info_source: adapters.stun_info_source, + public_ipv6_state, + limiter_state: Mutex::new(CoreLimiterState::default()), + stats_manager, + credentials, + trusted_keys: Arc::new(TrustedKeyMapManager::new()), + peer_events: tokio::sync::broadcast::channel(PEER_EVENT_CAPACITY).0, + events: adapters.events, + } + } + + fn snapshot(&self) -> Arc { + self.config.snapshot().peer.clone() + } + + pub fn stats_manager(&self) -> Arc { + self.stats_manager.clone() + } + + pub fn credential_manager(&self) -> Arc { + self.credentials.clone() + } + + fn record_control_metric( + &self, + network_name: &str, + bytes: u64, + bytes_metric: MetricName, + packets_metric: MetricName, + ) { + let labels = + LabelSet::new().with_label_type(LabelType::NetworkName(network_name.to_owned())); + self.stats_manager + .get_counter(bytes_metric, labels.clone()) + .add(bytes); + self.stats_manager.get_counter(packets_metric, labels).inc(); + } + + fn get_or_create_limiter(&self, key: &str, bps: u64) -> Option { + let mut state = self.limiter_state.lock().unwrap(); + if state.stopped { + return None; + } + let manager = state.manager.get_or_insert_with(TokenBucketManager::new); + Some( + manager.get_or_create( + key, + LimiterConfig { + burst_rate: None, + bps: Some(bps), + fill_duration_ms: None, + } + .into(), + ), + ) + } + + pub(crate) async fn stop(&self) { + let manager = { + let mut state = self.limiter_state.lock().unwrap(); + state.stopped = true; + state.manager.take() + }; + if let Some(manager) = manager { + manager.stop().await; + } + } +} + +#[derive(Default)] +struct CoreLimiterState { + manager: Option, + stopped: bool, +} + +fn config_ipv4(value: &IpPrefix) -> Option { + let IpAddr::V4(address) = value.address else { + return None; + }; + Ipv4Inet::new(address, value.prefix_len).ok() +} + +fn config_ipv4_cidr(value: &IpPrefix) -> Option { + let IpAddr::V4(address) = value.address else { + return None; + }; + Ipv4Cidr::new(address, value.prefix_len).ok() +} + +fn config_ipv6(value: &IpPrefix) -> Option { + let IpAddr::V6(address) = value.address else { + return None; + }; + Ipv6Inet::new(address, value.prefix_len).ok() +} + +/// Source of a trusted public key propagated by the OSPF route layer. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum TrustedKeySource { + OspfNode, + OspfCredential, +} + +#[derive(Debug, Clone)] +pub(crate) struct TrustedKeyMetadata { + pub source: TrustedKeySource, + pub expiry_unix: Option, +} + +impl TrustedKeyMetadata { + pub fn is_expired(&self) -> bool { + if let Some(expiry) = self.expiry_unix { + let now = SystemTime::now() + .duration_since(UNIX_EPOCH) + .unwrap() + .as_secs() as i64; + return now >= expiry; + } + false + } +} + +pub(crate) type TrustedKeyMap = HashMap, TrustedKeyMetadata>; + +pub(crate) struct TrustedKeyMapManager { + network_trusted_keys: DashMap>, +} + +impl TrustedKeyMapManager { + pub fn new() -> Self { + Self { + network_trusted_keys: DashMap::new(), + } + } + + pub fn update_trusted_keys(&self, network_name: &str, trusted_keys: TrustedKeyMap) { + match self.network_trusted_keys.entry(network_name.to_string()) { + dashmap::Entry::Vacant(entry) => { + entry.insert(ArcSwap::new(Arc::new(trusted_keys))); + } + dashmap::Entry::Occupied(entry) => { + entry.get().store(Arc::new(trusted_keys)); + } + } + } + + pub fn remove_trusted_keys(&self, network_name: &str) { + self.network_trusted_keys.remove(network_name); + shrink_dashmap(&self.network_trusted_keys, None); + } + + pub fn verify_trusted_key(&self, pubkey: &[u8], network_name: &str) -> bool { + self.verify_trusted_key_with_source(pubkey, network_name, None) + } + + pub fn verify_trusted_key_with_source( + &self, + pubkey: &[u8], + network_name: &str, + source: Option, + ) -> bool { + let Some(trusted_keys) = self + .network_trusted_keys + .get(network_name) + .map(|v| v.load_full()) + else { + return false; + }; + + let Some(metadata) = trusted_keys.get(&pubkey.to_vec()) else { + return false; + }; + + if let Some(source) = source { + metadata.source == source && !metadata.is_expired() + } else { + !metadata.is_expired() + } + } + + pub fn list_trusted_keys(&self, network_name: &str) -> Vec<(Vec, TrustedKeyMetadata)> { + let Some(trusted_keys) = self + .network_trusted_keys + .get(network_name) + .map(|v| v.load_full()) + else { + return Vec::new(); + }; + + let mut items = trusted_keys + .iter() + .filter(|(_, metadata)| !metadata.is_expired()) + .map(|(pubkey, metadata)| (pubkey.clone(), metadata.clone())) + .collect::>(); + items.sort_by(|left, right| left.0.cmp(&right.0)); + items + } +} + +impl Default for TrustedKeyMapManager { + fn default() -> Self { + Self::new() + } +} + +/// Runtime dependency interface for the peers module. +/// +/// `PeerContext` is intentionally scoped to `easytier-core::peers`; other core +/// modules should depend on their own narrow DTOs or traits instead of treating +/// this as a core-wide global context. +pub(crate) trait PeerContext: Send + Sync { + fn host_routing_policy(&self) -> HostRoutingPolicy { + HostRoutingPolicy::default() + } + + fn network_identity(&self) -> NetworkIdentity; + + fn network_name(&self) -> String { + self.network_identity().network_name + } + + fn flags(&self) -> FlagsInConfig { + FlagsInConfig::default() + } + + fn disable_relay_data(&self) -> bool { + self.flags().disable_relay_data + } + + fn secure_mode(&self) -> Option { + None + } + + fn stun_info(&self) -> StunInfo { + StunInfo::default() + } + + fn instance_id(&self) -> uuid::Uuid { + uuid::Uuid::nil() + } + + fn ipv4(&self) -> Option { + None + } + + fn ipv6(&self) -> Option { + None + } + + fn is_ip_local_ipv6(&self, ip: &std::net::Ipv6Addr) -> bool { + self.ipv6() + .map(|addr| addr.address() == *ip) + .unwrap_or(false) + } + + fn is_ip_local_virtual_ip(&self, ip: &IpAddr) -> bool { + match ip { + IpAddr::V4(v4) => self + .ipv4() + .map(|addr| addr.address() == *v4) + .unwrap_or(false), + IpAddr::V6(v6) => self.is_ip_local_ipv6(v6), + } + } + + fn p2p_only(&self) -> bool { + self.flags().p2p_only + } + + fn latency_first(&self) -> bool { + let flags = self.flags(); + flags.latency_first && !flags.p2p_only + } + + fn proxy_cidrs(&self) -> Vec { + Vec::new() + } + + fn proxy_networks(&self) -> Vec { + Vec::new() + } + + fn vpn_portal_cidr(&self) -> Option { + None + } + + fn hostname(&self) -> String { + String::new() + } + + fn feature_flags(&self) -> PeerFeatureFlag { + PeerFeatureFlag::default() + } + + fn set_avoid_relay_data_preference(&self, _avoid_relay_data: bool) -> bool { + false + } + + fn subscribe_runtime_changes(&self) -> Option> { + None + } + + fn easytier_version(&self) -> String { + env!("CARGO_PKG_VERSION").to_string() + } + + fn ospf_update_my_foreign_network_interval_sec(&self) -> u64 { + 10 + } + + fn max_direct_conns_per_peer_in_foreign_network(&self) -> usize { + 3 + } + + fn hmac_secret_digest(&self) -> bool { + false + } + + fn advertised_ipv6_public_addr_prefix(&self) -> Option { + None + } + + fn is_ip_in_same_network(&self, _ip: &IpAddr) -> bool { + false + } + + fn peer_groups(&self, _peer_id: PeerId) -> Vec { + Vec::new() + } + + fn acl_group_declarations(&self) -> Vec { + Vec::new() + } + + fn pinned_remote_static_pubkey(&self, _tunnel_info: Option<&TunnelInfo>) -> Option { + None + } + + fn secret_proof(&self, _challenge: &[u8]) -> Option> { + None + } + + fn secret_digest(&self, network_identity: &NetworkIdentity) -> Vec { + network_identity + .secret_digest() + .unwrap_or_default() + .to_vec() + } + + fn is_pubkey_trusted(&self, _pubkey: &[u8], _network_name: &str) -> bool { + false + } + + fn is_pubkey_trusted_with_source( + &self, + _pubkey: &[u8], + _network_name: &str, + _source: TrustedKeySource, + ) -> bool { + false + } + + fn list_trusted_keys(&self, _network_name: &str) -> Vec<(Vec, TrustedKeyMetadata)> { + Vec::new() + } + + fn trusted_credential_pubkeys( + &self, + _network_secret: &str, + ) -> Vec { + Vec::new() + } + + fn remove_expired_credentials(&self) -> bool { + false + } + + fn issue_credential_changed(&self) {} + + fn update_trusted_keys(&self, _keys: TrustedKeyMap, _network_name: &str) {} + + fn remove_trusted_keys(&self, _network_name: &str) {} + + fn record_control_tx(&self, _network_name: &str, _bytes: u64) {} + + fn record_control_rx(&self, _network_name: &str, _bytes: u64) {} + + fn recv_limiter( + &self, + _network_name: &str, + _is_foreign_network: bool, + ) -> Option { + None + } + + fn foreign_forward_limiter(&self, _network_name: &str) -> Option { + None + } + + fn issue_event(&self, _event: PeerEvent) {} + + fn subscribe_peer_events(&self) -> Option { + None + } +} + +pub(crate) type ArcPeerContext = Arc; + +pub(crate) fn secret_proof_from_secret(secret: &str, challenge: &[u8]) -> Option> { + let mut mac = Hmac::::new_from_slice(secret.as_bytes()).ok()?; + mac.update(SECRET_PROOF_PREFIX); + mac.update(challenge); + Some(mac) +} + +impl PeerContext for CorePeerContext { + fn max_direct_conns_per_peer_in_foreign_network(&self) -> usize { + self.snapshot().max_direct_conns_per_peer_in_foreign_network + } + + fn network_identity(&self) -> NetworkIdentity { + self.snapshot().runtime.network_identity.clone() + } + + fn flags(&self) -> FlagsInConfig { + self.snapshot().flags.clone() + } + + fn host_routing_policy(&self) -> HostRoutingPolicy { + self.snapshot().runtime.host_routing + } + + fn secure_mode(&self) -> Option { + self.snapshot().runtime.secure_mode.clone() + } + + fn stun_info(&self) -> StunInfo { + self.stun_info_source + .as_ref() + .map(|source| source.stun_info()) + .unwrap_or_else(|| self.snapshot().runtime.stun_info.clone()) + } + + fn instance_id(&self) -> uuid::Uuid { + self.snapshot() + .runtime + .core + .node + .instance_id + .map(uuid::Uuid::from_bytes) + .unwrap_or_else(uuid::Uuid::nil) + } + + fn ipv4(&self) -> Option { + self.snapshot() + .runtime + .core + .routes + .ipv4 + .as_ref() + .and_then(config_ipv4) + } + + fn ipv6(&self) -> Option { + self.snapshot() + .runtime + .core + .routes + .ipv6 + .as_ref() + .and_then(config_ipv6) + } + + fn is_ip_local_ipv6(&self, ip: &std::net::Ipv6Addr) -> bool { + self.ipv6().is_some_and(|address| address.address() == *ip) + || self.public_ipv6_state.public_ipv6_lease_contains(ip) + } + + fn proxy_cidrs(&self) -> Vec { + self.snapshot() + .runtime + .core + .routes + .proxy_networks + .iter() + .filter_map(|proxy| config_ipv4_cidr(proxy.mapped.as_ref().unwrap_or(&proxy.real))) + .collect() + } + + fn proxy_networks(&self) -> Vec { + self.snapshot().runtime.core.routes.proxy_networks.clone() + } + + fn vpn_portal_cidr(&self) -> Option { + self.snapshot().vpn_portal_cidr + } + + fn hostname(&self) -> String { + self.snapshot() + .runtime + .core + .node + .hostname + .clone() + .unwrap_or_default() + } + + fn feature_flags(&self) -> PeerFeatureFlag { + let snapshot = self.snapshot(); + let mut flags = snapshot.runtime.feature_flags; + flags.avoid_relay_data = snapshot.flags.disable_relay_data + || self.avoid_relay_data_preference.load(Ordering::Acquire); + flags.ipv6_public_addr_provider |= self.public_ipv6_state.public_ipv6_provider_enabled(); + flags + } + + fn set_avoid_relay_data_preference(&self, avoid_relay_data: bool) -> bool { + let before = self.feature_flags().avoid_relay_data; + self.avoid_relay_data_preference + .store(avoid_relay_data, Ordering::Release); + before != self.feature_flags().avoid_relay_data + } + + fn subscribe_runtime_changes(&self) -> Option> { + Some(self.config.subscribe_peer_runtime_changes()) + } + + fn easytier_version(&self) -> String { + self.snapshot().easytier_version.clone() + } + + fn ospf_update_my_foreign_network_interval_sec(&self) -> u64 { + self.snapshot().ospf_update_my_foreign_network_interval_sec + } + + fn hmac_secret_digest(&self) -> bool { + self.snapshot().hmac_secret_digest + } + + fn advertised_ipv6_public_addr_prefix(&self) -> Option { + self.public_ipv6_state.advertised_ipv6_public_addr_prefix() + } + + fn is_ip_in_same_network(&self, ip: &IpAddr) -> bool { + match ip { + IpAddr::V4(ip) => self.ipv4().is_some_and(|network| network.contains(ip)), + IpAddr::V6(ip) => self.ipv6().is_some_and(|network| network.contains(ip)), + } + } + + fn pinned_remote_static_pubkey(&self, tunnel_info: Option<&TunnelInfo>) -> Option { + let remote_url = tunnel_info + .and_then(|info| info.remote_addr.as_ref())? + .url + .parse::() + .ok()?; + self.snapshot() + .pinned_peers + .iter() + .find(|(uri, _)| *uri == remote_url) + .and_then(|(_, public_key)| public_key.clone()) + } + + fn secret_proof(&self, challenge: &[u8]) -> Option> { + let snapshot = self.snapshot(); + let secret = snapshot.runtime.network_identity.network_secret.as_ref()?; + secret_proof_from_secret(secret, challenge) + } + + fn secret_digest(&self, network_identity: &NetworkIdentity) -> Vec { + let snapshot = self.snapshot(); + if snapshot.hmac_secret_digest { + snapshot + .runtime + .network_identity + .network_secret + .as_deref() + .and_then(|secret| secret_proof_from_secret(secret, b"digest")) + .map(|mac| mac.finalize().into_bytes().to_vec()) + .unwrap_or_default() + } else { + network_identity + .secret_digest() + .unwrap_or_default() + .to_vec() + } + } + + fn peer_groups(&self, peer_id: PeerId) -> Vec { + self.snapshot() + .peer_group_memberships + .iter() + .map(|group| { + PeerGroupInfo::generate_with_proof( + group.group_name.clone(), + group.group_secret.clone(), + peer_id, + ) + }) + .collect() + } + + fn acl_group_declarations(&self) -> Vec { + self.snapshot().acl_group_declarations.clone() + } + + fn is_pubkey_trusted(&self, pubkey: &[u8], network_name: &str) -> bool { + if self.trusted_keys.verify_trusted_key(pubkey, network_name) { + return true; + } + network_name == self.snapshot().runtime.network_identity.network_name + && self.credentials.is_pubkey_trusted(pubkey) + } + + fn is_pubkey_trusted_with_source( + &self, + pubkey: &[u8], + network_name: &str, + source: TrustedKeySource, + ) -> bool { + self.trusted_keys + .verify_trusted_key_with_source(pubkey, network_name, Some(source)) + } + + fn list_trusted_keys(&self, network_name: &str) -> Vec<(Vec, TrustedKeyMetadata)> { + self.trusted_keys.list_trusted_keys(network_name) + } + + fn trusted_credential_pubkeys( + &self, + network_secret: &str, + ) -> Vec { + self.credentials.get_trusted_pubkeys(network_secret) + } + + fn remove_expired_credentials(&self) -> bool { + self.credentials.remove_expired_credentials() + } + + fn issue_credential_changed(&self) { + self.events.emit(CoreEvent::CredentialChanged); + } + + fn update_trusted_keys(&self, keys: TrustedKeyMap, network_name: &str) { + self.trusted_keys.update_trusted_keys(network_name, keys); + } + + fn remove_trusted_keys(&self, network_name: &str) { + self.trusted_keys.remove_trusted_keys(network_name); + } + + fn record_control_tx(&self, network_name: &str, bytes: u64) { + self.record_control_metric( + network_name, + bytes, + MetricName::TrafficControlBytesTx, + MetricName::TrafficControlPacketsTx, + ); + } + + fn record_control_rx(&self, network_name: &str, bytes: u64) { + self.record_control_metric( + network_name, + bytes, + MetricName::TrafficControlBytesRx, + MetricName::TrafficControlPacketsRx, + ); + } + + fn recv_limiter(&self, network_name: &str, is_foreign_network: bool) -> Option { + let limits = self.snapshot().traffic_limits(); + let (key, bps) = if is_foreign_network && let Some(limit) = limits.foreign_relay_bps { + (format!("peer:foreign:{network_name}:recv"), limit) + } else { + ("peer:instance:recv".to_owned(), limits.instance_recv_bps?) + }; + self.get_or_create_limiter(&key, bps) + } + + fn foreign_forward_limiter(&self, network_name: &str) -> Option { + let bps = self.snapshot().traffic_limits().foreign_relay_bps?; + self.get_or_create_limiter(&format!("peer:foreign:{network_name}:forward"), bps) + } + + fn issue_event(&self, event: PeerEvent) { + let context_event = match &event { + PeerEvent::PeerAdded(peer_id) => PeerContextEvent::PeerAdded(*peer_id), + PeerEvent::PeerRemoved(peer_id) => PeerContextEvent::PeerRemoved(*peer_id), + PeerEvent::PeerConnAdded(_) => PeerContextEvent::PeerConnAdded, + PeerEvent::PeerConnRemoved(_) => PeerContextEvent::PeerConnRemoved, + }; + let _ = self.peer_events.send(context_event); + let event = match event { + PeerEvent::PeerAdded(peer_id) => CoreEvent::PeerAdded(peer_id), + PeerEvent::PeerRemoved(peer_id) => CoreEvent::PeerRemoved(peer_id), + PeerEvent::PeerConnAdded(info) => CoreEvent::PeerConnAdded(info), + PeerEvent::PeerConnRemoved(info) => CoreEvent::PeerConnRemoved(info), + }; + self.events.emit(event); + } + + fn subscribe_peer_events(&self) -> Option { + Some(self.peer_events.subscribe()) + } +} + +#[cfg(test)] +pub(crate) mod tests { + #![allow(clippy::field_reassign_with_default)] + + use super::*; + use crate::config::runtime::{CoreRuntimeConfig, CoreRuntimeConfigStore}; + use crate::peers::test_support::{NoopPeerContext, PeerContextTestExt}; + use std::sync::atomic::{AtomicBool, Ordering}; + + impl PeerRuntimeSnapshot { + pub(crate) fn set_acl_groups(&mut self, acl: Option<&Acl>) { + (self.acl_group_declarations, self.peer_group_memberships) = peer_acl_groups(acl); + } + } + + impl CorePeerContext { + pub(crate) fn trusted_key_manager(&self) -> Arc { + self.trusted_keys.clone() + } + } + + impl PeerContextTestExt for CorePeerContext { + fn runtime_config(&self) -> PeerRuntimeConfig { + let mut runtime = self.snapshot().runtime.clone(); + runtime.stun_info = self.stun_info(); + runtime + } + } + + #[derive(Default)] + struct TestPeerEventSink { + events: Mutex>, + credential_changed: AtomicBool, + } + + impl CoreEventSink for TestPeerEventSink { + fn emit(&self, event: CoreEvent) { + if matches!(&event, CoreEvent::CredentialChanged) { + self.credential_changed.store(true, Ordering::Release); + } + self.events.lock().unwrap().push(event); + } + } + + fn test_core_context_adapters(events: Arc) -> CorePeerContextAdapters { + CorePeerContextAdapters { + stun_info_source: Some(Arc::new(())), + events, + credential_storage: None, + } + } + + fn submitted_snapshot(hostname: &str, disable_relay_data: bool) -> PeerRuntimeSnapshot { + let mut flags = FlagsInConfig::default(); + flags.disable_relay_data = disable_relay_data; + PeerRuntimeSnapshot::new( + PeerRuntimeConfig { + core: CoreConfig { + node: NodeConfig { + hostname: Some(hostname.to_owned()), + ..Default::default() + }, + ..Default::default() + }, + network_identity: NetworkIdentity::default(), + stun_info: StunInfo::default(), + feature_flags: PeerFeatureFlag::default(), + secure_mode: None, + host_routing: HostRoutingPolicy::default(), + }, + flags, + ) + } + + fn host_snapshot_input(flags: FlagsInConfig, acl: Option) -> PeerRuntimeSnapshotInput { + PeerRuntimeSnapshotInput { + node: NodeConfig { + peer_id: None, + instance_id: Some([7; 16]), + hostname: Some("host-node".to_owned()), + network_name: "host-network".to_owned(), + }, + routes: RouteConfig { + ipv4: Some(IpPrefix::new("10.20.0.7".parse().unwrap(), 16).unwrap()), + ..Default::default() + }, + network_identity: NetworkIdentity::new("host-network".to_owned(), "secret".to_owned()), + stun_info: StunInfo::default(), + flags, + secure_mode: None, + host_routing: HostRoutingPolicy { + local_exit_node_fallback: true, + }, + acl, + easytier_version: "host-version".to_owned(), + vpn_portal_cidr: Some("10.30.0.0/24".parse().unwrap()), + pinned_peers: vec![( + "tcp://192.0.2.10:11010".parse().unwrap(), + Some("peer-key".to_owned()), + )], + ospf_update_my_foreign_network_interval_sec: 17, + max_direct_conns_per_peer_in_foreign_network: 5, + hmac_secret_digest: true, + } + } + + #[test] + fn host_input_derives_peer_policy_features_and_traffic() { + let mut flags = FlagsInConfig::default(); + flags.disable_p2p = true; + flags.need_p2p = true; + flags.relay_all_peer_rpc = true; + flags.disable_relay_data = true; + flags.latency_first = true; + flags.enable_encryption = false; + flags.disable_kcp_input = true; + flags.disable_relay_kcp = true; + flags.disable_quic_input = true; + flags.disable_relay_quic = true; + flags.mtu = 1400; + flags.instance_recv_bps_limit = 0; + flags.foreign_relay_bps_limit = u64::MAX; + flags.relay_network_whitelist = "host-network".to_owned(); + + let snapshot = + PeerRuntimeSnapshot::from_host_input(host_snapshot_input(flags.clone(), None)); + let runtime = &snapshot.runtime; + + assert_eq!(snapshot.flags, flags); + assert!(!runtime.core.peer_policy.p2p_enabled); + assert!(runtime.core.peer_policy.relay_peer_rpc); + assert!(!runtime.core.peer_policy.relay_data); + assert!(runtime.core.peer_policy.latency_first); + assert!(!runtime.core.peer_policy.encryption_required); + assert_eq!(runtime.core.traffic.mtu, Some(1400)); + assert_eq!(runtime.core.traffic.instance_recv_bps_limit, Some(0)); + assert_eq!(runtime.core.traffic.foreign_relay_bps_limit, None); + assert!(!runtime.feature_flags.kcp_input); + assert!(runtime.feature_flags.no_relay_kcp); + assert!(runtime.feature_flags.support_conn_list_sync); + assert!(!runtime.feature_flags.quic_input); + assert!(runtime.feature_flags.no_relay_quic); + assert!(runtime.feature_flags.need_p2p); + assert!(runtime.feature_flags.disable_p2p); + assert!(runtime.feature_flags.avoid_relay_data); + assert!(!snapshot.avoid_relay_data_preference); + } + + #[test] + fn host_input_derives_acl_groups_and_preserves_explicit_inputs() { + let acl = Acl { + acl_v1: Some(easytier_proto::acl::AclV1 { + chains: Vec::new(), + group: Some(easytier_proto::acl::GroupInfo { + declares: vec![ + easytier_proto::acl::GroupIdentity { + group_name: "ops".to_owned(), + group_secret: "ops-secret".to_owned(), + }, + easytier_proto::acl::GroupIdentity { + group_name: "audit".to_owned(), + group_secret: "audit-secret".to_owned(), + }, + ], + members: vec!["ops".to_owned(), "undeclared".to_owned()], + }), + }), + }; + let mut flags = FlagsInConfig::default(); + flags.relay_network_whitelist = "other-network".to_owned(); + + let snapshot = PeerRuntimeSnapshot::from_host_input(host_snapshot_input(flags, Some(acl))); + + assert!(snapshot.avoid_relay_data_preference); + assert_eq!(snapshot.easytier_version, "host-version"); + assert_eq!( + snapshot.vpn_portal_cidr, + Some("10.30.0.0/24".parse().unwrap()) + ); + assert_eq!( + snapshot.pinned_peers, + vec![( + "tcp://192.0.2.10:11010".parse().unwrap(), + Some("peer-key".to_owned()) + )] + ); + assert_eq!(snapshot.ospf_update_my_foreign_network_interval_sec, 17); + assert_eq!(snapshot.max_direct_conns_per_peer_in_foreign_network, 5); + assert!(snapshot.hmac_secret_digest); + assert!(snapshot.runtime.host_routing.local_exit_node_fallback); + assert_eq!( + snapshot.acl_group_declarations, + vec![ + PeerGroupIdentity { + group_name: "ops".to_owned(), + group_secret: "ops-secret".to_owned(), + }, + PeerGroupIdentity { + group_name: "audit".to_owned(), + group_secret: "audit-secret".to_owned(), + }, + ] + ); + assert_eq!( + snapshot.peer_group_memberships, + vec![PeerGroupIdentity { + group_name: "ops".to_owned(), + group_secret: "ops-secret".to_owned(), + }] + ); + } + + #[test] + fn records_control_traffic_in_core_owned_metrics() { + let context = CorePeerContext::new( + CoreRuntimeConfigStore::new( + CoreRuntimeConfig::default(), + Arc::new(PeerRuntimeSnapshot::default()), + ), + Arc::new(()), + test_core_context_adapters(Arc::new(())), + ); + let labels = + LabelSet::new().with_label_type(LabelType::NetworkName("metrics-network".to_owned())); + + PeerContext::record_control_tx(&context, "metrics-network", 128); + + assert_eq!( + context + .stats_manager() + .get_metric(MetricName::TrafficControlBytesTx, &labels) + .unwrap() + .value, + 128 + ); + assert_eq!( + context + .stats_manager() + .get_metric(MetricName::TrafficControlPacketsTx, &labels) + .unwrap() + .value, + 1 + ); + } + + #[test] + fn explicit_traffic_limits_preserve_zero_and_override_flags() { + let mut runtime = PeerRuntimeSnapshot::default().runtime; + runtime.core.traffic.instance_recv_bps_limit = Some(0); + runtime.core.traffic.foreign_relay_bps_limit = Some(2048); + let mut flags = FlagsInConfig::default(); + flags.instance_recv_bps_limit = 1024; + flags.foreign_relay_bps_limit = 4096; + + let snapshot = PeerRuntimeSnapshot::new(runtime, flags); + + assert_eq!(snapshot.traffic_limits().instance_recv_bps, Some(0)); + assert_eq!(snapshot.traffic_limits().foreign_relay_bps, Some(2048)); + } + + #[test] + fn legacy_traffic_limits_ignore_unlimited_sentinels() { + let runtime = PeerRuntimeSnapshot::default().runtime; + let mut flags = FlagsInConfig::default(); + flags.instance_recv_bps_limit = 1024; + flags.foreign_relay_bps_limit = u64::MAX; + + let snapshot = PeerRuntimeSnapshot::new(runtime, flags); + + assert_eq!(snapshot.traffic_limits().instance_recv_bps, Some(1024)); + assert_eq!(snapshot.traffic_limits().foreign_relay_bps, None); + } + + #[test] + fn explicit_unlimited_limits_override_legacy_values() { + let mut runtime = PeerRuntimeSnapshot::default().runtime; + runtime.core.traffic.instance_recv_bps_limit = Some(u64::MAX); + runtime.core.traffic.foreign_relay_bps_limit = Some(u64::MAX); + let mut flags = FlagsInConfig::default(); + flags.instance_recv_bps_limit = 1024; + flags.foreign_relay_bps_limit = 2048; + + let snapshot = PeerRuntimeSnapshot::new(runtime, flags); + + assert_eq!(snapshot.traffic_limits(), PeerTrafficLimits::default()); + } + + #[test] + fn portable_traffic_limits_default_to_unlimited() { + let runtime = PeerRuntimeSnapshot::default().runtime; + + let snapshot = PeerRuntimeSnapshot::new(runtime, FlagsInConfig::default()); + + assert_eq!(snapshot.traffic_limits(), PeerTrafficLimits::default()); + } + + fn submitted_config(snapshot: PeerRuntimeSnapshot) -> CoreRuntimeConfigStore { + CoreRuntimeConfigStore::new(CoreRuntimeConfig::default(), Arc::new(snapshot)) + } + + #[test] + fn core_peer_context_separates_config_versions_from_live_support() { + let config = submitted_config(submitted_snapshot("before", false)); + let context = CorePeerContext::new( + config.clone(), + Arc::new(()), + test_core_context_adapters(Arc::new(())), + ); + + assert_eq!(context.hostname(), "before"); + assert!(!context.feature_flags().avoid_relay_data); + + context.set_avoid_relay_data_preference(true); + assert!(context.feature_flags().avoid_relay_data); + + config.update_peer(Arc::new(submitted_snapshot("after", true))); + context.set_avoid_relay_data_preference(false); + assert_eq!(context.hostname(), "after"); + assert!(context.feature_flags().avoid_relay_data); + + config.update_peer(Arc::new(submitted_snapshot("after", false))); + assert!(!context.feature_flags().avoid_relay_data); + } + + #[tokio::test] + async fn core_peer_context_owns_events_and_projects_them_to_sink() { + let config = submitted_config(submitted_snapshot("events", false)); + let sink = Arc::new(TestPeerEventSink::default()); + let context = CorePeerContext::new( + config, + Arc::new(()), + test_core_context_adapters(sink.clone()), + ); + let mut events = context.subscribe_peer_events().unwrap(); + + context.issue_event(PeerEvent::PeerAdded(7)); + + assert_eq!(events.recv().await.unwrap(), PeerContextEvent::PeerAdded(7)); + assert!(matches!( + sink.events.lock().unwrap().as_slice(), + [CoreEvent::PeerAdded(7)] + )); + } + + #[test] + fn core_peer_context_owns_trusted_keys_and_projects_credential_changes() { + let config = submitted_config(submitted_snapshot("trust", false)); + let events = Arc::new(TestPeerEventSink::default()); + let context = CorePeerContext::new( + config, + Arc::new(()), + test_core_context_adapters(events.clone()), + ); + let public_key = vec![7; 32]; + let mut keys = TrustedKeyMap::new(); + keys.insert( + public_key.clone(), + TrustedKeyMetadata { + source: TrustedKeySource::OspfNode, + expiry_unix: None, + }, + ); + + context.update_trusted_keys(keys, "foreign"); + context.issue_credential_changed(); + + assert!(context.is_pubkey_trusted(&public_key, "foreign")); + assert_eq!(context.list_trusted_keys("foreign").len(), 1); + assert!(events.credential_changed.load(Ordering::Acquire)); + } + + #[tokio::test] + async fn foreign_forward_limiter_is_independent_from_peer_receive_limiter() { + let mut snapshot = submitted_snapshot("limiter", false); + snapshot.runtime.core.traffic.foreign_relay_bps_limit = Some(1024); + let config = submitted_config(snapshot); + let context = CorePeerContext::new( + config, + Arc::new(()), + test_core_context_adapters(Arc::new(())), + ); + + let receive = context.recv_limiter("foreign", true).unwrap(); + let receive_again = context.recv_limiter("foreign", true).unwrap(); + let forward = context.foreign_forward_limiter("foreign").unwrap(); + assert!(Arc::ptr_eq(&receive, &receive_again)); + assert!(!Arc::ptr_eq(&receive, &forward)); + context.stop().await; + } + + #[tokio::test] + async fn foreign_forward_limiter_does_not_fall_back_to_instance_limit() { + let mut snapshot = submitted_snapshot("limiter", false); + snapshot.runtime.core.traffic.foreign_relay_bps_limit = None; + snapshot.runtime.core.traffic.instance_recv_bps_limit = Some(1024); + let config = submitted_config(snapshot); + let context = CorePeerContext::new( + config, + Arc::new(()), + test_core_context_adapters(Arc::new(())), + ); + + assert!(context.recv_limiter("foreign", true).is_some()); + assert!(context.foreign_forward_limiter("foreign").is_none()); + context.stop().await; + } + + #[test] + fn noop_peer_context_uses_runtime_secret_proof_prefix() { + let context = NoopPeerContext::new(NetworkIdentity { + network_name: "net".to_string(), + network_secret: Some("secret".to_string()), + network_secret_digest: None, + }); + + let proof = context + .secret_proof(b"challenge") + .unwrap() + .finalize() + .into_bytes() + .to_vec(); + let expected = secret_proof_from_secret("secret", b"challenge") + .unwrap() + .finalize() + .into_bytes() + .to_vec(); + + assert_eq!(proof, expected); + } + + struct RuntimeConfigContext { + instance_id: uuid::Uuid, + } + + impl PeerContext for RuntimeConfigContext { + fn network_identity(&self) -> NetworkIdentity { + NetworkIdentity { + network_name: "net".to_string(), + network_secret: Some("secret".to_string()), + network_secret_digest: None, + } + } + + fn instance_id(&self) -> uuid::Uuid { + self.instance_id + } + + fn ipv4(&self) -> Option { + Some("10.1.0.1/24".parse().unwrap()) + } + + fn ipv6(&self) -> Option { + Some("2001:db8::1/64".parse().unwrap()) + } + + fn hostname(&self) -> String { + "node-a".to_string() + } + } + + impl PeerContextTestExt for RuntimeConfigContext {} + + #[test] + fn runtime_config_preserves_peer_context_snapshot() { + let instance_id = uuid::Uuid::from_u128(0x11223344556677889900aabbccddeeff); + let context = RuntimeConfigContext { instance_id }; + + let config = context.runtime_config(); + + assert_eq!(config.network_identity.network_name, "net"); + assert_eq!(config.core.node.instance_id, Some(*instance_id.as_bytes())); + assert_eq!(config.core.node.hostname.as_deref(), Some("node-a")); + assert_eq!(config.core.node.network_name, "net"); + assert_eq!( + config.core.routes.ipv4, + Some(IpPrefix::new("10.1.0.1".parse().unwrap(), 24).unwrap()) + ); + assert_eq!( + config.core.routes.ipv6, + Some(IpPrefix::new("2001:db8::1".parse().unwrap(), 64).unwrap()) + ); + } + + #[test] + fn core_owned_peer_context_reads_normalized_snapshot() { + let instance_id = uuid::Uuid::from_u128(0x00112233445566778899aabbccddeeff); + let runtime = PeerRuntimeConfig { + core: CoreConfig { + node: NodeConfig { + peer_id: Some(7), + instance_id: Some(*instance_id.as_bytes()), + hostname: Some("config-node".to_owned()), + network_name: "config-net".to_owned(), + }, + routes: RouteConfig { + ipv4: Some(IpPrefix::new("10.20.0.7".parse().unwrap(), 16).unwrap()), + ipv6: Some(IpPrefix::new("2001:db8::7".parse().unwrap(), 64).unwrap()), + proxy_networks: vec![ + crate::config::ProxyNetworkConfig { + real: IpPrefix::new("10.40.0.0".parse().unwrap(), 16).unwrap(), + mapped: Some(IpPrefix::new("10.50.0.0".parse().unwrap(), 16).unwrap()), + }, + crate::config::ProxyNetworkConfig { + real: IpPrefix::new("10.60.0.0".parse().unwrap(), 16).unwrap(), + mapped: None, + }, + ], + ..Default::default() + }, + ..Default::default() + }, + network_identity: NetworkIdentity { + network_name: "config-net".to_owned(), + network_secret: Some("secret".to_owned()), + network_secret_digest: None, + }, + stun_info: StunInfo::default(), + feature_flags: PeerFeatureFlag::default(), + secure_mode: Some(SecureModeConfig { + enabled: true, + ..Default::default() + }), + host_routing: HostRoutingPolicy { + local_exit_node_fallback: true, + }, + }; + let mut flags = FlagsInConfig::default(); + flags.p2p_only = true; + let acl = Acl { + acl_v1: Some(easytier_proto::acl::AclV1 { + chains: Vec::new(), + group: Some(easytier_proto::acl::GroupInfo { + declares: vec![easytier_proto::acl::GroupIdentity { + group_name: "ops".to_string(), + group_secret: "group-secret".to_string(), + }], + members: vec!["ops".to_string()], + }), + }), + }; + let context = core_owned_context_with_acl(runtime.clone(), flags.clone(), None, Some(&acl)); + + assert_eq!(context.runtime_config().core, runtime.core); + assert_eq!(context.network_identity(), runtime.network_identity); + assert_eq!(context.flags(), flags); + assert_eq!(context.instance_id(), instance_id); + assert_eq!(context.hostname(), "config-node"); + assert_eq!(context.ipv4(), Some("10.20.0.7/16".parse().unwrap())); + assert_eq!(context.ipv6(), Some("2001:db8::7/64".parse().unwrap())); + assert!(context.is_ip_in_same_network(&"10.20.99.1".parse().unwrap())); + assert!(context.is_ip_in_same_network(&"2001:db8::99".parse().unwrap())); + assert!(!context.is_ip_in_same_network(&"10.21.0.1".parse().unwrap())); + assert_eq!( + context.proxy_cidrs(), + vec![ + "10.50.0.0/16".parse().unwrap(), + "10.60.0.0/16".parse().unwrap() + ] + ); + assert!(context.secure_mode().unwrap().enabled); + assert!(context.host_routing_policy().local_exit_node_fallback); + + let groups = context.peer_groups(7); + assert_eq!(groups.len(), 1); + assert_eq!(groups[0].group_name, "ops"); + assert!(groups[0].verify("group-secret", 7)); + assert_eq!( + context.acl_group_declarations(), + vec![PeerGroupIdentity { + group_name: "ops".to_string(), + group_secret: "group-secret".to_string(), + }] + ); + + let proof = context + .secret_proof(b"challenge") + .unwrap() + .finalize() + .into_bytes(); + let expected = secret_proof_from_secret("secret", b"challenge") + .unwrap() + .finalize() + .into_bytes(); + assert_eq!(proof, expected); + } + + #[test] + fn core_owned_peer_context_uses_live_stun_source_when_injected() { + struct TestStunInfoSource(StunInfo); + + impl PeerStunInfoSource for TestStunInfoSource { + fn stun_info(&self) -> StunInfo { + self.0.clone() + } + } + + let mut runtime = PeerRuntimeSnapshot::default().runtime; + runtime.stun_info.tcp_nat_type = 1; + let mut live = StunInfo::default(); + live.tcp_nat_type = 4; + let context = core_owned_context( + runtime, + FlagsInConfig::default(), + Some(Arc::new(TestStunInfoSource(live.clone()))), + ); + + assert_eq!(context.stun_info(), live); + assert_eq!(context.runtime_config().stun_info, live); + } + + fn core_owned_context( + runtime: PeerRuntimeConfig, + flags: FlagsInConfig, + stun_info_source: Option>, + ) -> CorePeerContext { + core_owned_context_with_acl(runtime, flags, stun_info_source, None) + } + + fn core_owned_context_with_acl( + runtime: PeerRuntimeConfig, + flags: FlagsInConfig, + stun_info_source: Option>, + acl: Option<&Acl>, + ) -> CorePeerContext { + let mut snapshot = PeerRuntimeSnapshot::new(runtime, flags); + snapshot.set_acl_groups(acl); + let config = CoreRuntimeConfigStore::new(CoreRuntimeConfig::default(), Arc::new(snapshot)); + CorePeerContext::new( + config, + Arc::new(()), + CorePeerContextAdapters { + stun_info_source, + events: Arc::new(()), + credential_storage: None, + }, + ) + } + + fn core_context_with_routes(ipv4: Option, ipv6: Option) -> CorePeerContext { + core_owned_context( + PeerRuntimeConfig { + core: CoreConfig { + routes: RouteConfig { + ipv4, + ipv6, + ..Default::default() + }, + ..Default::default() + }, + network_identity: NetworkIdentity::default(), + stun_info: StunInfo::default(), + feature_flags: PeerFeatureFlag::default(), + secure_mode: None, + host_routing: HostRoutingPolicy::default(), + }, + FlagsInConfig::default(), + None, + ) + } + + #[tokio::test] + async fn core_owned_peer_context_publishes_peer_events_per_instance() { + let context = core_context_with_routes(None, None); + let mut events = context.subscribe_peer_events().unwrap(); + + context.issue_event(PeerEvent::PeerAdded(7)); + assert_eq!(events.recv().await.unwrap(), PeerContextEvent::PeerAdded(7)); + context.issue_event(PeerEvent::PeerConnAdded(Default::default())); + assert_eq!( + events.recv().await.unwrap(), + PeerContextEvent::PeerConnAdded + ); + context.issue_event(PeerEvent::PeerConnRemoved(Default::default())); + assert_eq!( + events.recv().await.unwrap(), + PeerContextEvent::PeerConnRemoved + ); + context.issue_event(PeerEvent::PeerRemoved(7)); + assert_eq!( + events.recv().await.unwrap(), + PeerContextEvent::PeerRemoved(7) + ); + } + + #[test] + fn core_owned_peer_context_validates_route_family_and_prefix_edges() { + let mismatched = core_context_with_routes( + Some(IpPrefix { + address: "2001:db8::1".parse().unwrap(), + prefix_len: 64, + }), + Some(IpPrefix { + address: "10.20.0.1".parse().unwrap(), + prefix_len: 24, + }), + ); + assert_eq!(mismatched.ipv4(), None); + assert_eq!(mismatched.ipv6(), None); + assert!(!mismatched.is_ip_in_same_network(&"2001:db8::2".parse().unwrap())); + assert!(!mismatched.is_ip_in_same_network(&"10.20.0.2".parse().unwrap())); + + let edges = core_context_with_routes( + Some(IpPrefix::new("10.20.0.7".parse().unwrap(), 0).unwrap()), + Some(IpPrefix::new("2001:db8::7".parse().unwrap(), 128).unwrap()), + ); + assert!(edges.is_ip_in_same_network(&"203.0.113.1".parse().unwrap())); + assert!(edges.is_ip_in_same_network(&"2001:db8::7".parse().unwrap())); + assert!(!edges.is_ip_in_same_network(&"2001:db8::8".parse().unwrap())); + + let invalid = core_context_with_routes( + Some(IpPrefix { + address: "10.20.0.7".parse().unwrap(), + prefix_len: 33, + }), + None, + ); + assert_eq!(invalid.ipv4(), None); + assert!(!invalid.is_ip_in_same_network(&"10.20.0.7".parse().unwrap())); + } + + #[test] + fn trusted_key_manager_respects_source_filter() { + let manager = TrustedKeyMapManager::new(); + let network_name = "net"; + let pubkey = vec![1; 32]; + manager.update_trusted_keys( + network_name, + HashMap::from([( + pubkey.clone(), + TrustedKeyMetadata { + source: TrustedKeySource::OspfCredential, + expiry_unix: None, + }, + )]), + ); + + assert!(manager.verify_trusted_key(&pubkey, network_name)); + assert!(manager.verify_trusted_key_with_source( + &pubkey, + network_name, + Some(TrustedKeySource::OspfCredential), + )); + assert!(!manager.verify_trusted_key_with_source( + &pubkey, + network_name, + Some(TrustedKeySource::OspfNode), + )); + } +} diff --git a/easytier-core/src/peers/credential_manager.rs b/easytier-core/src/peers/credential_manager.rs new file mode 100644 index 00000000..b77d1f62 --- /dev/null +++ b/easytier-core/src/peers/credential_manager.rs @@ -0,0 +1,497 @@ +use std::{ + collections::HashMap, + sync::{Arc, Mutex}, + time::{Duration, SystemTime, UNIX_EPOCH}, +}; + +use base64::{Engine, engine::general_purpose::STANDARD as BASE64_STANDARD}; +use serde::{Deserialize, Serialize}; +use x25519_dalek::{PublicKey, StaticSecret}; + +use crate::proto::peer_rpc::{TrustedCredentialPubkey, TrustedCredentialPubkeyProof}; + +fn default_true() -> bool { + true +} + +fn current_unix_timestamp() -> i64 { + SystemTime::now() + .duration_since(UNIX_EPOCH) + .unwrap() + .as_secs() as i64 +} + +#[derive(Debug, Clone)] +pub struct CredentialCreateOptions { + pub groups: Vec, + pub allow_relay: bool, + pub allowed_proxy_cidrs: Vec, + pub ttl: Duration, + pub credential_id: Option, + pub reusable: bool, +} + +#[derive(Debug, Clone, Serialize, Deserialize)] +pub(crate) struct CredentialEntry { + pubkey: String, + #[serde(default)] + secret: String, + groups: Vec, + allow_relay: bool, + allowed_proxy_cidrs: Vec, + #[serde(default = "default_true")] + reusable: bool, + expiry_unix: i64, + created_at_unix: i64, +} + +impl CredentialEntry { + fn is_active_at(&self, now: i64) -> bool { + self.expiry_unix > now + } + + fn to_trusted_credential(&self) -> Option { + Some(TrustedCredentialPubkey { + pubkey: CredentialManager::decode_pubkey_b64(&self.pubkey)?, + groups: self.groups.clone(), + allow_relay: self.allow_relay, + expiry_unix: self.expiry_unix, + allowed_proxy_cidrs: self.allowed_proxy_cidrs.clone(), + reusable: Some(self.reusable), + }) + } + + fn to_credential_info(&self, credential_id: &str) -> CredentialInfo { + CredentialInfo { + credential_id: credential_id.to_string(), + groups: self.groups.clone(), + allow_relay: self.allow_relay, + expiry_unix: self.expiry_unix, + allowed_proxy_cidrs: self.allowed_proxy_cidrs.clone(), + reusable: Some(self.reusable), + } + } +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct CredentialInfo { + pub credential_id: String, + pub groups: Vec, + pub allow_relay: bool, + pub expiry_unix: i64, + pub allowed_proxy_cidrs: Vec, + pub reusable: Option, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct GeneratedCredential { + pub credential_id: String, + pub secret: String, + pub changed: bool, +} + +pub trait CredentialStorage: Send + Sync + 'static { + fn load(&self) -> anyhow::Result>; + fn store(&self, serialized_credentials: &str) -> anyhow::Result<()>; +} + +pub(crate) struct CredentialManager { + credentials: Mutex>, + storage: Option>, + storage_write: Mutex<()>, +} + +impl Default for CredentialManager { + fn default() -> Self { + Self::new() + } +} + +impl CredentialManager { + pub fn new() -> Self { + Self { + credentials: Mutex::new(HashMap::new()), + storage: None, + storage_write: Mutex::new(()), + } + } + + pub fn from_storage(storage: Arc) -> Self { + let credentials = match storage.load() { + Ok(Some(serialized)) => serde_json::from_str(&serialized).unwrap_or_else(|error| { + tracing::warn!(?error, "failed to parse stored credentials"); + HashMap::new() + }), + Ok(None) => HashMap::new(), + Err(error) => { + tracing::warn!(?error, "failed to load stored credentials"); + HashMap::new() + } + }; + Self { + credentials: Mutex::new(credentials), + storage: Some(storage), + storage_write: Mutex::new(()), + } + } + + pub fn with_entries(&self, f: impl FnOnce(&HashMap) -> R) -> R { + let credentials = self.credentials.lock().unwrap(); + f(&credentials) + } + + pub fn generate_credential_with_options( + &self, + groups: Vec, + allow_relay: bool, + allowed_proxy_cidrs: Vec, + ttl: Duration, + credential_id: Option, + reusable: bool, + ) -> GeneratedCredential { + self.remove_expired_credentials(); + self.generate_credential_with_options_after_cleanup( + groups, + allow_relay, + allowed_proxy_cidrs, + ttl, + credential_id, + reusable, + ) + } + + pub fn generate_credential_with_options_after_cleanup( + &self, + groups: Vec, + allow_relay: bool, + allowed_proxy_cidrs: Vec, + ttl: Duration, + credential_id: Option, + reusable: bool, + ) -> GeneratedCredential { + let generated = { + let mut credentials = self.credentials.lock().unwrap(); + let id = if let Some(id) = credential_id + .map(|x| x.trim().to_string()) + .filter(|x| !x.is_empty()) + { + if let Some(existing) = credentials.get(&id) + && !existing.secret.is_empty() + { + return GeneratedCredential { + credential_id: id, + secret: existing.secret.clone(), + changed: false, + }; + } + id + } else { + uuid::Uuid::new_v4().to_string() + }; + + let (entry, secret) = + Self::build_entry(groups, allow_relay, allowed_proxy_cidrs, reusable, ttl); + credentials.insert(id.clone(), entry); + GeneratedCredential { + credential_id: id, + secret, + changed: true, + } + }; + self.persist(); + generated + } + + fn build_entry( + groups: Vec, + allow_relay: bool, + allowed_proxy_cidrs: Vec, + reusable: bool, + ttl: Duration, + ) -> (CredentialEntry, String) { + let private = StaticSecret::random_from_rng(rand::rngs::OsRng); + let public = PublicKey::from(&private); + let pubkey = BASE64_STANDARD.encode(public.as_bytes()); + let secret = BASE64_STANDARD.encode(private.as_bytes()); + + let now = SystemTime::now() + .duration_since(UNIX_EPOCH) + .unwrap() + .as_secs() as i64; + let expiry_unix = now + ttl.as_secs() as i64; + + let entry = CredentialEntry { + pubkey, + secret: secret.clone(), + groups, + allow_relay, + allowed_proxy_cidrs, + reusable, + expiry_unix, + created_at_unix: now, + }; + (entry, secret) + } + + pub fn revoke_credential(&self, credential_id: &str) -> bool { + let removed = self + .credentials + .lock() + .unwrap() + .remove(credential_id) + .is_some(); + if removed { + self.persist(); + } + removed + } + + pub fn remove_expired_credentials(&self) -> bool { + self.remove_expired_credentials_at(current_unix_timestamp()) + } + + fn remove_expired_credentials_at(&self, now: i64) -> bool { + let mut credentials = self.credentials.lock().unwrap(); + let before = credentials.len(); + credentials.retain(|_, entry| entry.is_active_at(now)); + let changed = before != credentials.len(); + drop(credentials); + if changed { + self.persist(); + } + changed + } + + pub fn get_trusted_pubkeys(&self, network_secret: &str) -> Vec { + let now = current_unix_timestamp(); + + self.credentials + .lock() + .unwrap() + .values() + .filter(|entry| entry.is_active_at(now)) + .filter_map(|entry| { + entry.to_trusted_credential().map(|credential| { + TrustedCredentialPubkeyProof::new_signed(credential, network_secret) + }) + }) + .collect() + } + + pub fn is_pubkey_trusted(&self, pubkey: &[u8]) -> bool { + let now = current_unix_timestamp(); + + let encoded = BASE64_STANDARD.encode(pubkey); + self.credentials + .lock() + .unwrap() + .values() + .any(|entry| entry.pubkey == encoded && entry.is_active_at(now)) + } + + pub fn list_credentials(&self) -> Vec { + let now = current_unix_timestamp(); + + self.credentials + .lock() + .unwrap() + .iter() + .filter(|(_, entry)| entry.is_active_at(now)) + .map(|(id, entry)| entry.to_credential_info(id)) + .collect() + } + + fn decode_pubkey_b64(s: &str) -> Option> { + let decoded = BASE64_STANDARD.decode(s).ok()?; + if decoded.len() != 32 { + return None; + } + Some(decoded) + } + + fn persist(&self) { + let Some(storage) = &self.storage else { + return; + }; + let _storage_write = self.storage_write.lock().unwrap(); + let serialized = match self.with_entries(serde_json::to_string_pretty) { + Ok(serialized) => serialized, + Err(error) => { + tracing::warn!(?error, "failed to serialize credentials"); + return; + } + }; + if let Err(error) = storage.store(&serialized) { + tracing::warn!(?error, "failed to store credentials"); + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + + impl CredentialManager { + pub(crate) fn generate_credential( + &self, + groups: Vec, + allow_relay: bool, + allowed_proxy_cidrs: Vec, + ttl: Duration, + ) -> GeneratedCredential { + self.generate_credential_with_options( + groups, + allow_relay, + allowed_proxy_cidrs, + ttl, + None, + true, + ) + } + + fn generate_credential_with_id( + &self, + groups: Vec, + allow_relay: bool, + allowed_proxy_cidrs: Vec, + ttl: Duration, + credential_id: Option, + ) -> GeneratedCredential { + self.generate_credential_with_options( + groups, + allow_relay, + allowed_proxy_cidrs, + ttl, + credential_id, + true, + ) + } + } + + #[derive(Default)] + struct MemoryCredentialStorage { + serialized: Mutex>, + } + + impl CredentialStorage for MemoryCredentialStorage { + fn load(&self) -> anyhow::Result> { + Ok(self.serialized.lock().unwrap().clone()) + } + + fn store(&self, serialized_credentials: &str) -> anyhow::Result<()> { + *self.serialized.lock().unwrap() = Some(serialized_credentials.to_owned()); + Ok(()) + } + } + + #[test] + fn generate_and_revoke_credential() { + let mgr = CredentialManager::new(); + let generated = mgr.generate_credential( + vec!["guest".to_string()], + false, + vec![], + Duration::from_secs(3600), + ); + + assert!(!generated.credential_id.is_empty()); + assert!(!generated.secret.is_empty()); + assert!(generated.changed); + assert!(uuid::Uuid::parse_str(&generated.credential_id).is_ok()); + + let privkey_bytes: [u8; 32] = BASE64_STANDARD + .decode(&generated.secret) + .unwrap() + .try_into() + .unwrap(); + let private = StaticSecret::from(privkey_bytes); + let pubkey_bytes = PublicKey::from(&private).as_bytes().to_vec(); + assert!(mgr.is_pubkey_trusted(&pubkey_bytes)); + + let trusted = mgr.get_trusted_pubkeys("sec"); + assert_eq!(trusted.len(), 1); + assert_eq!( + trusted[0].credential.as_ref().unwrap().groups, + vec!["guest".to_string()] + ); + assert_eq!(trusted[0].credential.as_ref().unwrap().reusable, Some(true)); + + assert!(mgr.revoke_credential(&generated.credential_id)); + assert!(!mgr.is_pubkey_trusted(&pubkey_bytes)); + assert!(mgr.get_trusted_pubkeys("sec").is_empty()); + } + + #[test] + fn fixed_id_reuses_existing_secret() { + let mgr = CredentialManager::new(); + let fixed_id = "fixed-credential-id".to_string(); + let first = mgr.generate_credential_with_id( + vec!["group-a".to_string()], + false, + vec!["10.0.0.0/24".to_string()], + Duration::from_secs(3600), + Some(fixed_id.clone()), + ); + let second = mgr.generate_credential_with_id( + vec!["group-b".to_string()], + true, + vec!["192.168.0.0/16".to_string()], + Duration::from_secs(7200), + Some(fixed_id.clone()), + ); + + assert_eq!(first.credential_id, fixed_id); + assert_eq!(second.credential_id, fixed_id); + assert_eq!(first.secret, second.secret); + assert!(first.changed); + assert!(!second.changed); + + let list = mgr.list_credentials(); + assert_eq!(list.len(), 1); + assert_eq!(list[0].credential_id, fixed_id); + assert_eq!(list[0].groups, vec!["group-a".to_string()]); + assert!(!list[0].allow_relay); + assert_eq!(list[0].allowed_proxy_cidrs, vec!["10.0.0.0/24".to_string()]); + assert_eq!(list[0].reusable, Some(true)); + } + + #[test] + fn expired_credentials_are_filtered() { + let mgr = CredentialManager::new(); + mgr.generate_credential(vec![], false, vec![], Duration::from_secs(3600)); + mgr.generate_credential(vec![], false, vec![], Duration::from_secs(0)); + + assert_eq!(mgr.list_credentials().len(), 1); + assert!(mgr.remove_expired_credentials()); + assert_eq!(mgr.list_credentials().len(), 1); + } + + #[test] + fn injected_storage_loads_and_persists_mutations() { + let storage = Arc::new(MemoryCredentialStorage::default()); + let manager = CredentialManager::from_storage(storage.clone()); + let generated = + manager.generate_credential(vec![], false, vec![], Duration::from_secs(3600)); + + let reloaded = CredentialManager::from_storage(storage.clone()); + assert_eq!( + reloaded.list_credentials()[0].credential_id, + generated.credential_id + ); + + assert!(manager.revoke_credential(&generated.credential_id)); + let reloaded = CredentialManager::from_storage(storage); + assert!(reloaded.list_credentials().is_empty()); + } + + #[test] + fn malformed_storage_starts_with_empty_credentials() { + let storage = Arc::new(MemoryCredentialStorage { + serialized: Mutex::new(Some("not json".to_owned())), + }); + + let manager = CredentialManager::from_storage(storage); + + assert!(manager.list_credentials().is_empty()); + } +} diff --git a/easytier-core/src/peers/error.rs b/easytier-core/src/peers/error.rs new file mode 100644 index 00000000..8209e7f4 --- /dev/null +++ b/easytier-core/src/peers/error.rs @@ -0,0 +1,31 @@ +use crate::{config::PeerId, tunnel::TunnelError}; + +#[derive(Debug, thiserror::Error)] +pub enum Error { + #[error("wait response error: {0}")] + WaitRespError(String), + #[error("secret key error: {0}")] + SecretKeyError(String), + #[error("peer has no connection: {0}")] + PeerNoConnectionError(PeerId), + #[error("route error: {0:?}")] + RouteError(Option), + #[error("not found")] + NotFound, + #[error(transparent)] + Tunnel(#[from] TunnelError), + #[error(transparent)] + Other(#[from] anyhow::Error), +} + +impl From for Error { + fn from(value: snow::Error) -> Self { + Self::WaitRespError(value.to_string()) + } +} + +impl From for Error { + fn from(value: crate::foundation::time::error::Elapsed) -> Self { + Self::WaitRespError(value.to_string()) + } +} diff --git a/easytier/src/peers/foreign_network_client.rs b/easytier-core/src/peers/foreign_network/client.rs similarity index 81% rename from easytier/src/peers/foreign_network_client.rs rename to easytier-core/src/peers/foreign_network/client.rs index 4c8d6765..cf97dbf7 100644 --- a/easytier/src/peers/foreign_network_client.rs +++ b/easytier-core/src/peers/foreign_network/client.rs @@ -1,39 +1,33 @@ use std::sync::{Arc, Mutex}; -use crate::{ - common::{PeerId, error::Error, global_ctx::ArcGlobalCtx}, - tunnel::packet_def::ZCPacket, -}; use tokio_util::task::AbortOnDropHandle; -use super::{PacketRecvChan, peer_conn::PeerConn, peer_map::PeerMap, peer_rpc::PeerRpcManager}; +use crate::{config::PeerId, packet::ZCPacket}; + +use crate::peers::{ + PacketRecvChan, + conn::{peer_conn::PeerConn, peer_map::PeerMap}, + context::ArcPeerContext, + error::Error, +}; pub struct ForeignNetworkClient { - global_ctx: ArcGlobalCtx, - peer_rpc: Arc, - my_peer_id: PeerId, - peer_map: Arc, task: Mutex>>, } impl ForeignNetworkClient { - pub fn new( - global_ctx: ArcGlobalCtx, + pub(crate) fn new( + context: ArcPeerContext, packet_sender_to_mgr: PacketRecvChan, - peer_rpc: Arc, my_peer_id: PeerId, ) -> Self { let peer_map = Arc::new(PeerMap::new( packet_sender_to_mgr, - global_ctx.clone(), + context.clone(), my_peer_id, )); Self { - global_ctx, - peer_rpc, - my_peer_id, - peer_map, task: Mutex::new(None), } @@ -44,6 +38,10 @@ impl ForeignNetworkClient { self.peer_map.add_new_peer_conn(peer_conn).await } + pub fn is_client_url_alive(&self, url: &url::Url) -> bool { + self.peer_map.is_client_url_alive(url) + } + pub fn has_next_hop(&self, peer_id: PeerId) -> bool { self.get_next_hop(peer_id).is_some() } @@ -85,7 +83,8 @@ impl ForeignNetworkClient { let peer_map = Arc::downgrade(&self.peer_map); *self.task.lock().unwrap() = Some(AbortOnDropHandle::new(tokio::spawn(async move { loop { - tokio::time::sleep(tokio::time::Duration::from_secs(1)).await; + crate::foundation::time::sleep(crate::foundation::time::Duration::from_secs(1)) + .await; let Some(peer_map) = peer_map.upgrade() else { break; }; diff --git a/easytier-core/src/peers/foreign_network/mod.rs b/easytier-core/src/peers/foreign_network/mod.rs new file mode 100644 index 00000000..8983371c --- /dev/null +++ b/easytier-core/src/peers/foreign_network/mod.rs @@ -0,0 +1,1693 @@ +//! Foreign network management: the manager owning foreign-network entries +//! and the client connecting to foreign networks. + +pub(crate) mod client; + +use std::sync::atomic::{AtomicBool, Ordering}; +use std::{ + future, + sync::{Arc, Weak}, + time::SystemTime, +}; + +use dashmap::{DashMap, DashSet}; +use easytier_proto::common::FlagsInConfig; +use guarden::{Guard, defer}; +use tokio::sync::{ + Mutex, RwLock, RwLockReadGuard, + mpsc::{self, UnboundedReceiver, UnboundedSender}, +}; +use tokio::task::JoinSet; + +use crate::{ + config::peers::{PeerRuntimeConfig, PeerRuntimeSnapshot}, + config::runtime::{CoreRuntimeConfig, CoreRuntimeConfigStore}, + config::{CoreConfig, NodeConfig, PeerId}, + foundation::{ + stats::{CounterHandle, LabelSet, LabelType, MetricName, StatsManager}, + task::reap_joinset_background, + token_bucket::ArcByteLimiter, + }, + packet::{PacketType, ZCPacket}, + peers::peer_center::instance::{PeerCenterInstance, PeerCenterPeerManagerTrait}, + peers::{PacketRecvChan, PacketRecvChanReceiver, recv_packet_from_chan}, + proto::core_peer::peer::{PeerConnInfo, Route as CoreRoute}, + socket::SocketContext, +}; + +use super::{ + conn::{ + peer_conn::{PeerConn, PeerConnId}, + peer_map::{PeerMap, direct_peer_info}, + peer_session::PeerSessionStore, + }, + context::{ + ArcPeerContext, CorePeerContext, CorePeerContextAdapters, NetworkIdentity, PeerContext, + PeerContextEvent, PeerStunInfoSource, TrustedKeySource, + }, + error::Error, + peer_rpc::{PeerRpcManager, PeerRpcManagerTransport}, + public_ipv6::{DisabledPublicIpv6Runtime, PublicIpv6Runtime}, + relay_peer_map::RelayPeerMap, + route::{NextHopPolicy, Route, RouteInterface, peer_ospf_route::PeerRoute}, + traffic_metrics::{ + TrafficKind, TrafficMetricRecorder, is_relay_data_packet_type, traffic_kind, + }, + util::shrink_dashmap, + whitelist::check_network_in_relay_whitelist, +}; +use crate::proto::peer_rpc::{PeerIdentityType, PeerInfoForGlobalMap}; + +pub const PUBLIC_SERVER_HOSTNAME_PREFIX: &str = "PublicServer_"; + +pub(crate) fn desired_foreign_avoid_relay_data( + parent_context: &ArcPeerContext, + relay_data: bool, +) -> bool { + !relay_data || parent_context.feature_flags().avoid_relay_data +} + +pub(crate) fn sync_foreign_avoid_relay_data( + parent_context: &ArcPeerContext, + foreign_context: &ArcPeerContext, + relay_data: bool, +) -> bool { + let desired = desired_foreign_avoid_relay_data(parent_context, relay_data); + if foreign_context.feature_flags().avoid_relay_data == desired { + return false; + } + foreign_context.set_avoid_relay_data_preference(desired) +} + +fn build_foreign_peer_context( + network: &NetworkIdentity, + parent_context: &Arc, + relay_data: bool, + mut flags: FlagsInConfig, +) -> Arc { + let parent_context_dyn: ArcPeerContext = parent_context.clone(); + let parent_flags = parent_context_dyn.flags(); + flags.disable_relay_kcp = !parent_flags.enable_relay_foreign_network_kcp; + flags.disable_relay_quic = !parent_flags.enable_relay_foreign_network_quic; + flags.socket_mark = parent_flags.socket_mark; + + let mut feature_flags = parent_context_dyn.feature_flags(); + feature_flags.is_public_server = true; + feature_flags.avoid_relay_data = + desired_foreign_avoid_relay_data(&parent_context_dyn, relay_data); + + let instance_id = uuid::Uuid::new_v4(); + let runtime = PeerRuntimeConfig { + core: CoreConfig { + node: NodeConfig { + instance_id: Some(*instance_id.as_bytes()), + hostname: Some(format!( + "{PUBLIC_SERVER_HOSTNAME_PREFIX}{}", + parent_context_dyn.hostname() + )), + network_name: network.network_name.clone(), + ..Default::default() + }, + ..Default::default() + }, + network_identity: network.clone(), + stun_info: parent_context_dyn.stun_info(), + feature_flags, + secure_mode: parent_context_dyn.secure_mode(), + host_routing: parent_context_dyn.host_routing_policy(), + }; + let mut snapshot = PeerRuntimeSnapshot::new(runtime, flags); + snapshot.easytier_version = parent_context_dyn.easytier_version(); + snapshot.ospf_update_my_foreign_network_interval_sec = + parent_context_dyn.ospf_update_my_foreign_network_interval_sec(); + snapshot.max_direct_conns_per_peer_in_foreign_network = + parent_context_dyn.max_direct_conns_per_peer_in_foreign_network(); + snapshot.hmac_secret_digest = parent_context_dyn.hmac_secret_digest(); + + Arc::new(CorePeerContext::new_foreign( + CoreRuntimeConfigStore::new(CoreRuntimeConfig::default(), Arc::new(snapshot)), + CorePeerContextAdapters { + stun_info_source: Some(Arc::new(ParentStunInfoSource(parent_context_dyn))), + events: Arc::new(()), + credential_storage: None, + }, + parent_context, + )) +} + +struct ParentStunInfoSource(ArcPeerContext); + +impl PeerStunInfoSource for ParentStunInfoSource { + fn stun_info(&self) -> crate::proto::common::StunInfo { + self.0.stun_info() + } +} + +/// Adapts a foreign network's peer map and RPC manager to the peer-center +/// trait so each foreign network runs its own peer-center instance. +pub struct PeerMapWithPeerRpcManager { + pub peer_map: Arc, + pub rpc_mgr: Arc, + pub network_name: String, +} + +#[async_trait::async_trait] +impl PeerCenterPeerManagerTrait for PeerMapWithPeerRpcManager { + async fn list_peers(&self) -> PeerInfoForGlobalMap { + // TODO: currently latency between public server cannot be calculated because one public-server pair + // has no connection between them. (hard to get latency from peer manager because it's hard to transform the peer id) + // but it's fine because we don't want too much traffic between public servers. + direct_peer_info(std::slice::from_ref(&self.peer_map)).await + } + + fn my_peer_id(&self) -> PeerId { + self.peer_map.my_peer_id() + } + + fn network_name(&self) -> String { + self.network_name.clone() + } + + fn get_rpc_mgr(&self) -> Weak { + Arc::downgrade(&self.rpc_mgr) + } + + async fn list_routes(&self) -> Vec { + self.peer_map.list_route_infos().await + } +} + +#[derive(Clone, Debug, Default)] +pub(crate) struct ForeignNetworkRouteInfo { + pub network_name: String, + pub peer_ids: Vec, + pub network_secret_digest: Vec, + pub my_peer_id_for_this_network: PeerId, +} + +struct ForeignNetworkRouteInterface { + my_peer_id: PeerId, + peer_map: Weak, + network_identity: NetworkIdentity, + global_peer_map: Weak, +} + +#[async_trait::async_trait] +impl RouteInterface for ForeignNetworkRouteInterface { + async fn list_peers(&self) -> Vec { + let Some(peer_map) = self.peer_map.upgrade() else { + return vec![]; + }; + + let mut global = if let Some(global_peer_map) = self.global_peer_map.upgrade() { + global_peer_map + .list_peers_own_foreign_network(&self.network_identity) + .await + } else { + vec![] + }; + let local = peer_map.list_peers_with_conn().await; + global.extend(local.iter().cloned()); + global + .into_iter() + .filter(|peer_id| *peer_id != self.my_peer_id) + .collect() + } + + fn my_peer_id(&self) -> PeerId { + self.my_peer_id + } + + fn need_periodic_requery_peers(&self) -> bool { + true + } + + async fn get_peer_identity_type(&self, peer_id: PeerId) -> Option { + let peer_map = self.peer_map.upgrade()?; + peer_map.get_peer_identity_type(peer_id) + } + + async fn get_peer_public_key(&self, peer_id: PeerId) -> Option> { + let peer_map = self.peer_map.upgrade()?; + peer_map.get_peer_public_key(peer_id) + } + + async fn close_peer(&self, peer_id: PeerId) { + if let Some(peer_map) = self.peer_map.upgrade() { + let _ = peer_map.close_peer(peer_id).await; + } + } +} + +struct RpcTransport { + my_peer_id: PeerId, + peer_map: Weak, + + packet_recv: Mutex>, +} + +#[async_trait::async_trait] +impl PeerRpcManagerTransport for RpcTransport { + fn my_peer_id(&self) -> PeerId { + self.my_peer_id + } + + async fn send(&self, msg: ZCPacket, dst_peer_id: PeerId) -> anyhow::Result<()> { + tracing::debug!( + "foreign network manager send rpc to peer: {:?}", + dst_peer_id + ); + let peer_map = self + .peer_map + .upgrade() + .ok_or(anyhow::anyhow!("peer map is gone"))?; + + // send to ourselves so we can handle it in forward logic. + peer_map.send_msg_directly(msg, self.my_peer_id).await?; + Ok(()) + } + + async fn recv(&self) -> anyhow::Result { + if let Some(packet) = self.packet_recv.lock().await.recv().await { + tracing::trace!("recv rpc packet in foreign network manager rpc transport"); + Ok(packet) + } else { + Err(anyhow::anyhow!("unknown data store error")) + } + } +} + +impl Drop for RpcTransport { + fn drop(&mut self) { + tracing::debug!( + "drop rpc transport for foreign network manager, my_peer_id: {:?}", + self.my_peer_id + ); + } +} + +#[derive(Clone, Debug)] +pub struct ForeignNetworkTrustedKeyInfo { + pub pubkey: Vec, + pub source: TrustedKeySource, + pub expiry_unix: Option, +} + +#[derive(Clone, Debug, Default)] +pub struct ForeignNetworkPeerInfo { + pub peer_id: PeerId, + pub conns: Vec, +} + +#[derive(Clone, Debug, Default)] +pub struct ForeignNetworkEntryInfo { + pub network_secret_digest: Vec, + pub my_peer_id_for_this_network: PeerId, + pub peers: Vec, + pub trusted_keys: Vec, +} + +#[auto_impl::auto_impl(&, Arc)] +pub(crate) trait ForeignNetworkRpcRegistrar: Send + Sync + 'static { + fn register_peer_rpc_services( + &self, + _peer_rpc: &Arc, + _network_name: &str, + _socket_context: SocketContext, + ) { + } +} + +impl ForeignNetworkRpcRegistrar for () {} + +async fn abort_and_join_persistent_tasks(tasks: &std::sync::Mutex>) { + tasks.lock().unwrap().abort_all(); + future::poll_fn(|cx| { + let mut tasks = tasks.lock().unwrap(); + loop { + match tasks.poll_join_next(cx) { + std::task::Poll::Ready(Some(_)) => continue, + std::task::Poll::Ready(None) => return std::task::Poll::Ready(()), + std::task::Poll::Pending => return std::task::Poll::Pending, + } + } + }) + .await; +} + +struct ForeignNetworkEntry { + my_peer_id: PeerId, + + parent_context: ArcPeerContext, + peer_context: Arc, + network: NetworkIdentity, + peer_map: Arc, + relay_peer_map: Arc, + relay_data: bool, + pm_packet_sender: Mutex>, + + peer_rpc: Arc, + rpc_sender: UnboundedSender, + + packet_recv: Mutex>, + + bps_limiter: Option, + + peer_center: Arc, + + traffic_metrics: Arc, + event_handler_started: AtomicBool, + + tasks: Mutex>, + + lock: Mutex<()>, +} + +impl ForeignNetworkEntry { + #[allow(clippy::too_many_arguments)] + fn new( + network: NetworkIdentity, + my_peer_id: PeerId, + rpc_registrar: Arc, + parent_context: Arc, + foreign_context_default_flags: FlagsInConfig, + relay_data: bool, + peer_session_store: Arc, + pm_packet_sender: PacketRecvChan, + ) -> Self { + let parent_context_dyn: ArcPeerContext = parent_context.clone(); + let peer_context = build_foreign_peer_context( + &network, + &parent_context, + relay_data, + foreign_context_default_flags, + ); + let socket_mark = peer_context.flags().socket_mark; + let stats_mgr = peer_context.stats_manager(); + let network_name = network.network_name.clone(); + + let (packet_sender, packet_recv) = super::create_packet_recv_chan(); + + let peer_map = Arc::new(PeerMap::new( + packet_sender, + peer_context.clone(), + my_peer_id, + )); + let traffic_metrics = Arc::new(TrafficMetricRecorder::new( + my_peer_id, + Arc::new(super::traffic_metrics::LogicalTrafficMetrics::new( + stats_mgr.clone(), + network_name.clone(), + MetricName::TrafficBytesTx, + MetricName::TrafficPacketsTx, + MetricName::TrafficBytesTxByInstance, + MetricName::TrafficPacketsTxByInstance, + super::traffic_metrics::InstanceLabelKind::To, + )), + Arc::new(super::traffic_metrics::LogicalTrafficMetrics::new( + stats_mgr.clone(), + network_name.clone(), + MetricName::TrafficControlBytesTx, + MetricName::TrafficControlPacketsTx, + MetricName::TrafficControlBytesTxByInstance, + MetricName::TrafficControlPacketsTxByInstance, + super::traffic_metrics::InstanceLabelKind::To, + )), + Arc::new(super::traffic_metrics::LogicalTrafficMetrics::new( + stats_mgr.clone(), + network_name.clone(), + MetricName::TrafficBytesRx, + MetricName::TrafficPacketsRx, + MetricName::TrafficBytesRxByInstance, + MetricName::TrafficPacketsRxByInstance, + super::traffic_metrics::InstanceLabelKind::From, + )), + Arc::new(super::traffic_metrics::LogicalTrafficMetrics::new( + stats_mgr.clone(), + network_name.clone(), + MetricName::TrafficControlBytesRx, + MetricName::TrafficControlPacketsRx, + MetricName::TrafficControlBytesRxByInstance, + MetricName::TrafficControlPacketsRxByInstance, + super::traffic_metrics::InstanceLabelKind::From, + )), + { + let peer_map = Arc::downgrade(&peer_map); + move |peer_id| { + let peer_map = peer_map.clone(); + async move { + let peer_map = peer_map.upgrade()?; + peer_map + .get_route_peer_info(peer_id) + .await + .as_ref() + .and_then(super::traffic_metrics::route_peer_info_instance_id) + } + } + }, + )); + let relay_peer_map = super::relay_peer_map::new_relay_peer_map( + peer_map.clone(), + None, + peer_context.clone(), + my_peer_id, + peer_session_store.clone(), + ); + + let (rpc_transport_sender, rpc_packet_recv) = mpsc::unbounded_channel(); + let peer_rpc = Arc::new(PeerRpcManager::new(RpcTransport { + my_peer_id, + peer_map: Arc::downgrade(&peer_map), + packet_recv: Mutex::new(rpc_packet_recv), + })); + + rpc_registrar.register_peer_rpc_services( + &peer_rpc, + &network.network_name, + SocketContext::default().with_socket_mark(socket_mark), + ); + + let bps_limiter = parent_context.foreign_forward_limiter(&network.network_name); + + let peer_center = Arc::new(PeerCenterInstance::new(Arc::new( + PeerMapWithPeerRpcManager { + peer_map: peer_map.clone(), + rpc_mgr: peer_rpc.clone(), + network_name: peer_context.network_name(), + }, + ))); + + Self { + my_peer_id, + + parent_context: parent_context_dyn, + peer_context, + network, + peer_map, + relay_peer_map, + relay_data, + pm_packet_sender: Mutex::new(Some(pm_packet_sender)), + + peer_rpc, + rpc_sender: rpc_transport_sender, + + packet_recv: Mutex::new(Some(packet_recv)), + + bps_limiter, + + traffic_metrics, + event_handler_started: AtomicBool::new(false), + + tasks: Mutex::new(JoinSet::new()), + + peer_center, + + lock: Mutex::new(()), + } + } + + async fn prepare_route(&self, global_peer_map: Weak) { + let public_ipv6_runtime: Arc = + Arc::new(DisabledPublicIpv6Runtime::new( + self.peer_context.instance_id(), + self.network.network_name.clone(), + )); + let route = PeerRoute::new( + self.my_peer_id, + self.peer_context.clone(), + public_ipv6_runtime, + self.peer_rpc.clone(), + ); + route + .open(Box::new(ForeignNetworkRouteInterface { + my_peer_id: self.my_peer_id, + peer_map: Arc::downgrade(&self.peer_map), + network_identity: self.network.clone(), + global_peer_map, + })) + .await + .unwrap(); + + route + .set_route_cost_fn(self.peer_center.get_cost_calculator()) + .await; + + self.peer_map.add_route(route).await; + } + + async fn start_packet_recv(&self) { + let packet_recv = self.packet_recv.lock().await.take().unwrap(); + let pm_sender = self.pm_packet_sender.lock().await.take().unwrap(); + let router = ForeignNetworkPacketRouter::new( + self.my_peer_id, + packet_recv, + self.rpc_sender.clone(), + self.peer_map.clone(), + self.relay_peer_map.clone(), + self.traffic_metrics.clone(), + self.parent_context.clone(), + self.relay_data, + pm_sender, + self.network.network_name.clone(), + self.bps_limiter.clone(), + self.peer_context.stats_manager(), + ); + + self.tasks.lock().await.spawn(router.run()); + } + + async fn run_relay_session_gc_routine(&self) { + let relay_peer_map = self.relay_peer_map.clone(); + self.tasks.lock().await.spawn(async move { + loop { + relay_peer_map.evict_idle_sessions(std::time::Duration::from_secs(60)); + crate::foundation::time::sleep(std::time::Duration::from_secs(30)).await; + } + }); + } + + async fn run_parent_feature_flag_sync_routine(&self) { + let parent_context = self.parent_context.clone(); + let foreign_peer_context: ArcPeerContext = self.peer_context.clone(); + let relay_data = self.relay_data; + let runtime_changes = parent_context.subscribe_runtime_changes(); + sync_foreign_avoid_relay_data(&parent_context, &foreign_peer_context, relay_data); + let Some(mut runtime_changes) = runtime_changes else { + return; + }; + self.tasks.lock().await.spawn(async move { + loop { + if runtime_changes.changed().await.is_err() { + break; + } + sync_foreign_avoid_relay_data(&parent_context, &foreign_peer_context, relay_data); + } + }); + } + + async fn prepare(&self, global_peer_map: Weak) { + self.prepare_route(global_peer_map).await; + self.start_packet_recv().await; + self.run_relay_session_gc_routine().await; + self.run_parent_feature_flag_sync_routine().await; + self.peer_rpc.run(); + self.peer_center.init().await; + } + + async fn stop(&self) { + let mut tasks = self.tasks.lock().await; + tasks.abort_all(); + while tasks.join_next().await.is_some() {} + drop(tasks); + self.peer_center.stop().await; + self.peer_rpc.stop().await; + self.peer_map.clear_resources().await; + self.peer_context.stop().await; + } +} + +impl Drop for ForeignNetworkEntry { + fn drop(&mut self) { + self.peer_rpc + .rpc_server() + .registry() + .unregister_by_domain(&self.network.network_name); + self.peer_context + .remove_trusted_keys(&self.network.network_name); + + tracing::debug!(self.my_peer_id, ?self.network, "drop foreign network entry"); + } +} + +struct ForeignNetworkManagerData { + network_peer_maps: DashMap>, + peer_network_map: DashMap>, + network_peer_last_update: DashMap, + global_peer_map: Weak, + lock: std::sync::Mutex<()>, +} + +impl ForeignNetworkManagerData { + fn get_peer_network(&self, peer_id: PeerId) -> Option> { + self.peer_network_map.get(&peer_id).map(|v| v.clone()) + } + + fn get_network_entry(&self, network_name: &str) -> Option> { + self.network_peer_maps.get(network_name).map(|v| v.clone()) + } + + fn remove_peer(&self, peer_id: PeerId, network_name: &String) { + let _l = self.lock.lock().unwrap(); + self.peer_network_map.remove_if(&peer_id, |_, v| { + let _ = v.remove(network_name); + v.is_empty() + }); + if self + .network_peer_maps + .remove_if(network_name, |_, v| v.peer_map.is_empty()) + .is_some() + { + self.network_peer_last_update.remove(network_name); + } + shrink_dashmap(&self.peer_network_map, None); + shrink_dashmap(&self.network_peer_maps, None); + shrink_dashmap(&self.network_peer_last_update, None); + } + + async fn clear_no_conn_peer(&self, network_name: &String) { + let Some(peer_map) = self + .network_peer_maps + .get(network_name) + .map(|v| v.peer_map.clone()) + else { + return; + }; + peer_map.clean_peer_without_conn().await; + } + + fn remove_network_if_current( + &self, + network_name: &String, + expected_entry: &Weak, + ) { + let _l = self.lock.lock().unwrap(); + let Some(expected_entry) = expected_entry.upgrade() else { + return; + }; + let old = self + .network_peer_maps + .remove_if(network_name, |_, entry| Arc::ptr_eq(entry, &expected_entry)); + let Some((_, old)) = old else { + return; + }; + + old.traffic_metrics.clear_peer_cache(); + let to_remove_peers = old.peer_map.list_peers(); + for p in to_remove_peers { + self.peer_network_map.remove_if(&p, |_, v| { + v.remove(network_name); + v.is_empty() + }); + } + self.network_peer_last_update.remove(network_name); + shrink_dashmap(&self.peer_network_map, None); + shrink_dashmap(&self.network_peer_maps, None); + shrink_dashmap(&self.network_peer_last_update, None); + } + + #[allow(clippy::too_many_arguments)] + async fn get_or_insert_entry( + &self, + network_identity: &NetworkIdentity, + my_peer_id: PeerId, + dst_peer_id: PeerId, + relay_data: bool, + rpc_registrar: Arc, + parent_context: Arc, + foreign_context_default_flags: FlagsInConfig, + peer_session_store: Arc, + pm_packet_sender: &PacketRecvChan, + ) -> (Arc, bool) { + let mut new_added = false; + + let l = self.lock.lock().unwrap(); + let entry = self + .network_peer_maps + .entry(network_identity.network_name.clone()) + .or_insert_with(|| { + new_added = true; + Arc::new(ForeignNetworkEntry::new( + network_identity.clone(), + my_peer_id, + rpc_registrar, + parent_context, + foreign_context_default_flags, + relay_data, + peer_session_store, + pm_packet_sender.clone(), + )) + }) + .clone(); + + self.peer_network_map + .entry(dst_peer_id) + .or_default() + .insert(network_identity.network_name.clone()); + + self.network_peer_last_update + .insert(network_identity.network_name.clone(), SystemTime::now()); + + drop(l); + + if new_added { + entry.prepare(self.global_peer_map.clone()).await; + } + + (entry, new_added) + } +} + +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +enum ForeignNetworkManagerState { + Running, + Stopping, + Stopped, +} + +pub(crate) struct ForeignNetworkManager { + rpc_registrar: Arc, + parent_context: Arc, + foreign_context_default_flags: FlagsInConfig, + peer_session_store: Arc, + packet_sender_to_mgr: PacketRecvChan, + + data: Arc, + + tasks: Arc>>, + task_reaper: Mutex>>, + lifecycle: RwLock, +} + +impl ForeignNetworkManager { + fn network_secret_digest_is_empty(network: &NetworkIdentity) -> bool { + network + .network_secret_digest + .as_ref() + .is_none_or(|d| d.iter().all(|b| *b == 0)) + } + + fn should_reject_credential_trust_path(identity_type: PeerIdentityType) -> bool { + matches!(identity_type, PeerIdentityType::Admin) + } + + fn credential_pubkey_is_trusted( + entry: &ForeignNetworkEntry, + remote_static_pubkey: &[u8], + ) -> bool { + remote_static_pubkey.len() == 32 + && entry.peer_context.is_pubkey_trusted_with_source( + remote_static_pubkey, + &entry.network.network_name, + TrustedKeySource::OspfCredential, + ) + } + + pub fn new( + rpc_registrar: Arc, + parent_context: Arc, + foreign_context_default_flags: FlagsInConfig, + peer_session_store: Arc, + packet_sender_to_mgr: PacketRecvChan, + global_peer_map: Weak, + ) -> Self { + let data = Arc::new(ForeignNetworkManagerData { + network_peer_maps: DashMap::new(), + peer_network_map: DashMap::new(), + network_peer_last_update: DashMap::new(), + global_peer_map, + lock: std::sync::Mutex::new(()), + }); + + let tasks = Arc::new(std::sync::Mutex::new(JoinSet::new())); + let task_reaper = tokio::spawn(reap_joinset_background( + tasks.clone(), + "ForeignNetworkManager", + )); + + Self { + rpc_registrar, + parent_context, + foreign_context_default_flags, + peer_session_store, + packet_sender_to_mgr, + data, + tasks, + task_reaper: Mutex::new(Some(task_reaper)), + lifecycle: RwLock::new(ForeignNetworkManagerState::Running), + } + } + + async fn admission_guard( + &self, + ) -> Result, Error> { + let guard = self.lifecycle.read().await; + if *guard != ForeignNetworkManagerState::Running { + return Err(anyhow::anyhow!("foreign network manager is stopping").into()); + } + Ok(guard) + } + + pub async fn stop(&self) { + let mut lifecycle = self.lifecycle.write().await; + if *lifecycle == ForeignNetworkManagerState::Stopped { + return; + } + // Set this before the first await so cancellation permanently closes + // admission. A later stop call can safely resume the idempotent + // teardown while admissions remain rejected. + *lifecycle = ForeignNetworkManagerState::Stopping; + + let mut reaper = self.task_reaper.lock().await; + if let Some(reaper) = reaper.as_mut() { + reaper.abort(); + let _ = reaper.await; + } + reaper.take(); + drop(reaper); + abort_and_join_persistent_tasks(&self.tasks).await; + + let entries = self + .data + .network_peer_maps + .iter() + .map(|entry| entry.value().clone()) + .collect::>(); + for entry in entries { + entry.stop().await; + } + // Keep entries discoverable until every nested graph is stopped. If + // this future is cancelled, the next stop call can snapshot and retry + // the remaining idempotent teardown instead of losing ownership. + self.data.peer_network_map.clear(); + self.data.network_peer_maps.clear(); + self.data.network_peer_last_update.clear(); + *lifecycle = ForeignNetworkManagerState::Stopped; + } + + pub fn get_network_peer_id(&self, network_name: &str) -> Option { + self.data + .network_peer_maps + .get(network_name) + .map(|v| v.my_peer_id) + } + + pub fn is_existing_credential_pubkey_trusted( + &self, + network_name: &str, + remote_static_pubkey: &[u8], + ) -> bool { + self.data + .get_network_entry(network_name) + .is_some_and(|entry| Self::credential_pubkey_is_trusted(&entry, remote_static_pubkey)) + } + + pub async fn add_peer_conn(&self, peer_conn: PeerConn) -> Result<(), Error> { + let _admission = self.admission_guard().await?; + let conn_info = peer_conn.get_conn_info(); + let peer_network = peer_conn.get_network_identity(); + tracing::info!(peer_conn = ?conn_info, network = ?peer_network, "add new peer conn in foreign network manager"); + + let parent_flags = self.parent_context.flags(); + let ret = check_network_in_relay_whitelist( + &parent_flags.relay_network_whitelist, + &peer_network.network_name, + ); + if ret.is_err() && !parent_flags.relay_all_peer_rpc { + return ret.map_err(Error::Other); + } + + let peer_digest_empty = Self::network_secret_digest_is_empty(&peer_network); + if peer_digest_empty + && self + .data + .get_network_entry(&peer_network.network_name) + .is_none() + { + return Err(anyhow::anyhow!( + "foreign network {} is not established by a secret-verified peer yet", + peer_network.network_name + ) + .into()); + } + + let (entry, new_added) = self + .data + .get_or_insert_entry( + &peer_network, + peer_conn.get_my_peer_id(), + peer_conn.get_peer_id(), + ret.is_ok(), + self.rpc_registrar.clone(), + self.parent_context.clone(), + self.foreign_context_default_flags.clone(), + self.peer_session_store.clone(), + &self.packet_sender_to_mgr, + ) + .await; + + defer!(rollback_new_entry => sync [ + data = self.data.clone(), + network_name = entry.network.network_name.clone(), + peer_id = peer_conn.get_peer_id(), + should_rollback = new_added + ] { + if should_rollback { + tracing::warn!( + %network_name, + "rollback newly added foreign network entry after add_peer_conn returned error" + ); + data.remove_peer(peer_id, &network_name); + } + }); + + self.ensure_event_handler_started(&entry); + + let same_identity = entry.network == peer_network; + let peer_identity_type = peer_conn.get_peer_identity_type(); + let credential_peer_trusted = peer_digest_empty + && Self::credential_pubkey_is_trusted(&entry, &conn_info.noise_remote_static_pubkey); + let credential_identity_mismatch = credential_peer_trusted + && Self::should_reject_credential_trust_path(peer_identity_type); + + let _g = entry.lock.lock().await; + + if (!(same_identity || credential_peer_trusted)) + || credential_identity_mismatch + || entry.my_peer_id != peer_conn.get_my_peer_id() + { + let err = if entry.my_peer_id != peer_conn.get_my_peer_id() { + anyhow::anyhow!( + "my peer id not match. exp: {:?} real: {:?}, need retry connect", + entry.my_peer_id, + peer_conn.get_my_peer_id() + ) + } else if credential_identity_mismatch { + anyhow::anyhow!( + "credential-trusted foreign peer has invalid identity type: {:?}", + peer_identity_type + ) + } else { + anyhow::anyhow!( + "foreign peer identity not trusted. exp: {:?} real: {:?}, remote_pubkey_len: {}, credential_trusted: {}", + entry.network, + peer_network, + conn_info.noise_remote_static_pubkey.len(), + credential_peer_trusted, + ) + }; + tracing::error!(?err, "foreign network entry not match, disconnect peer"); + return Err(err.into()); + } + + if !new_added && let Some(peer) = entry.peer_map.get_peer_by_id(peer_conn.get_peer_id()) { + let direct_conns_len = peer.get_directly_connections().len(); + let max_count = self + .parent_context + .max_direct_conns_per_peer_in_foreign_network(); + if direct_conns_len >= max_count { + return Err(anyhow::anyhow!( + "too many direct conns, cur: {}, max: {}", + direct_conns_len, + max_count + ) + .into()); + } + } + + entry.peer_map.add_new_peer_conn(peer_conn).await?; + let _ = rollback_new_entry.defuse(); + Ok(()) + } + + fn ensure_event_handler_started(&self, entry: &Arc) { + if entry.event_handler_started.swap(true, Ordering::AcqRel) { + return; + } + + let Some(mut s) = entry.peer_context.subscribe_peer_events() else { + return; + }; + let data = self.data.clone(); + let network_name = entry.network.network_name.clone(); + let entry_for_cleanup = Arc::downgrade(entry); + let traffic_metrics = Arc::downgrade(&entry.traffic_metrics); + self.tasks.lock().unwrap().spawn(async move { + while let Ok(e) = s.recv().await { + match &e { + PeerContextEvent::PeerRemoved(peer_id) => { + tracing::info!(?e, "remove peer from foreign network manager"); + if let Some(traffic_metrics) = traffic_metrics.upgrade() { + traffic_metrics.remove_peer(*peer_id); + } + data.network_peer_last_update + .insert(network_name.clone(), SystemTime::now()); + data.remove_peer(*peer_id, &network_name); + } + PeerContextEvent::PeerConnRemoved => { + tracing::info!(?e, "clear no conn peer from foreign network manager"); + data.clear_no_conn_peer(&network_name).await; + } + PeerContextEvent::PeerAdded(_) => { + tracing::info!(?e, "add peer to foreign network manager"); + data.network_peer_last_update + .insert(network_name.clone(), SystemTime::now()); + } + _ => continue, + } + } + tracing::error!("global event handler at foreign network manager exit"); + if let Some(traffic_metrics) = traffic_metrics.upgrade() { + traffic_metrics.clear_peer_cache(); + } + data.remove_network_if_current(&network_name, &entry_for_cleanup); + }); + } + + pub async fn list_foreign_network_infos( + &self, + include_trusted_keys: bool, + ) -> std::collections::HashMap { + let mut ret = std::collections::HashMap::new(); + let networks = self + .data + .network_peer_maps + .iter() + .map(|v| v.key().clone()) + .collect::>(); + + for network_name in networks { + let Some(item) = self + .data + .network_peer_maps + .get(&network_name) + .map(|v| v.clone()) + else { + continue; + }; + + let mut entry = ForeignNetworkEntryInfo { + network_secret_digest: item + .network + .network_secret_digest + .unwrap_or_default() + .to_vec(), + my_peer_id_for_this_network: item.my_peer_id, + peers: Default::default(), + trusted_keys: if include_trusted_keys { + item.peer_context + .list_trusted_keys(&item.network.network_name) + .into_iter() + .map(|(pubkey, metadata)| ForeignNetworkTrustedKeyInfo { + pubkey, + source: metadata.source, + expiry_unix: metadata.expiry_unix, + }) + .collect() + } else { + Default::default() + }, + }; + for peer in item.peer_map.list_peers() { + let peer_info = ForeignNetworkPeerInfo { + peer_id: peer, + conns: item.peer_map.list_peer_conns(peer).await.unwrap_or(vec![]), + }; + entry.peers.push(peer_info); + } + + ret.insert(network_name, entry); + } + ret + } + + pub async fn list_foreign_network_route_infos(&self) -> Vec { + self.list_foreign_network_infos(false) + .await + .into_iter() + .map(|(network_name, info)| ForeignNetworkRouteInfo { + network_name, + peer_ids: info.peers.into_iter().map(|peer| peer.peer_id).collect(), + network_secret_digest: info.network_secret_digest, + my_peer_id_for_this_network: info.my_peer_id_for_this_network, + }) + .collect() + } + + pub fn get_foreign_network_last_update(&self, network_name: &str) -> Option { + self.data + .network_peer_last_update + .get(network_name) + .map(|v| *v) + } + + pub async fn forward_foreign_network_packet( + &self, + network_name: &str, + dst_peer_id: PeerId, + msg: ZCPacket, + ) -> Result<(), Error> { + if let Some(entry) = self.data.get_network_entry(network_name) { + let packet_type = msg + .peer_manager_header() + .map(|hdr| hdr.packet_type) + .unwrap_or(0); + let msg_len = msg.buf_len() as u64; + let send_result = entry + .peer_map + .send_msg(msg, dst_peer_id, NextHopPolicy::LeastHop) + .await; + if send_result.is_ok() { + entry + .traffic_metrics + .record_tx(dst_peer_id, packet_type, msg_len) + .await; + } + send_result + } else { + Err(Error::RouteError(Some("network not found".to_string()))) + } + } + + pub async fn close_peer_conn( + &self, + peer_id: PeerId, + conn_id: &PeerConnId, + ) -> Result<(), Error> { + let network_names = self.data.get_peer_network(peer_id).unwrap_or_default(); + for network_name in network_names { + if let Some(entry) = self.data.get_network_entry(&network_name) { + let ret = entry.peer_map.close_peer_conn(peer_id, conn_id).await; + if ret.is_ok() || !matches!(ret.as_ref().unwrap_err(), Error::NotFound) { + return ret; + } + } + } + Err(Error::NotFound) + } +} + +impl Drop for ForeignNetworkManager { + fn drop(&mut self) { + if let Ok(mut reaper) = self.task_reaper.try_lock() + && let Some(reaper) = reaper.take() + { + reaper.abort(); + } + self.tasks.lock().unwrap().abort_all(); + self.data.peer_network_map.clear(); + self.data.network_peer_maps.clear(); + } +} + +struct ForeignNetworkForwardCounters { + forward_data_bytes: CounterHandle, + forward_data_packets: CounterHandle, + forward_control_bytes: CounterHandle, + forward_control_packets: CounterHandle, + rx_bytes: CounterHandle, + rx_packets: CounterHandle, +} + +pub(crate) struct ForeignNetworkPacketRouter { + my_node_id: PeerId, + packet_recv: PacketRecvChanReceiver, + rpc_sender: UnboundedSender, + peer_map: Arc, + relay_peer_map: Arc, + traffic_metrics: Arc, + parent_context: ArcPeerContext, + relay_data: bool, + pm_sender: PacketRecvChan, + network_name: String, + bps_limiter: Option, + counters: ForeignNetworkForwardCounters, +} + +impl ForeignNetworkPacketRouter { + #[allow(clippy::too_many_arguments)] + pub fn new( + my_node_id: PeerId, + packet_recv: PacketRecvChanReceiver, + rpc_sender: UnboundedSender, + peer_map: Arc, + relay_peer_map: Arc, + traffic_metrics: Arc, + parent_context: ArcPeerContext, + relay_data: bool, + pm_sender: PacketRecvChan, + network_name: String, + bps_limiter: Option, + stats_mgr: Arc, + ) -> Self { + let label_set = + LabelSet::new().with_label_type(LabelType::NetworkName(network_name.clone())); + let counters = ForeignNetworkForwardCounters { + forward_data_bytes: stats_mgr + .get_counter(MetricName::TrafficBytesForwarded, label_set.clone()), + forward_data_packets: stats_mgr + .get_counter(MetricName::TrafficPacketsForwarded, label_set.clone()), + forward_control_bytes: stats_mgr + .get_counter(MetricName::TrafficControlBytesForwarded, label_set.clone()), + forward_control_packets: stats_mgr.get_counter( + MetricName::TrafficControlPacketsForwarded, + label_set.clone(), + ), + rx_bytes: stats_mgr.get_counter(MetricName::TrafficBytesSelfRx, label_set.clone()), + rx_packets: stats_mgr.get_counter(MetricName::TrafficPacketsRx, label_set), + }; + + Self { + my_node_id, + packet_recv, + rpc_sender, + peer_map, + relay_peer_map, + traffic_metrics, + parent_context, + relay_data, + pm_sender, + network_name, + bps_limiter, + counters, + } + } + + pub async fn run(self) { + let Self { + my_node_id, + mut packet_recv, + rpc_sender, + peer_map, + relay_peer_map, + traffic_metrics, + parent_context, + relay_data, + pm_sender, + network_name, + bps_limiter, + counters, + } = self; + + while let Ok(mut zc_packet) = recv_packet_from_chan(&mut packet_recv).await { + let buf_len = zc_packet.buf_len(); + let Some(hdr) = zc_packet.peer_manager_header() else { + tracing::warn!("invalid packet, skip"); + continue; + }; + tracing::trace!(?hdr, "recv packet in foreign network manager"); + let from_peer_id = hdr.from_peer_id.get(); + let packet_type = hdr.packet_type; + let len = hdr.len.get(); + let to_peer_id = hdr.to_peer_id.get(); + let is_local_delivery = to_peer_id == my_node_id; + let is_locally_originated = from_peer_id == my_node_id; + if is_local_delivery && !is_locally_originated { + traffic_metrics + .record_rx(from_peer_id, packet_type, buf_len as u64) + .await; + } + if is_local_delivery { + if packet_type == PacketType::RelayHandshake as u8 + || packet_type == PacketType::RelayHandshakeAck as u8 + { + let _ = relay_peer_map.handle_handshake_packet(zc_packet).await; + continue; + } + + if relay_peer_map.is_secure_mode_enabled() && hdr.is_encrypted() { + match relay_peer_map.decrypt_if_needed(&mut zc_packet).await { + Ok(true) => {} + Ok(false) => { + tracing::error!("secure session not found"); + continue; + } + Err(e) => { + tracing::error!(?e, "secure decrypt failed"); + continue; + } + } + } + + if packet_type == PacketType::TaRpc as u8 + || packet_type == PacketType::RpcReq as u8 + || packet_type == PacketType::RpcResp as u8 + { + counters.rx_bytes.add(buf_len as u64); + counters.rx_packets.inc(); + rpc_sender.send(zc_packet).unwrap(); + continue; + } + tracing::trace!( + ?packet_type, + ?len, + ?from_peer_id, + ?to_peer_id, + "ignore packet in foreign network" + ); + } else { + if is_relay_data_packet_type(packet_type) { + let disable_relay_data = parent_context.disable_relay_data(); + if !relay_data || disable_relay_data { + tracing::debug!( + ?from_peer_id, + ?to_peer_id, + packet_type, + disable_relay_data, + "drop foreign network relay data" + ); + continue; + } + if let Some(bps_limiter) = bps_limiter.as_ref() + && !bps_limiter.try_consume(len.into()) + { + continue; + } + } + + match traffic_kind(packet_type) { + TrafficKind::Data => { + counters.forward_data_bytes.add(buf_len as u64); + counters.forward_data_packets.inc(); + } + TrafficKind::Control => { + counters.forward_control_bytes.add(buf_len as u64); + counters.forward_control_packets.inc(); + } + } + + let gateway_peer_id = peer_map + .get_gateway_peer_id(to_peer_id, NextHopPolicy::LeastHop) + .await; + + match gateway_peer_id { + Some(peer_id) if peer_map.has_peer(peer_id) => { + if peer_id != to_peer_id && hdr.from_peer_id.get() == my_node_id { + if let Err(e) = relay_peer_map + .send_msg(zc_packet, to_peer_id, NextHopPolicy::LeastHop) + .await + { + tracing::error!( + ?e, + "send packet to foreign peer inside relay peer map failed" + ); + } else if is_locally_originated { + traffic_metrics + .record_tx(to_peer_id, packet_type, buf_len as u64) + .await; + } + } else if let Err(e) = peer_map.send_msg_directly(zc_packet, peer_id).await + { + tracing::error!( + ?e, + "send packet to foreign peer inside peer map failed" + ); + } else if is_locally_originated { + traffic_metrics + .record_tx(to_peer_id, packet_type, buf_len as u64) + .await; + } + } + _ => { + let mut foreign_packet = ZCPacket::new_for_foreign_network( + &network_name, + to_peer_id, + &zc_packet, + ); + let via_peer = gateway_peer_id.unwrap_or(to_peer_id); + foreign_packet.fill_peer_manager_hdr( + my_node_id, + via_peer, + PacketType::ForeignNetworkPacket as u8, + ); + if let Err(e) = pm_sender.send(foreign_packet).await { + tracing::error!("send packet to peer with pm failed: {:?}", e); + } else if is_locally_originated { + traffic_metrics + .record_tx(to_peer_id, packet_type, buf_len as u64) + .await; + } + } + }; + } + } + } +} + +#[cfg(test)] +mod tests { + #![allow(clippy::field_reassign_with_default)] + + use std::sync::{ + Arc, + atomic::{AtomicBool, Ordering}, + }; + + use easytier_proto::common::{FlagsInConfig, PeerFeatureFlag}; + + use super::{ + ForeignNetworkManager, abort_and_join_persistent_tasks, build_foreign_peer_context, + desired_foreign_avoid_relay_data, sync_foreign_avoid_relay_data, + }; + + use crate::peers::whitelist::check_network_in_relay_whitelist; + + use crate::{ + config::peers::PeerRuntimeSnapshot, + config::runtime::{CoreRuntimeConfig, CoreRuntimeConfigStore}, + foundation::stats::{LabelSet, LabelType, MetricName}, + peers::{ + context::{ + ArcPeerContext, CorePeerContext, CorePeerContextAdapters, NetworkIdentity, + PeerContext, + }, + error::Error, + }, + proto::common::StunInfo, + }; + + impl ForeignNetworkManager { + pub(crate) async fn is_stopped_for_test(&self) -> bool { + *self.lifecycle.read().await == super::ForeignNetworkManagerState::Stopped + && self.task_reaper.lock().await.is_none() + && self.tasks.lock().unwrap().is_empty() + && self.data.network_peer_maps.is_empty() + && self.data.peer_network_map.is_empty() + } + + pub(crate) async fn admission_is_open_for_test(&self) -> bool { + self.admission_guard().await.is_ok() + } + + pub(crate) async fn hold_admission_for_test( + &self, + entered: Arc, + release: Arc, + ) -> Result<(), Error> { + let _admission = self.admission_guard().await?; + entered.notify_one(); + release.notified().await; + Ok(()) + } + } + + #[tokio::test] + async fn cancelled_join_wait_keeps_persistent_task_ownership() { + struct DropFlag(Arc); + + impl Drop for DropFlag { + fn drop(&mut self) { + self.0.store(true, Ordering::Release); + } + } + + let tasks = std::sync::Mutex::new(tokio::task::JoinSet::new()); + let entered = Arc::new(tokio::sync::Notify::new()); + let dropped = Arc::new(AtomicBool::new(false)); + let task_entered = entered.clone(); + let task_dropped = dropped.clone(); + tasks.lock().unwrap().spawn(async move { + let _drop = DropFlag(task_dropped); + task_entered.notify_one(); + std::future::pending::<()>().await; + }); + entered.notified().await; + + { + let first_wait = abort_and_join_persistent_tasks(&tasks); + tokio::pin!(first_wait); + assert!(matches!( + futures::poll!(first_wait.as_mut()), + std::task::Poll::Pending + )); + } + + abort_and_join_persistent_tasks(&tasks).await; + assert!(dropped.load(Ordering::Acquire)); + assert!(tasks.lock().unwrap().is_empty()); + } + + struct FeatureContext { + avoid_relay_data: AtomicBool, + flags: FlagsInConfig, + hostname: String, + } + + impl FeatureContext { + fn new(avoid_relay_data: bool) -> Self { + Self { + avoid_relay_data: AtomicBool::new(avoid_relay_data), + flags: FlagsInConfig::default(), + hostname: String::new(), + } + } + } + + impl PeerContext for FeatureContext { + fn network_identity(&self) -> NetworkIdentity { + NetworkIdentity::default() + } + + fn feature_flags(&self) -> PeerFeatureFlag { + PeerFeatureFlag { + avoid_relay_data: self.avoid_relay_data.load(Ordering::Acquire), + ..Default::default() + } + } + + fn flags(&self) -> FlagsInConfig { + self.flags.clone() + } + + fn hostname(&self) -> String { + self.hostname.clone() + } + + fn set_avoid_relay_data_preference(&self, avoid_relay_data: bool) -> bool { + self.avoid_relay_data + .swap(avoid_relay_data, Ordering::AcqRel) + != avoid_relay_data + } + } + + #[test] + fn relay_whitelist_supports_exact_wildcard_and_empty_rules() { + assert!(check_network_in_relay_whitelist("net1 net2*", "net1").is_ok()); + assert!(check_network_in_relay_whitelist("net1 net2*", "net2-west").is_ok()); + assert!(check_network_in_relay_whitelist("*", "any-network").is_ok()); + assert!(check_network_in_relay_whitelist("", "net1").is_err()); + assert!(check_network_in_relay_whitelist("net1 net2*", "net3").is_err()); + } + + #[test] + fn foreign_avoid_relay_data_policy_tracks_parent_and_relay_permission() { + let parent_state = Arc::new(FeatureContext::new(false)); + let parent: ArcPeerContext = parent_state.clone(); + let foreign_state = Arc::new(FeatureContext::new(true)); + let foreign: ArcPeerContext = foreign_state.clone(); + + assert!(!desired_foreign_avoid_relay_data(&parent, true)); + assert!(sync_foreign_avoid_relay_data(&parent, &foreign, true)); + assert!(!foreign.feature_flags().avoid_relay_data); + assert!(!sync_foreign_avoid_relay_data(&parent, &foreign, true)); + + parent_state.avoid_relay_data.store(true, Ordering::Release); + assert!(sync_foreign_avoid_relay_data(&parent, &foreign, true)); + assert!(foreign.feature_flags().avoid_relay_data); + + parent_state + .avoid_relay_data + .store(false, Ordering::Release); + assert!(!sync_foreign_avoid_relay_data(&parent, &foreign, false)); + assert!(foreign.feature_flags().avoid_relay_data); + } + + #[test] + fn foreign_context_resources_are_assembled_in_core() { + let mut parent_snapshot = PeerRuntimeSnapshot::default(); + parent_snapshot.runtime.core.node.hostname = Some("parent".to_owned()); + parent_snapshot.runtime.stun_info = StunInfo { + public_ip: vec!["198.51.100.1".to_owned()], + ..Default::default() + }; + parent_snapshot + .runtime + .host_routing + .local_exit_node_fallback = true; + parent_snapshot.flags.enable_relay_foreign_network_kcp = true; + parent_snapshot.flags.enable_relay_foreign_network_quic = false; + parent_snapshot.flags.socket_mark = Some(7); + parent_snapshot.hmac_secret_digest = true; + parent_snapshot.runtime.feature_flags = PeerFeatureFlag { + kcp_input: false, + no_relay_kcp: false, + support_conn_list_sync: false, + quic_input: false, + no_relay_quic: false, + need_p2p: true, + disable_p2p: true, + ipv6_public_addr_provider: true, + ..Default::default() + }; + let parent_config = CoreRuntimeConfigStore::new( + CoreRuntimeConfig::default(), + Arc::new(parent_snapshot.clone()), + ); + let parent = Arc::new(CorePeerContext::new( + parent_config.clone(), + Arc::new(()), + CorePeerContextAdapters { + stun_info_source: None, + events: Arc::new(()), + credential_storage: None, + }, + )); + let network = NetworkIdentity { + network_name: "foreign".to_owned(), + network_secret: Some("secret".to_owned()), + network_secret_digest: None, + }; + let mut defaults = FlagsInConfig::default(); + defaults.mtu = 1400; + defaults.relay_network_whitelist = "baseline".to_owned(); + let foreign = build_foreign_peer_context(&network, &parent, false, defaults); + + assert!(Arc::ptr_eq( + &parent.stats_manager(), + &foreign.stats_manager(), + )); + assert!(!Arc::ptr_eq( + &parent.credential_manager(), + &foreign.credential_manager(), + )); + assert!(!Arc::ptr_eq( + &parent.trusted_key_manager(), + &foreign.trusted_key_manager(), + )); + assert_eq!(foreign.network_name(), "foreign"); + assert_eq!(foreign.hostname(), "PublicServer_parent"); + assert!(foreign.secure_mode().is_none()); + assert!(!foreign.flags().disable_relay_kcp); + assert!(foreign.flags().disable_relay_quic); + assert_eq!(foreign.flags().socket_mark, Some(7)); + assert_eq!(foreign.flags().mtu, 1400); + assert_eq!(foreign.flags().relay_network_whitelist, "baseline"); + assert!(foreign.feature_flags().is_public_server); + assert!(foreign.feature_flags().avoid_relay_data); + assert!(!foreign.feature_flags().kcp_input); + assert!(!foreign.feature_flags().no_relay_kcp); + assert!(!foreign.feature_flags().support_conn_list_sync); + assert!(!foreign.feature_flags().quic_input); + assert!(!foreign.feature_flags().no_relay_quic); + assert!(foreign.feature_flags().need_p2p); + assert!(foreign.feature_flags().disable_p2p); + assert!(foreign.feature_flags().ipv6_public_addr_provider); + assert_eq!(foreign.stun_info().public_ip, vec!["198.51.100.1"]); + assert!(foreign.host_routing_policy().local_exit_node_fallback); + assert!(foreign.hmac_secret_digest()); + + parent_snapshot.runtime.stun_info.public_ip = vec!["203.0.113.2".to_owned()]; + parent_config.update_peer(Arc::new(parent_snapshot)); + assert_eq!(foreign.stun_info().public_ip, vec!["203.0.113.2"]); + + foreign.record_control_tx("foreign", 64); + let labels = LabelSet::new().with_label_type(LabelType::NetworkName("foreign".to_owned())); + assert_eq!( + parent + .stats_manager() + .get_metric(MetricName::TrafficControlBytesTx, &labels) + .unwrap() + .value, + 64 + ); + } +} diff --git a/easytier-core/src/peers/mod.rs b/easytier-core/src/peers/mod.rs new file mode 100644 index 00000000..244c3a8b --- /dev/null +++ b/easytier-core/src/peers/mod.rs @@ -0,0 +1,60 @@ +pub(crate) mod acl; +pub(crate) mod admission; +pub(crate) mod conn; +pub mod context; +pub mod credential_manager; +pub mod error; +pub mod foreign_network; +pub mod peer_center; +pub mod peer_manager; +pub(crate) mod peer_rpc; +pub mod public_ipv6; +pub(crate) mod relay_peer_map; +pub(crate) mod route; +pub(crate) mod traffic_metrics; +mod util; +pub(crate) mod whitelist; + +#[cfg(test)] +pub(crate) mod test_support; +#[cfg(test)] +mod tests; + +use crate::packet::ZCPacket; + +pub type PacketRecvChan = tokio::sync::mpsc::Sender; +pub type PacketRecvChanReceiver = tokio::sync::mpsc::Receiver; + +pub fn create_packet_recv_chan() -> (PacketRecvChan, PacketRecvChanReceiver) { + tokio::sync::mpsc::channel(128) +} + +pub async fn recv_packet_from_chan( + packet_recv_chan_receiver: &mut PacketRecvChanReceiver, +) -> Result { + packet_recv_chan_receiver + .recv() + .await + .ok_or(anyhow::anyhow!("recv_packet_from_chan failed")) +} + +#[async_trait::async_trait] +#[auto_impl::auto_impl(Arc)] +pub trait PeerPacketFilter { + async fn try_process_packet_from_peer(&self, zc_packet: ZCPacket) -> Option { + Some(zc_packet) + } +} + +#[async_trait::async_trait] +#[auto_impl::auto_impl(Arc)] +pub trait NicPacketFilter { + async fn try_process_packet_from_nic(&self, data: &mut ZCPacket) -> bool; + + fn id(&self) -> String { + format!("{:p}", self) + } +} + +pub type BoxPeerPacketFilter = Box; +pub type BoxNicPacketFilter = Box; diff --git a/easytier/src/peer_center/instance.rs b/easytier-core/src/peers/peer_center/instance.rs similarity index 67% rename from easytier/src/peer_center/instance.rs rename to easytier-core/src/peers/peer_center/instance.rs index 81bf32b3..2e8525d0 100644 --- a/easytier/src/peer_center/instance.rs +++ b/easytier-core/src/peers/peer_center/instance.rs @@ -12,19 +12,17 @@ use tokio::task::JoinSet; use tracing::Instrument; use crate::{ - common::{PeerId, global_ctx::GlobalCtx}, + config::PeerId, peers::{ - peer_manager::PeerManager, - peer_map::PeerMap, peer_rpc::PeerRpcManager, - route_trait::{RouteCostCalculator, RouteCostCalculatorInterface}, - rpc_service::PeerManagerRpcService, + route::{RouteCostCalculator, RouteCostCalculatorInterface}, }, proto::{ + core_peer::peer::Route as CoreRoute, peer_rpc::{ - DirectConnectedPeerInfo, GetGlobalPeerMapRequest, GetGlobalPeerMapResponse, - GlobalPeerMap, PeerCenterRpc, PeerCenterRpcClientFactory, PeerCenterRpcServer, - PeerInfoForGlobalMap, ReportPeersRequest, ReportPeersResponse, + GetGlobalPeerMapRequest, GetGlobalPeerMapResponse, GlobalPeerMap, PeerCenterRpc, + PeerCenterRpcClientFactory, PeerCenterRpcServer, PeerInfoForGlobalMap, + ReportPeersRequest, ReportPeersResponse, }, rpc_types::{self, controller::BaseController}, }, @@ -37,9 +35,9 @@ use super::{Digest, Error, server::PeerCenterServer}; pub trait PeerCenterPeerManagerTrait: Send + Sync + 'static { async fn list_peers(&self) -> PeerInfoForGlobalMap; fn my_peer_id(&self) -> PeerId; - fn get_global_ctx(&self) -> Arc; + fn network_name(&self) -> String; fn get_rpc_mgr(&self) -> Weak; - async fn list_routes(&self) -> Vec; + async fn list_routes(&self) -> Vec; } struct PeerCenterBase { @@ -49,11 +47,7 @@ struct PeerCenterBase { lock: Arc>, } -// static SERVICE_ID: u32 = 5; for compatibility with the original code -static SERVICE_ID: u32 = 50; - struct PeridicJobCtx { - peer_mgr: Arc, my_peer_id: PeerId, center_peer: AtomicCell, job_ctx: T, @@ -66,7 +60,7 @@ impl PeerCenterBase { }; rpc_mgr.rpc_server().registry().register( PeerCenterRpcServer::new(PeerCenterServer::new()), - &self.peer_mgr.get_global_ctx().get_network_name(), + &self.peer_mgr.network_name(), ); Ok(()) } @@ -110,7 +104,6 @@ impl PeerCenterBase { self.tasks.lock().await.spawn( async move { let ctx = Arc::new(PeridicJobCtx { - peer_mgr: peer_mgr.clone(), my_peer_id, center_peer: AtomicCell::new(PeerId::default()), job_ctx, @@ -118,7 +111,7 @@ impl PeerCenterBase { loop { let Some(center_peer) = Self::select_center_peer(&peer_mgr).await else { tracing::trace!("no center peer found, sleep 1 second"); - tokio::time::sleep(Duration::from_secs(1)).await; + crate::foundation::time::sleep(Duration::from_secs(1)).await; continue; }; let Some(rpc_mgr) = peer_mgr.get_rpc_mgr().upgrade() else { @@ -134,19 +127,20 @@ impl PeerCenterBase { .scoped_client::>( my_peer_id, center_peer, - peer_mgr.get_global_ctx().get_network_name(), + peer_mgr.network_name(), ); let ret = job_fn(stub, ctx.clone()).await; drop(_g); let Ok(sleep_time_ms) = ret else { tracing::error!("periodic job to center server rpc failed: {:?}", ret); - tokio::time::sleep(Duration::from_secs(3)).await; + crate::foundation::time::sleep(Duration::from_secs(3)).await; continue; }; if sleep_time_ms > 0 { - tokio::time::sleep(Duration::from_millis(sleep_time_ms as u64)).await; + crate::foundation::time::sleep(Duration::from_millis(sleep_time_ms as u64)) + .await; } } } @@ -163,6 +157,12 @@ impl PeerCenterBase { lock: Arc::new(Mutex::new(())), } } + + async fn stop(&self) { + let mut tasks = self.tasks.lock().await; + tasks.abort_all(); + while tasks.join_next().await.is_some() {} + } } #[derive(Clone)] @@ -216,12 +216,23 @@ impl PeerCenterInstance { } } + pub fn global_peer_map_snapshot(&self) -> GetGlobalPeerMapResponse { + GetGlobalPeerMapResponse { + global_peer_map: self.global_peer_map.read().unwrap().map.clone(), + digest: Some(self.global_peer_map_digest.load()), + } + } + pub async fn init(&self) { self.client.init().await.unwrap(); self.init_get_global_info_job().await; self.init_report_peers_job().await; } + pub async fn stop(&self) { + self.client.stop().await; + } + async fn init_get_global_info_job(&self) { struct Ctx { global_peer_map: Arc>, @@ -407,156 +418,3 @@ impl PeerCenterInstance { }) } } - -#[async_trait::async_trait] -impl PeerCenterPeerManagerTrait for PeerManager { - async fn list_peers(&self) -> PeerInfoForGlobalMap { - PeerManagerRpcService::list_peers(self).await.into() - } - - fn my_peer_id(&self) -> PeerId { - self.get_peer_map().my_peer_id() - } - - fn get_global_ctx(&self) -> Arc { - self.get_peer_map().get_global_ctx() - } - - fn get_rpc_mgr(&self) -> Weak { - Arc::downgrade(&self.get_peer_rpc_mgr()) - } - - async fn list_routes(&self) -> Vec { - self.list_routes().await - } -} - -pub struct PeerMapWithPeerRpcManager { - pub peer_map: Arc, - pub rpc_mgr: Arc, -} - -#[async_trait::async_trait] -impl PeerCenterPeerManagerTrait for PeerMapWithPeerRpcManager { - async fn list_peers(&self) -> PeerInfoForGlobalMap { - // TODO: currently latency between public server cannot be calculated because one public-server pair - // has no connection between them. (hard to get latency from peer manager because it's hard to transfrom the peer id) - // but it's fine because we don't want to too much traffic between public servers. - let peers = self.peer_map.list_peers(); - let mut ret = PeerInfoForGlobalMap::default(); - for peer in peers { - if let Some(conns) = self.peer_map.list_peer_conns(peer).await { - let Some(min_lat) = conns - .iter() - .map(|conn| conn.stats.as_ref().unwrap().latency_us) - .min() - else { - continue; - }; - - ret.direct_peers.insert( - peer, - DirectConnectedPeerInfo { - latency_ms: std::cmp::max(1, (min_lat as u32 / 1000) as i32), - }, - ); - } - } - - ret - } - - fn my_peer_id(&self) -> PeerId { - self.peer_map.my_peer_id() - } - - fn get_global_ctx(&self) -> Arc { - self.peer_map.get_global_ctx() - } - - fn get_rpc_mgr(&self) -> Weak { - Arc::downgrade(&self.rpc_mgr) - } - - async fn list_routes(&self) -> Vec { - self.peer_map.list_route_infos().await - } -} - -#[cfg(test)] -mod tests { - use crate::{ - peers::tests::{connect_peer_manager, create_mock_peer_manager, wait_route_appear}, - tunnel::common::tests::wait_for_condition, - }; - - use super::*; - - #[tokio::test] - async fn test_peer_center_instance() { - let peer_mgr_a = create_mock_peer_manager().await; - let peer_mgr_b = create_mock_peer_manager().await; - let peer_mgr_c = create_mock_peer_manager().await; - - let peer_center_a = PeerCenterInstance::new(peer_mgr_a.clone()); - let peer_center_b = PeerCenterInstance::new(peer_mgr_b.clone()); - let peer_center_c = PeerCenterInstance::new(peer_mgr_c.clone()); - - let peer_centers = [&peer_center_a, &peer_center_b, &peer_center_c]; - for pc in peer_centers.iter() { - pc.init().await; - } - - connect_peer_manager(peer_mgr_a.clone(), peer_mgr_b.clone()).await; - connect_peer_manager(peer_mgr_b.clone(), peer_mgr_c.clone()).await; - - wait_route_appear(peer_mgr_a.clone(), peer_mgr_c.clone()) - .await - .unwrap(); - - let mut digest = None; - for pc in peer_centers.iter() { - let rpc_service = pc.get_rpc_service(); - wait_for_condition( - || async { rpc_service.global_peer_map.read().unwrap().map.len() == 3 }, - Duration::from_secs(20), - ) - .await; - - println!("rpc service ready, {:#?}", rpc_service.global_peer_map); - - if let Some(prev) = digest { - let v = rpc_service.global_peer_map_digest.load(); - assert_eq!(prev, v); - digest = Some(prev); - } else { - digest = Some(rpc_service.global_peer_map_digest.load()); - } - - let mut route_cost = pc.get_cost_calculator(); - assert!(route_cost.need_update()); - - route_cost.begin_update(); - assert!( - route_cost.calculate_cost(peer_mgr_a.my_peer_id(), peer_mgr_b.my_peer_id()) < 30 - ); - assert!( - route_cost.calculate_cost(peer_mgr_b.my_peer_id(), peer_mgr_a.my_peer_id()) < 30 - ); - assert!( - route_cost.calculate_cost(peer_mgr_b.my_peer_id(), peer_mgr_c.my_peer_id()) < 30 - ); - assert!( - route_cost.calculate_cost(peer_mgr_c.my_peer_id(), peer_mgr_b.my_peer_id()) < 30 - ); - assert!( - route_cost.calculate_cost(peer_mgr_c.my_peer_id(), peer_mgr_a.my_peer_id()) > 50 - ); - assert!( - route_cost.calculate_cost(peer_mgr_a.my_peer_id(), peer_mgr_c.my_peer_id()) > 50 - ); - route_cost.end_update(); - assert!(!route_cost.need_update()); - } - } -} diff --git a/easytier-core/src/peers/peer_center/mod.rs b/easytier-core/src/peers/peer_center/mod.rs new file mode 100644 index 00000000..26296fb2 --- /dev/null +++ b/easytier-core/src/peers/peer_center/mod.rs @@ -0,0 +1,21 @@ +// peer_center is used to collect peer info into one peer node. +// the center node is selected with the following rules: +// 1. has smallest peer id +// 2. TODO: has allow_to_be_center peer feature +// peer center is not guaranteed to be stable and can be changed when peer enter or leave. +// it's used to reduce the cost to exchange infos between peers. + +pub mod instance; +mod server; + +#[derive(thiserror::Error, Debug, serde::Deserialize, serde::Serialize)] +pub enum Error { + #[error("Digest not match, need provide full peer info to center server.")] + DigestMismatch, + #[error("Not center server")] + NotCenterServer, + #[error("Instance shutdown")] + Shutdown, +} + +pub type Digest = u64; diff --git a/easytier/src/peer_center/server.rs b/easytier-core/src/peers/peer_center/server.rs similarity index 97% rename from easytier/src/peer_center/server.rs rename to easytier-core/src/peers/peer_center/server.rs index 10129f13..1fdac7a1 100644 --- a/easytier/src/peer_center/server.rs +++ b/easytier-core/src/peers/peer_center/server.rs @@ -9,7 +9,7 @@ use dashmap::DashMap; use tokio::task::JoinSet; use crate::{ - common::PeerId, + config::PeerId, proto::{ peer_rpc::{ DirectConnectedPeerInfo, GetGlobalPeerMapRequest, GetGlobalPeerMapResponse, @@ -44,7 +44,7 @@ struct PeerCenterServerData { #[derive(Clone, Debug)] pub struct PeerCenterServer { data: Arc, - tasks: Arc>, + _tasks: Arc>, } impl PeerCenterServer { @@ -54,7 +54,7 @@ impl PeerCenterServer { let mut tasks = JoinSet::new(); tasks.spawn(async move { loop { - tokio::time::sleep(std::time::Duration::from_secs(10)).await; + crate::foundation::time::sleep(std::time::Duration::from_secs(10)).await; let Some(data) = weak_data.upgrade() else { break; }; @@ -64,7 +64,7 @@ impl PeerCenterServer { PeerCenterServer { data, - tasks: Arc::new(tasks), + _tasks: Arc::new(tasks), } } diff --git a/easytier-core/src/peers/peer_manager.rs b/easytier-core/src/peers/peer_manager.rs new file mode 100644 index 00000000..e12d50f5 --- /dev/null +++ b/easytier-core/src/peers/peer_manager.rs @@ -0,0 +1,3933 @@ +use std::time::{Duration, SystemTime}; +use std::{ + collections::BTreeSet, + net::{IpAddr, Ipv4Addr, Ipv6Addr}, + sync::{ + Arc, Weak, + atomic::{AtomicBool, Ordering}, + }, +}; + +use anyhow::Context; +use dashmap::DashMap; +use parking_lot::RwLock as SyncRwLock; +use quanta::Instant; +use serde::{Deserialize, Serialize}; +use tokio::sync::{ + Mutex, RwLock, + mpsc::{self, UnboundedReceiver, UnboundedSender}, +}; +use tokio::task::JoinSet; +use url::Url; + +use crate::{ + config::peers::{HostRoutingPolicy, PeerRuntimeConfig, PeerRuntimeSnapshot}, + config::runtime::CoreRuntimeConfigStore, + config::{P2pPolicyFlags, PeerId, ProxyNetworkConfig}, + events::CoreEventSink, + foundation::task::ExternalTaskSignal, + packet::{ + CompressorAlgo, PacketType, ZCPacket, + compressor::{Compressor as _, DefaultCompressor}, + }, + proto::common::{FlagsInConfig, PeerFeatureFlag, StunInfo, Url as ProtoUrl}, + proto::core_peer::peer::{ListPublicIpv6InfoResponse, PeerConnInfo, Route as CoreRoute}, + tunnel::{ + Tunnel, + encrypt::{ + Encryptor, NullCipher, create_encryptor, derive_key_128, derive_key_256, + validate_algorithm, + }, + }, +}; + +use super::{ + BoxNicPacketFilter, BoxPeerPacketFilter, PacketRecvChan, PacketRecvChanReceiver, + PeerPacketFilter, + acl::AclFilter, + conn::{ + peer_conn::{PeerConn, PeerConnId}, + peer_map::{PeerMap, direct_peer_info}, + peer_session::PeerSessionStore, + }, + context::{ + ArcPeerContext, CorePeerContext, CorePeerContextAdapters, NetworkIdentity, PeerContext, + PeerStunInfoSource, + }, + credential_manager::{CredentialManager, CredentialStorage}, + error::Error, + foreign_network::client::ForeignNetworkClient, + foreign_network::{ForeignNetworkEntryInfo, ForeignNetworkManager, ForeignNetworkRpcRegistrar}, + peer_center::instance::PeerCenterPeerManagerTrait, + peer_rpc::{PeerRpcManager, PeerRpcManagerTransport}, + public_ipv6::{CorePublicIpv6Runtime, PublicIpv6Runtime}, + recv_packet_from_chan, + relay_peer_map::RelayPeerMap, + route::{ + ArcRoute, DisabledRoute, ForeignNetworkRouteInfoMap, NextHopPolicy, Route, RouteInterface, + peer_ospf_route::PeerRoute, + }, + traffic_metrics::{ + InstanceLabelKind, LogicalTrafficMetrics, TrafficKind, TrafficMetricRecorder, + route_peer_info_instance_id, traffic_kind, + }, + util::shrink_dashmap, +}; +use crate::foundation::stats::{CounterHandle, LabelSet, LabelType, MetricName, StatsManager}; +use crate::proto::peer_rpc::{ + ForeignNetworkRouteInfoEntry, ForeignNetworkRouteInfoKey, GetIpListResponse, PeerIdentityType, + PeerInfoForGlobalMap, RouteForeignNetworkInfos, RouteForeignNetworkSummary, +}; + +#[derive(Debug, Clone)] +pub struct PeerSnapshot { + pub peer_id: PeerId, + pub default_conn_id: Option, + pub directly_connected_conns: Vec, + pub conns: Vec, +} + +#[derive(Debug, Clone)] +pub struct NodeSnapshot { + pub peer_id: PeerId, + pub ipv4_addr: Option, + pub proxy_networks: Vec, + pub hostname: String, + pub stun_info: StunInfo, + pub instance_id: uuid::Uuid, + pub listeners: Vec, + pub version: String, + pub feature_flags: PeerFeatureFlag, + pub ip_list: GetIpListResponse, + pub public_ipv6_addr: Option, + pub ipv6_public_addr_prefix: Option, +} + +struct RpcTransport { + my_peer_id: PeerId, + peers: Weak, + + packet_recv: Mutex>, + peer_rpc_tspt_sender: UnboundedSender, + + encryptor: Arc, + is_secure_mode_enabled: bool, +} + +impl RpcTransport { + pub fn new( + my_peer_id: PeerId, + peers: Weak, + encryptor: Arc, + is_secure_mode_enabled: bool, + ) -> Arc { + let (peer_rpc_tspt_sender, peer_rpc_tspt_recv) = mpsc::unbounded_channel(); + Arc::new(Self { + my_peer_id, + peers, + packet_recv: Mutex::new(peer_rpc_tspt_recv), + peer_rpc_tspt_sender, + encryptor, + is_secure_mode_enabled, + }) + } + + pub fn packet_sender(&self) -> UnboundedSender { + self.peer_rpc_tspt_sender.clone() + } +} + +#[async_trait::async_trait] +impl PeerRpcManagerTransport for RpcTransport { + fn my_peer_id(&self) -> PeerId { + self.my_peer_id + } + + async fn send(&self, mut msg: ZCPacket, dst_peer_id: PeerId) -> anyhow::Result<()> { + let peers = self + .peers + .upgrade() + .ok_or_else(|| anyhow::anyhow!("peer map is gone"))?; + // NOTE: if route info is not exchanged, this will return None. treat it as public server. + let is_dst_peer_public_server = peers + .get_route_peer_info(dst_peer_id) + .await + .and_then(|x| x.feature_flag.map(|x| x.is_public_server)) + // if dst is directly connected, it's must not public server + .unwrap_or(!peers.has_peer(dst_peer_id)); + if !is_dst_peer_public_server && !self.is_secure_mode_enabled { + self.encryptor + .encrypt(&mut msg) + .with_context(|| "encrypt failed")?; + } + // send to self and this packet will be forwarded in peer_recv loop + peers.send_msg_directly(msg, self.my_peer_id).await?; + Ok(()) + } + + async fn recv(&self) -> anyhow::Result { + if let Some(o) = self.packet_recv.lock().await.recv().await { + Ok(o) + } else { + Err(anyhow::anyhow!("rpc transport is closed")) + } + } +} + +pub(crate) fn get_next_hop_policy(is_latency_first: bool) -> NextHopPolicy { + if is_latency_first { + NextHopPolicy::LeastCost + } else { + NextHopPolicy::LeastHop + } +} + +fn random_peer_id() -> PeerId { + loop { + let peer_id = rand::random(); + if peer_id != 0 { + return peer_id; + } + } +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] +pub enum RouteAlgoType { + Ospf, + None, +} + +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct PortablePeerManagerConfig { + pub snapshot: PeerRuntimeSnapshot, + pub route_algo: RouteAlgoType, + pub exit_nodes: Vec, + /// Defaults inherited by peer contexts created for foreign networks. + /// + /// This is explicit because those contexts participate in the same + /// handshake as the parent but do not inherit all parent policy flags. + pub foreign_context_default_flags: FlagsInConfig, +} + +impl PortablePeerManagerConfig { + pub fn new(mut runtime: PeerRuntimeConfig) -> Self { + let policy = &runtime.core.peer_policy; + let traffic = &runtime.core.traffic; + let flags = FlagsInConfig { + enable_encryption: policy.encryption_required, + encryption_algorithm: crate::config::EncryptionAlgorithm::default().to_string(), + disable_p2p: !policy.p2p_enabled, + relay_all_peer_rpc: policy.relay_peer_rpc, + disable_relay_data: !policy.relay_data, + latency_first: policy.latency_first, + data_compress_algo: crate::proto::common::CompressionAlgoPb::None.into(), + mtu: traffic.mtu.map(u32::from).unwrap_or_default(), + instance_recv_bps_limit: traffic.instance_recv_bps_limit.unwrap_or_default(), + foreign_relay_bps_limit: traffic.foreign_relay_bps_limit.unwrap_or_default(), + ..Default::default() + }; + runtime.feature_flags.disable_p2p = flags.disable_p2p; + runtime.feature_flags.avoid_relay_data |= flags.disable_relay_data; + let foreign_context_default_flags = flags.clone(); + Self { + snapshot: PeerRuntimeSnapshot::new(runtime, flags), + route_algo: RouteAlgoType::Ospf, + exit_nodes: Vec::new(), + foreign_context_default_flags, + } + } +} + +fn validate_portable_routes(routes: &crate::config::RouteConfig) -> anyhow::Result<()> { + if !routes.advertised_routes.is_empty() { + anyhow::bail!("portable peer manager does not support advertised routes yet"); + } + if !routes.foreign_networks.is_empty() { + anyhow::bail!("portable peer manager does not support foreign networks yet"); + } + if let Some(prefix) = &routes.ipv4 + && (!matches!(prefix.address, IpAddr::V4(_)) || prefix.prefix_len > 32) + { + anyhow::bail!("routes.ipv4 must contain a valid IPv4 prefix"); + } + if let Some(prefix) = &routes.ipv6 + && (!matches!(prefix.address, IpAddr::V6(_)) || prefix.prefix_len > 128) + { + anyhow::bail!("routes.ipv6 must contain a valid IPv6 prefix"); + } + for proxy in &routes.proxy_networks { + for (field, prefix) in [ + ("real", Some(&proxy.real)), + ("mapped", proxy.mapped.as_ref()), + ] { + let Some(prefix) = prefix else { + continue; + }; + let IpAddr::V4(address) = prefix.address else { + anyhow::bail!("proxy network {field} prefix must be IPv4"); + }; + if cidr::Ipv4Cidr::new(address, prefix.prefix_len).is_err() { + anyhow::bail!("proxy network {field} must be a valid IPv4 network prefix"); + } + } + } + Ok(()) +} + +pub(crate) enum RouteAlgoInst { + Ospf(Arc), + None, +} + +impl Clone for RouteAlgoInst { + fn clone(&self) -> Self { + match self { + RouteAlgoInst::Ospf(route) => RouteAlgoInst::Ospf(route.clone()), + RouteAlgoInst::None => RouteAlgoInst::None, + } + } +} + +impl RouteAlgoInst { + pub fn new( + route_algo: RouteAlgoType, + my_peer_id: PeerId, + context: ArcPeerContext, + public_ipv6_runtime: Arc, + peer_rpc_mgr: Arc, + ) -> Self { + match route_algo { + RouteAlgoType::Ospf => RouteAlgoInst::Ospf(PeerRoute::new( + my_peer_id, + context, + public_ipv6_runtime, + peer_rpc_mgr, + )), + RouteAlgoType::None => RouteAlgoInst::None, + } + } + + pub fn ospf_route(&self) -> Option> { + match self { + RouteAlgoInst::Ospf(route) => Some(route.clone()), + RouteAlgoInst::None => None, + } + } + + pub fn route_arc(&self) -> ArcRoute { + match self { + RouteAlgoInst::Ospf(route) => route.clone(), + RouteAlgoInst::None => Arc::new(DisabledRoute), + } + } +} + +fn network_secret_digest_is_empty(network: &NetworkIdentity) -> bool { + network + .network_secret_digest + .as_ref() + .is_none_or(|d| d.iter().all(|b| *b == 0)) +} + +pub(crate) async fn add_new_peer_conn( + peer_map: &PeerMap, + local_identity: &NetworkIdentity, + local_secure_mode: bool, + peer_conn: PeerConn, +) -> Result { + let peer_identity = peer_conn.get_network_identity(); + let conn_info = peer_conn.get_conn_info(); + let peer_secure_mode = !conn_info.noise_remote_static_pubkey.is_empty(); + + if local_secure_mode != peer_secure_mode { + return Err(Error::SecretKeyError( + "same-network peers must use the same secure mode".to_string(), + )); + } + + // For credential nodes, network_secret_digest is either None or all-zeros + // (all-zeros when received over the wire via handshake). + // In this case, only compare network_name. + let my_digest_empty = network_secret_digest_is_empty(local_identity); + let peer_digest_empty = network_secret_digest_is_empty(&peer_identity); + + let identity_ok = if my_digest_empty || peer_digest_empty { + // Credential node: only check network_name + local_identity.network_name == peer_identity.network_name + } else { + local_identity == &peer_identity + }; + + if !identity_ok { + return Err(Error::SecretKeyError( + "network identity not match".to_string(), + )); + } + let peer_id = peer_conn.get_peer_id(); + peer_map.add_new_peer_conn(peer_conn).await?; + Ok(peer_id) +} + +pub(crate) async fn close_untrusted_credential_peers( + peer_map: &PeerMap, + network_name: &str, + mut is_pubkey_trusted: F, +) where + F: FnMut(&[u8], &str) -> bool + Send, +{ + for peer_id in peer_map.list_peers() { + if !matches!( + peer_map.get_peer_identity_type(peer_id), + Some(PeerIdentityType::Credential) + ) { + continue; + } + let Some(peer) = peer_map.get_peer_by_id(peer_id) else { + continue; + }; + let Some(pubkey) = peer.get_peer_public_key() else { + continue; + }; + + if is_pubkey_trusted(&pubkey, network_name) { + continue; + } + + tracing::warn!(?peer_id, "closing untrusted credential peer"); + if let Err(e) = peer_map.close_peer(peer_id).await { + tracing::warn!(?e, ?peer_id, "failed to close untrusted credential peer"); + } + } +} + +struct NicPacketProcessor { + nic_channel: PacketRecvChan, +} + +#[async_trait::async_trait] +impl PeerPacketFilter for NicPacketProcessor { + async fn try_process_packet_from_peer(&self, packet: ZCPacket) -> Option { + let hdr = packet.peer_manager_header().unwrap(); + if hdr.packet_type == PacketType::Data as u8 && !hdr.is_not_send_to_tun() { + if hdr.is_encrypted() || hdr.is_compressed() { + tracing::warn!( + from_peer_id = hdr.from_peer_id.get(), + to_peer_id = hdr.to_peer_id.get(), + encrypted = hdr.is_encrypted(), + compressed = hdr.is_compressed(), + "dropping packet before nic because it is not fully decoded" + ); + return None; + } + tracing::trace!(?packet, "send packet to nic channel"); + let _ = self.nic_channel.send(packet).await; + None + } else { + Some(packet) + } + } +} + +struct PeerRpcPacketProcessor { + peer_rpc_tspt_sender: UnboundedSender, +} + +#[async_trait::async_trait] +impl PeerPacketFilter for PeerRpcPacketProcessor { + async fn try_process_packet_from_peer(&self, packet: ZCPacket) -> Option { + let hdr = packet.peer_manager_header().unwrap(); + if hdr.packet_type == PacketType::TaRpc as u8 + || hdr.packet_type == PacketType::RpcReq as u8 + || hdr.packet_type == PacketType::RpcResp as u8 + { + self.peer_rpc_tspt_sender.send(packet).unwrap(); + None + } else { + Some(packet) + } + } +} + +pub(crate) struct PeerPipelineEntry { + active: Arc, + filter: Arc>>>, +} + +pub(crate) struct NicPipelineEntry { + active: Arc, + filter: Arc>>>, +} + +#[derive(Clone)] +pub(crate) struct PipelineRegistrationGuard { + active: Arc, + release_filter: Arc, +} + +impl PipelineRegistrationGuard { + pub fn close(&self) { + self.active.store(false, Ordering::Release); + (self.release_filter)(); + } +} + +impl Drop for PipelineRegistrationGuard { + fn drop(&mut self) { + self.close(); + } +} + +fn permanent_peer_pipeline_entry(filter: BoxPeerPacketFilter) -> Arc { + Arc::new(PeerPipelineEntry { + active: Arc::new(AtomicBool::new(true)), + filter: Arc::new(SyncRwLock::new(Some(Arc::from(filter)))), + }) +} + +fn permanent_nic_pipeline_entry(filter: BoxNicPacketFilter) -> Arc { + Arc::new(NicPipelineEntry { + active: Arc::new(AtomicBool::new(true)), + filter: Arc::new(SyncRwLock::new(Some(Arc::from(filter)))), + }) +} + +fn managed_peer_pipeline_entry( + filter: BoxPeerPacketFilter, +) -> (Arc, PipelineRegistrationGuard) { + let active = Arc::new(AtomicBool::new(true)); + let filter = Arc::new(SyncRwLock::new(Some(Arc::from(filter)))); + let release_filter = filter.clone(); + ( + Arc::new(PeerPipelineEntry { + active: active.clone(), + filter, + }), + PipelineRegistrationGuard { + active, + release_filter: Arc::new(move || { + let filter = release_filter.write().take(); + drop(filter); + }), + }, + ) +} + +#[cfg(any(feature = "proxy-packet", test))] +fn managed_nic_pipeline_entry( + filter: BoxNicPacketFilter, +) -> (Arc, PipelineRegistrationGuard) { + let active = Arc::new(AtomicBool::new(true)); + let filter = Arc::new(SyncRwLock::new(Some(Arc::from(filter)))); + let release_filter = filter.clone(); + ( + Arc::new(NicPipelineEntry { + active: active.clone(), + filter, + }), + PipelineRegistrationGuard { + active, + release_filter: Arc::new(move || { + let filter = release_filter.write().take(); + drop(filter); + }), + }, + ) +} + +#[cfg(any(feature = "proxy-packet", test))] +async fn remove_managed_nic_pipeline_entry( + pipeline: &RwLock>>, + registration: &PipelineRegistrationGuard, +) { + registration.close(); + pipeline + .write() + .await + .retain(|entry| !Arc::ptr_eq(&entry.active, ®istration.active)); +} + +async fn init_packet_process_pipeline( + peer_packet_process_pipeline: &RwLock>>, + nic_channel: PacketRecvChan, + peer_rpc_tspt_sender: UnboundedSender, +) { + // for tun/tap ip/eth packet. + peer_packet_process_pipeline + .write() + .await + .push(permanent_peer_pipeline_entry(Box::new( + NicPacketProcessor { nic_channel }, + ))); + + // for peer rpc packet + peer_packet_process_pipeline + .write() + .await + .push(permanent_peer_pipeline_entry(Box::new( + PeerRpcPacketProcessor { + peer_rpc_tspt_sender, + }, + ))); +} + +async fn add_route( + peer_packet_process_pipeline: &RwLock>>, + peers: Arc, + foreign_network_client: Arc, + foreign_network_manager: Arc, + my_peer_id: PeerId, + route: Arc, +) where + T: Route + PeerPacketFilter + Send + Sync + 'static, +{ + // for route + peer_packet_process_pipeline + .write() + .await + .push(permanent_peer_pipeline_entry(Box::new(route.clone()))); + + let _route_id = route + .open(Box::new(PeerManagerRouteInterface { + my_peer_id, + peers: Arc::downgrade(&peers), + foreign_network_client: Arc::downgrade(&foreign_network_client), + foreign_network_manager: Arc::downgrade(&foreign_network_manager), + })) + .await + .unwrap(); + + let arc_route: ArcRoute = route; + peers.add_route(arc_route).await; +} + +pub(crate) struct PeerManagerTrafficCounters { + pub self_tx_packets: CounterHandle, + pub self_tx_bytes: CounterHandle, + pub compress_tx_bytes_before: CounterHandle, + pub compress_tx_bytes_after: CounterHandle, +} + +pub struct PeerManagerCore { + my_peer_id: PeerId, + tasks: Mutex>, + packet_recv: Arc>>, + peers: Arc, + peer_rpc_mgr: Arc, + peer_rpc_tspt: Arc, + peer_packet_process_pipeline: Arc>>>, + nic_packet_process_pipeline: Arc>>>, + nic_channel: PacketRecvChan, + route_algo_inst: RouteAlgoInst, + foreign_network_client: Arc, + foreign_network_manager: Arc, + relay_peer_map: Arc, + peer_connection_admission: PeerConnectionAdmission, + outbound_packet_router: PeerOutboundPacketRouter, + recent_traffic: RecentTrafficTracker, + peer_session_store: Arc, + encryptor: Arc, + data_compress_algo: CompressorAlgo, + exit_nodes: Arc>>, + acl_filter: Arc, + context: Arc, + is_secure_mode_enabled: bool, + route: ArcRoute, + traffic_metrics: Arc, + network_name: String, + counters: PeerManagerTrafficCounters, +} + +fn check_resolved_remote_addr_not_from_virtual_network( + context: &ArcPeerContext, + resolved_remote_addr: Option, +) -> Result<(), Error> { + let Some(remote_addr) = resolved_remote_addr.map(Url::from) else { + return Ok(()); + }; + let Some(ip) = remote_addr.host().and_then(|host| match host { + url::Host::Ipv4(ip) => Some(IpAddr::V4(ip)), + url::Host::Ipv6(ip) => Some(IpAddr::V6(ip)), + url::Host::Domain(host) => host.parse().ok(), + }) else { + return Ok(()); + }; + + // If no-tun is enabled, the src ip of packet in virtual network is converted to loopback. + // TCP/QUIC/KCP proxy listeners already filter those connections. + if ip.is_loopback() || !context.is_ip_in_same_network(&ip) { + return Ok(()); + } + + Err(anyhow::anyhow!( + "tunnel src {} is from the same network (ignore this error please)", + ip + ) + .into()) +} + +impl PeerManagerCore { + #[allow(clippy::too_many_arguments)] + pub(crate) fn new( + mut config: PortablePeerManagerConfig, + runtime_config: CoreRuntimeConfigStore, + stun_info_source: Arc, + nic_channel: PacketRecvChan, + public_ipv6_runtime: Arc, + events: Arc, + credential_storage: Option>, + foreign_rpc_registrar: Arc, + ) -> anyhow::Result { + let runtime = &mut config.snapshot.runtime; + let flags = &config.snapshot.flags; + let network_name = runtime.network_identity.network_name.clone(); + if network_name.is_empty() { + anyhow::bail!("network identity name cannot be empty"); + } + match runtime.core.node.network_name.as_str() { + "" => runtime.core.node.network_name = network_name.clone(), + configured if configured != network_name => anyhow::bail!( + "core node network name {configured:?} does not match identity {network_name:?}" + ), + _ => {} + } + validate_portable_routes(&runtime.core.routes)?; + + if let (Some(_), Some(expected_digest)) = ( + runtime.network_identity.network_secret.as_ref(), + runtime.network_identity.network_secret_digest.as_ref(), + ) { + let mut identity = runtime.network_identity.clone(); + identity.network_secret_digest = None; + let derived_digest = identity + .secret_digest() + .expect("identity with a secret should derive a digest"); + if &derived_digest != expected_digest { + anyhow::bail!("network secret does not match the configured digest"); + } + } + if runtime.network_identity.network_secret.is_none() + && runtime + .network_identity + .network_secret_digest + .as_ref() + .is_some_and(|digest| digest.iter().any(|byte| *byte != 0)) + { + anyhow::bail!("digest-only local identity requires credential key capabilities"); + } + let is_secure_mode_enabled = runtime + .secure_mode + .as_ref() + .is_some_and(|secure| secure.enabled); + let is_credential_peer = runtime.network_identity.network_secret.is_none(); + if is_credential_peer && !is_secure_mode_enabled { + anyhow::bail!("credential peer identity requires secure mode and a local keypair"); + } + runtime.feature_flags.is_credential_peer = is_credential_peer; + if let Some(secure) = runtime.secure_mode.as_ref().filter(|secure| secure.enabled) { + let private_key = secure.private_key()?; + let public_key = secure.public_key()?; + let derived_public = x25519_dalek::PublicKey::from(&private_key); + if derived_public.as_bytes() != public_key.as_bytes() { + anyhow::bail!("secure mode public key does not match its private key"); + } + } + let data_compress_algo = CompressorAlgo::try_from(flags.data_compress_algo())?; + data_compress_algo.ensure_available()?; + if tokio::runtime::Handle::try_current().is_err() { + anyhow::bail!("portable peer manager construction requires an entered Tokio runtime"); + } + + // Peer IDs identify one live process incarnation. They are never + // configuration: every new PeerManager gets a fresh runtime identity. + let my_peer_id = random_peer_id(); + runtime.core.node.peer_id = Some(my_peer_id); + let instance_id = runtime + .core + .node + .instance_id + .map(uuid::Uuid::from_bytes) + .unwrap_or_else(uuid::Uuid::new_v4); + runtime.core.node.instance_id = Some(*instance_id.as_bytes()); + + let secret = runtime + .network_identity + .network_secret + .as_deref() + .unwrap_or_default(); + let encryptor: Arc = if flags.enable_encryption { + validate_algorithm(&flags.encryption_algorithm)?; + create_encryptor( + &flags.encryption_algorithm, + derive_key_128(secret), + derive_key_256(secret), + ) + } else { + Arc::new(NullCipher) + }; + runtime.feature_flags.disable_p2p = flags.disable_p2p; + runtime.feature_flags.need_p2p = flags.need_p2p; + runtime.feature_flags.avoid_relay_data |= flags.disable_relay_data; + runtime_config.update_peer(Arc::new(config.snapshot.clone())); + let public_ipv6_state = public_ipv6_runtime.clone(); + let public_ipv6_runtime: Arc = public_ipv6_runtime; + let context = Arc::new(CorePeerContext::new( + runtime_config, + public_ipv6_state, + CorePeerContextAdapters { + stun_info_source: Some(stun_info_source), + events, + credential_storage, + }, + )); + Ok(Self::assemble( + config.route_algo, + my_peer_id, + context, + public_ipv6_runtime, + nic_channel, + encryptor, + is_secure_mode_enabled, + data_compress_algo, + config.exit_nodes, + config.foreign_context_default_flags, + foreign_rpc_registrar, + )) + } + + #[allow(clippy::too_many_arguments)] + fn assemble( + route_algo: RouteAlgoType, + my_peer_id: PeerId, + core_context: Arc, + public_ipv6_runtime: Arc, + nic_channel: PacketRecvChan, + encryptor: Arc, + is_secure_mode_enabled: bool, + data_compress_algo: CompressorAlgo, + exit_nodes: Vec, + foreign_context_default_flags: FlagsInConfig, + foreign_rpc_registrar: Arc, + ) -> Self { + let stats_manager = core_context.stats_manager(); + let acl_filter = Arc::new(AclFilter::new()); + let context: ArcPeerContext = core_context.clone(); + let (packet_send, packet_recv) = super::create_packet_recv_chan(); + let peers = Arc::new(PeerMap::new( + packet_send.clone(), + context.clone(), + my_peer_id, + )); + let peer_session_store = Arc::new(PeerSessionStore::new()); + + let rpc_tspt = RpcTransport::new( + my_peer_id, + Arc::downgrade(&peers), + encryptor.clone(), + is_secure_mode_enabled, + ); + let peer_rpc_mgr = Arc::new(super::peer_rpc::PeerRpcManager::new_with_stats_manager( + rpc_tspt.clone(), + stats_manager.clone(), + )); + + let route_algo_inst = RouteAlgoInst::new( + route_algo, + my_peer_id, + context.clone(), + public_ipv6_runtime, + peer_rpc_mgr.clone(), + ); + + let foreign_network_manager = Arc::new(ForeignNetworkManager::new( + foreign_rpc_registrar, + core_context.clone(), + foreign_context_default_flags, + peer_session_store.clone(), + packet_send.clone(), + Arc::downgrade(&peers), + )); + let foreign_network_client = Arc::new(ForeignNetworkClient::new( + context.clone(), + packet_send, + my_peer_id, + )); + + let network_name = context.network_name(); + let traffic_tx_metrics = Arc::new(LogicalTrafficMetrics::new( + stats_manager.clone(), + network_name.clone(), + MetricName::TrafficBytesTx, + MetricName::TrafficPacketsTx, + MetricName::TrafficBytesTxByInstance, + MetricName::TrafficPacketsTxByInstance, + InstanceLabelKind::To, + )); + let traffic_control_tx_metrics = Arc::new(LogicalTrafficMetrics::new( + stats_manager.clone(), + network_name.clone(), + MetricName::TrafficControlBytesTx, + MetricName::TrafficControlPacketsTx, + MetricName::TrafficControlBytesTxByInstance, + MetricName::TrafficControlPacketsTxByInstance, + InstanceLabelKind::To, + )); + let self_tx_counters = PeerManagerTrafficCounters { + self_tx_packets: stats_manager.get_counter( + MetricName::TrafficPacketsSelfTx, + LabelSet::new().with_label_type(LabelType::NetworkName(network_name.clone())), + ), + self_tx_bytes: stats_manager.get_counter( + MetricName::TrafficBytesSelfTx, + LabelSet::new().with_label_type(LabelType::NetworkName(network_name.clone())), + ), + compress_tx_bytes_before: stats_manager.get_counter( + MetricName::CompressionBytesTxBefore, + LabelSet::new().with_label_type(LabelType::NetworkName(network_name.clone())), + ), + compress_tx_bytes_after: stats_manager.get_counter( + MetricName::CompressionBytesTxAfter, + LabelSet::new().with_label_type(LabelType::NetworkName(network_name.clone())), + ), + }; + let traffic_rx_metrics = Arc::new(LogicalTrafficMetrics::new( + stats_manager.clone(), + network_name.clone(), + MetricName::TrafficBytesRx, + MetricName::TrafficPacketsRx, + MetricName::TrafficBytesRxByInstance, + MetricName::TrafficPacketsRxByInstance, + InstanceLabelKind::From, + )); + let traffic_control_rx_metrics = Arc::new(LogicalTrafficMetrics::new( + stats_manager.clone(), + network_name.clone(), + MetricName::TrafficControlBytesRx, + MetricName::TrafficControlPacketsRx, + MetricName::TrafficControlBytesRxByInstance, + MetricName::TrafficControlPacketsRxByInstance, + InstanceLabelKind::From, + )); + let route_algo_inst_for_metrics = route_algo_inst.clone(); + let traffic_metrics = Arc::new(TrafficMetricRecorder::new( + my_peer_id, + traffic_tx_metrics, + traffic_control_tx_metrics, + traffic_rx_metrics, + traffic_control_rx_metrics, + move |peer_id| { + let route_algo_inst = route_algo_inst_for_metrics.clone(); + async move { + match route_algo_inst.ospf_route() { + Some(route) => route + .get_peer_info(peer_id) + .await + .as_ref() + .and_then(route_peer_info_instance_id), + None => None, + } + } + }, + )); + let peer_packet_process_pipeline = Arc::new(RwLock::new(Vec::new())); + let nic_packet_process_pipeline = Arc::new(RwLock::new(Vec::new())); + let exit_nodes = Arc::new(RwLock::new(exit_nodes)); + let relay_peer_map = super::relay_peer_map::new_relay_peer_map( + peers.clone(), + Some(foreign_network_client.clone()), + context.clone(), + my_peer_id, + peer_session_store.clone(), + ); + let recent_traffic = RecentTrafficTracker::new(my_peer_id); + let peer_connection_admission = PeerConnectionAdmission::new( + my_peer_id, + context.clone(), + peers.clone(), + foreign_network_client.clone(), + foreign_network_manager.clone(), + peer_session_store.clone(), + recent_traffic.clone(), + ); + let route = route_algo_inst.route_arc(); + let outbound_packet_router = PeerOutboundPacketRouter::new( + my_peer_id, + context.clone(), + peers.clone(), + route.clone(), + foreign_network_client.clone(), + relay_peer_map.clone(), + nic_packet_process_pipeline.clone(), + encryptor.clone(), + data_compress_algo, + exit_nodes.clone(), + recent_traffic.clone(), + traffic_metrics.clone(), + acl_filter.clone(), + is_secure_mode_enabled, + self_tx_counters.self_tx_packets.clone(), + self_tx_counters.self_tx_bytes.clone(), + self_tx_counters.compress_tx_bytes_before.clone(), + self_tx_counters.compress_tx_bytes_after.clone(), + ); + + Self { + my_peer_id, + tasks: Mutex::new(JoinSet::new()), + packet_recv: Arc::new(Mutex::new(Some(packet_recv))), + peers, + peer_rpc_mgr, + peer_rpc_tspt: rpc_tspt, + peer_packet_process_pipeline, + nic_packet_process_pipeline, + nic_channel, + route_algo_inst, + foreign_network_client, + foreign_network_manager, + relay_peer_map, + peer_connection_admission, + outbound_packet_router, + recent_traffic, + peer_session_store, + encryptor, + data_compress_algo, + exit_nodes, + acl_filter, + context: core_context, + is_secure_mode_enabled, + route, + traffic_metrics, + network_name, + counters: self_tx_counters, + } + } + + pub fn my_peer_id(&self) -> PeerId { + self.my_peer_id + } + + pub(crate) fn credential_manager(&self) -> Arc { + self.context.credential_manager() + } + + pub fn stats_manager(&self) -> Arc { + self.context.stats_manager() + } + + pub fn can_manage_credentials(&self) -> bool { + self.context.network_identity().network_secret.is_some() + } + + pub(crate) fn set_avoid_relay_data_preference(&self, avoid_relay_data: bool) { + self.context + .set_avoid_relay_data_preference(avoid_relay_data); + } + + pub fn notify_credential_changed(&self) { + self.context.issue_credential_changed(); + } + + pub async fn list_peer_snapshots(&self) -> Vec { + let foreign_peer_map = self.foreign_network_client.get_peer_map(); + let mut peers = self.peers.list_peers(); + peers.extend(foreign_peer_map.list_peers()); + + let mut snapshots = Vec::with_capacity(peers.len()); + for peer_id in peers { + let conns = if let Some(conns) = self.peers.list_peer_conns(peer_id).await { + conns + } else { + foreign_peer_map + .list_peer_conns(peer_id) + .await + .unwrap_or_default() + }; + snapshots.push(PeerSnapshot { + peer_id, + default_conn_id: self.peers.get_peer_default_conn_id(peer_id).await, + directly_connected_conns: self + .peers + .get_directly_connections_by_peer_id(peer_id) + .into_iter() + .collect(), + conns, + }); + } + snapshots + } + + pub(crate) fn instance_id(&self) -> uuid::Uuid { + self.context.instance_id() + } + + pub(crate) async fn node_snapshot(&self, listeners: Vec) -> NodeSnapshot { + NodeSnapshot { + peer_id: self.my_peer_id, + ipv4_addr: self.context.ipv4(), + proxy_networks: self.context.proxy_networks(), + hostname: self.context.hostname(), + stun_info: self.context.stun_info(), + instance_id: self.context.instance_id(), + listeners, + version: self.context.easytier_version(), + feature_flags: self.context.feature_flags(), + ip_list: GetIpListResponse::default(), + public_ipv6_addr: self.get_route().get_my_public_ipv6_addr().await, + ipv6_public_addr_prefix: self.context.advertised_ipv6_public_addr_prefix().map( + |prefix| { + cidr::Ipv6Inet::new(prefix.first_address(), prefix.network_length()).unwrap() + }, + ), + } + } + + pub async fn list_route_snapshots(&self) -> Vec { + self.get_route().list_routes().await + } + + pub async fn list_public_ipv6_routes(&self) -> BTreeSet { + self.get_route().list_public_ipv6_routes().await + } + + pub async fn public_ipv6_addr(&self) -> Option { + self.get_route().get_my_public_ipv6_addr().await + } + + pub async fn dump_route(&self) -> String { + self.get_route().dump().await + } + + pub async fn local_public_ipv6_info(&self) -> ListPublicIpv6InfoResponse { + self.get_route().get_local_public_ipv6_info().await + } + + pub async fn foreign_network_route_infos(&self) -> RouteForeignNetworkInfos { + self.get_route().list_foreign_network_info().await + } + + pub async fn list_foreign_network_infos( + &self, + include_trusted_keys: bool, + ) -> std::collections::HashMap { + self.foreign_network_manager + .list_foreign_network_infos(include_trusted_keys) + .await + } + + pub async fn foreign_network_route_summary(&self) -> RouteForeignNetworkSummary { + self.get_route().get_foreign_network_summary().await + } + + pub fn acl_stats(&self) -> crate::proto::acl::AclStats { + self.acl_filter.get_stats() + } + + pub fn acl_filter(&self) -> Arc { + self.acl_filter.clone() + } + + pub fn network_name(&self) -> &str { + &self.network_name + } + + pub(crate) fn dns_route_identity( + &self, + ) -> (String, Option, String) { + ( + self.context.hostname(), + self.context.ipv4().map(Into::into), + self.context.flags().tld_dns_zone, + ) + } + + pub fn p2p_policy_flags(&self) -> P2pPolicyFlags { + let flags = self.context.flags(); + P2pPolicyFlags { + disable_udp_hole_punching: flags.disable_udp_hole_punching, + disable_sym_hole_punching: flags.disable_sym_hole_punching, + disable_upnp: flags.disable_upnp, + lazy_p2p: flags.lazy_p2p, + disable_p2p: flags.disable_p2p, + need_p2p: flags.need_p2p, + } + } + + pub fn tcp_hole_punching_disabled(&self) -> bool { + self.context.flags().disable_tcp_hole_punching + } + + pub fn get_peer_map(&self) -> Arc { + self.peers.clone() + } + + pub fn get_relay_peer_map(&self) -> Arc { + self.relay_peer_map.clone() + } + + pub fn get_peer_rpc_mgr(&self) -> Arc { + self.peer_rpc_mgr.clone() + } + + pub fn get_peer_session_store(&self) -> Arc { + self.peer_session_store.clone() + } + + pub fn get_nic_channel(&self) -> PacketRecvChan { + self.nic_channel.clone() + } + + pub(crate) fn is_local_virtual_ip(&self, ip: &IpAddr) -> bool { + self.context.is_ip_local_virtual_ip(ip) + } + + pub fn get_foreign_network_client(&self) -> Arc { + self.foreign_network_client.clone() + } + + pub async fn is_easytier_managed_ipv6(&self, ip: &std::net::Ipv6Addr) -> bool { + if self.context.is_ip_local_ipv6(ip) { + return true; + } + self.route + .list_public_ipv6_routes() + .await + .iter() + .any(|route| route.address() == *ip) + } + + pub fn traffic_metrics(&self) -> Arc { + self.traffic_metrics.clone() + } + + pub fn get_route(&self) -> ArcRoute { + self.route.clone() + } + + pub fn mark_recent_traffic(&self, dst_peer_id: PeerId) { + let flags = self.context.flags(); + self.recent_traffic + .mark(dst_peer_id, flags.disable_p2p, flags.lazy_p2p, |peer_id| { + self.has_directly_connected_conn(peer_id) + }); + } + + pub fn has_recent_traffic(&self, peer_id: PeerId, now: Instant) -> bool { + self.recent_traffic.has(peer_id, now, |peer_id| { + self.has_directly_connected_conn(peer_id) + }) + } + + pub fn clear_recent_traffic(&self, peer_id: PeerId) { + self.recent_traffic.clear(peer_id); + } + + pub fn p2p_demand_notify(&self) -> Arc { + self.recent_traffic.p2p_demand_notify() + } + + pub fn gc_recent_traffic(&self) { + self.recent_traffic.gc(Instant::now(), |peer_id| { + self.has_directly_connected_conn(peer_id) + }); + } + + pub fn has_directly_connected_conn(&self, peer_id: PeerId) -> bool { + if let Some(peer) = self.peers.get_peer_by_id(peer_id) { + peer.has_directly_connected_conn() + } else { + self.foreign_network_client.get_peer_map().has_peer(peer_id) + } + } + + pub async fn add_client_tunnel( + &self, + tunnel: Box, + is_directly_connected: bool, + ) -> Result<(PeerId, PeerConnId), Error> { + self.peer_connection_admission + .add_client_tunnel(tunnel, is_directly_connected) + .await + } + + pub async fn add_client_tunnel_with_peer_id_hint( + &self, + tunnel: Box, + is_directly_connected: bool, + peer_id_hint: Option, + ) -> Result<(PeerId, PeerConnId), Error> { + self.peer_connection_admission + .add_client_tunnel_with_peer_id_hint(tunnel, is_directly_connected, peer_id_hint) + .await + } + + pub async fn add_tunnel_as_server( + &self, + tunnel: Box, + is_directly_connected: bool, + ) -> Result<(), Error> { + self.peer_connection_admission + .add_tunnel_as_server(tunnel, is_directly_connected) + .await + } + + pub async fn add_packet_process_pipeline(&self, pipeline: BoxPeerPacketFilter) { + // newest pipeline will be executed first + self.peer_packet_process_pipeline + .write() + .await + .push(permanent_peer_pipeline_entry(pipeline)); + } + + pub async fn add_nic_packet_process_pipeline(&self, pipeline: BoxNicPacketFilter) { + // newest pipeline will be executed first + self.nic_packet_process_pipeline + .write() + .await + .push(permanent_nic_pipeline_entry(pipeline)); + } + + pub(crate) async fn add_managed_packet_process_pipeline( + &self, + pipeline: BoxPeerPacketFilter, + ) -> PipelineRegistrationGuard { + let (entry, guard) = managed_peer_pipeline_entry(pipeline); + let mut pipelines = self.peer_packet_process_pipeline.write().await; + pipelines.retain(|pipeline| pipeline.active.load(Ordering::Acquire)); + pipelines.push(entry); + guard + } + + #[cfg(feature = "proxy-packet")] + pub(crate) async fn add_managed_nic_packet_process_pipeline( + &self, + pipeline: BoxNicPacketFilter, + ) -> PipelineRegistrationGuard { + let (entry, guard) = managed_nic_pipeline_entry(pipeline); + let mut pipelines = self.nic_packet_process_pipeline.write().await; + pipelines.retain(|pipeline| pipeline.active.load(Ordering::Acquire)); + pipelines.push(entry); + guard + } + + #[cfg(feature = "proxy-packet")] + pub(crate) async fn remove_managed_nic_packet_process_pipeline( + &self, + registration: &PipelineRegistrationGuard, + ) { + remove_managed_nic_pipeline_entry(&self.nic_packet_process_pipeline, registration).await; + } + + pub async fn add_route(&self, route: Arc) + where + T: Route + PeerPacketFilter + Send + Sync + 'static, + { + add_route( + self.peer_packet_process_pipeline.as_ref(), + self.peers.clone(), + self.foreign_network_client.clone(), + self.foreign_network_manager.clone(), + self.my_peer_id, + route, + ) + .await; + } + + pub async fn remove_nic_packet_process_pipeline(&self, id: String) -> Result<(), Error> { + let mut pipelines = self.nic_packet_process_pipeline.write().await; + if let Some(pos) = pipelines.iter().position(|pipeline| { + let filter = pipeline.filter.read().clone(); + filter.is_some_and(|filter| filter.id() == id) + }) { + pipelines.remove(pos); + Ok(()) + } else { + Err(Error::NotFound) + } + } + + pub async fn send_msg_for_proxy( + &self, + msg: ZCPacket, + dst_peer_id: PeerId, + ) -> Result<(), Error> { + self.outbound_packet_router + .send_msg_for_proxy(msg, dst_peer_id) + .await + } + + pub async fn get_msg_dst_peer(&self, addr: &IpAddr) -> (Vec, bool) { + self.outbound_packet_router.get_msg_dst_peer(addr).await + } + + pub async fn get_msg_dst_peer_ipv4(&self, ipv4_addr: &Ipv4Addr) -> (Vec, bool) { + self.outbound_packet_router + .get_msg_dst_peer_ipv4(ipv4_addr) + .await + } + + pub async fn get_msg_dst_peer_ipv6(&self, ipv6_addr: &Ipv6Addr) -> (Vec, bool) { + self.outbound_packet_router + .get_msg_dst_peer_ipv6(ipv6_addr) + .await + } + + pub async fn send_msg_by_ip( + &self, + msg: ZCPacket, + ip_addr: IpAddr, + not_send_to_self: bool, + ) -> Result<(), Error> { + self.outbound_packet_router + .send_msg_by_ip(msg, ip_addr, not_send_to_self) + .await + } + + pub async fn check_allow_kcp_to_dst(&self, dst_ip: &IpAddr) -> bool { + self.outbound_packet_router + .check_allow_kcp_to_dst(dst_ip) + .await + } + + pub async fn check_allow_quic_to_dst(&self, dst_ip: &IpAddr) -> bool { + self.outbound_packet_router + .check_allow_quic_to_dst(dst_ip) + .await + } + + pub async fn update_exit_nodes(&self, exit_nodes: Vec) { + *self.exit_nodes.write().await = exit_nodes; + } + + pub(crate) fn reload_acl(&self, acl: Option<&crate::proto::acl::Acl>) { + // ACL rule effects are staged separately from configuration publication. + // Keep the submitted group snapshot unchanged so CoreInstance can detect + // the group change and refresh route trust state when the complete + // runtime configuration is published. + self.acl_filter.reload_rules(acl); + } + + pub async fn wait(&self) { + while !self.tasks.lock().await.is_empty() { + crate::foundation::time::sleep(std::time::Duration::from_secs(1)).await; + } + } + + pub(crate) async fn stop(&self) { + let mut tasks = { + let mut task_slot = self.tasks.lock().await; + std::mem::replace(&mut *task_slot, JoinSet::new()) + }; + tasks.abort_all(); + while tasks.join_next().await.is_some() {} + self.foreign_network_manager.stop().await; + self.route.close().await; + self.peer_rpc_mgr.stop().await; + self.context.stop().await; + self.context.stats_manager().stop_cleanup_task().await; + self.acl_filter.stop_cleanup_task().await; + } + + pub(crate) async fn clear_resources(&self) { + self.stop().await; + let mut peer_pipeline = self.peer_packet_process_pipeline.write().await; + peer_pipeline.clear(); + let mut nic_pipeline = self.nic_packet_process_pipeline.write().await; + nic_pipeline.clear(); + + self.peer_rpc_mgr.rpc_server().registry().unregister_all(); + } + + pub async fn close_peer_conn( + &self, + peer_id: PeerId, + conn_id: &PeerConnId, + ) -> Result<(), Error> { + close_peer_conn( + self.peers.as_ref(), + &self.foreign_network_client, + &self.foreign_network_manager, + peer_id, + conn_id, + ) + .await + } + + async fn start_peer_recv(&self) { + let packet_recv = self.packet_recv.lock().await.take().unwrap(); + let is_credential_node = + self.context.network_identity().network_secret.is_none() && self.is_secure_mode_enabled; + let router = PeerPacketRouter::new( + packet_recv, + self.my_peer_id, + self.peers.clone(), + self.peer_packet_process_pipeline.clone(), + self.foreign_network_client.clone(), + self.relay_peer_map.clone(), + self.foreign_network_manager.clone(), + self.encryptor.clone(), + self.data_compress_algo, + self.acl_filter.clone(), + self.context.clone(), + self.is_secure_mode_enabled, + self.route.clone(), + is_credential_node, + self.traffic_metrics.clone(), + self.context.stats_manager(), + self.network_name.clone(), + self.counters.self_tx_packets.clone(), + self.counters.self_tx_bytes.clone(), + self.counters.compress_tx_bytes_before.clone(), + self.counters.compress_tx_bytes_after.clone(), + ); + + self.tasks.lock().await.spawn(router.run()); + } + + pub(crate) async fn run(&self) -> Result<(), Error> { + self.context.stats_manager().start_cleanup_task(); + + if let Some(route) = self.route_algo_inst.ospf_route() { + self.add_route(route).await; + } + + init_packet_process_pipeline( + self.peer_packet_process_pipeline.as_ref(), + self.nic_channel.clone(), + self.peer_rpc_tspt.packet_sender(), + ) + .await; + self.peer_rpc_mgr.run(); + + self.start_peer_recv().await; + PeerMaintenanceTasks::new( + self.peers.clone(), + self.relay_peer_map.clone(), + self.recent_traffic.clone(), + self.foreign_network_client.clone(), + self.peer_session_store.clone(), + self.context.clone(), + self.traffic_metrics.clone(), + ) + .spawn_into(&self.tasks) + .await; + + self.foreign_network_client.run().await; + + Ok(()) + } +} + +#[async_trait::async_trait] +impl PeerCenterPeerManagerTrait for PeerManagerCore { + async fn list_peers(&self) -> PeerInfoForGlobalMap { + direct_peer_info(&[ + self.get_peer_map(), + self.get_foreign_network_client().get_peer_map(), + ]) + .await + } + + fn my_peer_id(&self) -> PeerId { + self.my_peer_id() + } + + fn network_name(&self) -> String { + self.network_name().to_owned() + } + + fn get_rpc_mgr(&self) -> Weak { + Arc::downgrade(&self.get_peer_rpc_mgr()) + } + + async fn list_routes(&self) -> Vec { + self.get_route().list_routes().await + } +} + +pub(crate) async fn close_peer_conn( + peers: &PeerMap, + foreign_network_client: &ForeignNetworkClient, + foreign_network_manager: &ForeignNetworkManager, + peer_id: PeerId, + conn_id: &PeerConnId, +) -> Result<(), Error> { + let ret = peers.close_peer_conn(peer_id, conn_id).await; + tracing::info!("close_peer_conn in peer map: {:?}", ret); + if ret.is_ok() || !matches!(ret.as_ref().unwrap_err(), Error::NotFound) { + return ret; + } + + let ret = foreign_network_client + .get_peer_map() + .close_peer_conn(peer_id, conn_id) + .await; + tracing::info!("close_peer_conn in foreign network client: {:?}", ret); + if ret.is_ok() || !matches!(ret.as_ref().unwrap_err(), Error::NotFound) { + return ret; + } + + let ret = foreign_network_manager + .close_peer_conn(peer_id, conn_id) + .await; + tracing::info!("close_peer_conn in foreign network manager done: {:?}", ret); + ret +} + +pub(crate) struct PeerConnectionAdmission { + my_peer_id: PeerId, + context: ArcPeerContext, + peers: Arc, + foreign_network_client: Arc, + foreign_network_manager: Arc, + peer_session_store: Arc, + recent_traffic: RecentTrafficTracker, + reserved_my_peer_id_map: DashMap, +} + +impl PeerConnectionAdmission { + fn new( + my_peer_id: PeerId, + context: ArcPeerContext, + peers: Arc, + foreign_network_client: Arc, + foreign_network_manager: Arc, + peer_session_store: Arc, + recent_traffic: RecentTrafficTracker, + ) -> Self { + Self { + my_peer_id, + context, + peers, + foreign_network_client, + foreign_network_manager, + peer_session_store, + recent_traffic, + reserved_my_peer_id_map: DashMap::new(), + } + } + + pub async fn add_client_tunnel( + &self, + tunnel: Box, + is_directly_connected: bool, + ) -> Result<(PeerId, PeerConnId), Error> { + self.add_client_tunnel_with_peer_id_hint(tunnel, is_directly_connected, None) + .await + } + + pub async fn add_client_tunnel_with_peer_id_hint( + &self, + tunnel: Box, + is_directly_connected: bool, + peer_id_hint: Option, + ) -> Result<(PeerId, PeerConnId), Error> { + let mut peer = PeerConn::new_with_peer_id_hint( + self.my_peer_id, + self.context.clone(), + tunnel, + peer_id_hint, + self.peer_session_store.clone(), + ); + peer.set_is_hole_punched(!is_directly_connected); + peer.do_handshake_as_client().await?; + let conn_id = peer.get_conn_id(); + let peer_id = peer.get_peer_id(); + let local_identity = self.context.network_identity(); + if peer.get_network_identity().network_name == local_identity.network_name { + let local_secure_mode = self + .context + .secure_mode() + .as_ref() + .map(|cfg| cfg.enabled) + .unwrap_or(false); + let peer_id = add_new_peer_conn( + self.peers.as_ref(), + &local_identity, + local_secure_mode, + peer, + ) + .await?; + self.recent_traffic.clear(peer_id); + } else { + self.foreign_network_client.add_new_peer_conn(peer).await?; + } + Ok((peer_id, conn_id)) + } + + fn release_reserved_peer_id(&self, network_name: &str) { + self.reserved_my_peer_id_map.remove(network_name); + shrink_dashmap(&self.reserved_my_peer_id_map, None); + } + + #[tracing::instrument(ret, skip(self, tunnel))] + pub async fn add_tunnel_as_server( + &self, + tunnel: Box, + is_directly_connected: bool, + ) -> Result<(), Error> { + tracing::info!("add tunnel as server start"); + let resolved_remote_addr = tunnel.info().and_then(|info| info.resolved_remote_addr); + check_resolved_remote_addr_not_from_virtual_network(&self.context, resolved_remote_addr)?; + + let mut conn = PeerConn::new( + self.my_peer_id, + self.context.clone(), + tunnel, + self.peer_session_store.clone(), + ); + let mut reserved_peer_id_network_name = None; + let handshake_ret = conn + .do_handshake_as_server_ext(|peer, network_name: &str| { + if network_name == self.context.network_identity().network_name { + return Ok(()); + } + + let mut peer_id = self + .foreign_network_manager + .get_network_peer_id(network_name); + if peer_id.is_none() { + reserved_peer_id_network_name = Some(network_name.to_string()); + peer_id = Some( + *self + .reserved_my_peer_id_map + .entry(network_name.to_string()) + .or_insert_with(random_peer_id) + .value(), + ); + } + peer.set_peer_id(peer_id.unwrap()); + + tracing::info!( + ?peer_id, + ?network_name, + "handshake as server with foreign network, new peer id: {}, peer id in foreign manager: {:?}", + peer.get_my_peer_id(), peer_id + ); + + Ok(()) + }) + .await; + + if let Err(err) = handshake_ret { + if let Some(network_name) = reserved_peer_id_network_name { + self.release_reserved_peer_id(&network_name); + } + return Err(err); + } + + let peer_identity = conn.get_network_identity(); + let peer_network_name = peer_identity.network_name.clone(); + let local_identity = self.context.network_identity(); + let is_local_network = peer_network_name == local_identity.network_name; + let trusted_foreign_credential = + matches!(conn.get_peer_identity_type(), PeerIdentityType::Credential) + && self + .foreign_network_manager + .is_existing_credential_pubkey_trusted( + &peer_network_name, + &conn.get_conn_info().noise_remote_static_pubkey, + ); + let foreign_network_allowed = + conn.matches_local_network_secret() || trusted_foreign_credential; + + if !is_local_network && self.context.flags().private_mode && !foreign_network_allowed { + self.release_reserved_peer_id(&peer_network_name); + return Err(Error::SecretKeyError( + "private mode is turned on, foreign network secret mismatch".to_string(), + )); + } + + conn.set_is_hole_punched(!is_directly_connected); + + let add_peer_ret = if is_local_network { + let local_secure_mode = self + .context + .secure_mode() + .as_ref() + .map(|cfg| cfg.enabled) + .unwrap_or(false); + match add_new_peer_conn( + self.peers.as_ref(), + &local_identity, + local_secure_mode, + conn, + ) + .await + { + Ok(peer_id) => { + self.recent_traffic.clear(peer_id); + Ok(()) + } + Err(err) => Err(err), + } + } else { + self.foreign_network_manager.add_peer_conn(conn).await + }; + + if let Err(err) = add_peer_ret { + self.release_reserved_peer_id(&peer_network_name); + return Err(err); + } + + self.release_reserved_peer_id(&peer_network_name); + + tracing::info!("add tunnel as server done"); + Ok(()) + } +} + +// Keep lazy-p2p demand alive across the 5s task rescan interval and a full on-demand +// connect attempt, without retaining extra per-task state in the hot path. +pub(crate) const RECENT_HAVE_TRAFFIC_TTL: Duration = Duration::from_secs(30); + +pub(crate) fn should_mark_recent_traffic_for_fanout(total_dst_peers: usize) -> bool { + total_dst_peers <= 1 +} + +fn gc_recent_traffic_entries( + recent_have_traffic: &DashMap, + now: Instant, + mut has_directly_connected_conn: F, +) where + F: FnMut(PeerId) -> bool, +{ + let mut to_remove = Vec::new(); + for entry in recent_have_traffic.iter() { + let peer_id = *entry.key(); + let expired = now.saturating_duration_since(*entry.value()) > RECENT_HAVE_TRAFFIC_TTL; + if expired || has_directly_connected_conn(peer_id) { + to_remove.push(peer_id); + } + } + + if !to_remove.is_empty() { + for peer_id in to_remove { + recent_have_traffic.remove(&peer_id); + } + shrink_dashmap(recent_have_traffic, None); + } +} + +#[derive(Clone)] +pub(crate) struct RecentTrafficTracker { + my_peer_id: PeerId, + recent_have_traffic: Arc>, + p2p_demand_notify: Arc, +} + +impl RecentTrafficTracker { + pub fn new(my_peer_id: PeerId) -> Self { + Self { + my_peer_id, + recent_have_traffic: Arc::new(DashMap::new()), + p2p_demand_notify: Arc::new(ExternalTaskSignal::new()), + } + } + + pub fn mark( + &self, + dst_peer_id: PeerId, + disable_p2p: bool, + lazy_p2p: bool, + mut has_directly_connected_conn: F, + ) where + F: FnMut(PeerId) -> bool, + { + if dst_peer_id == self.my_peer_id { + return; + } + + if disable_p2p || !lazy_p2p || has_directly_connected_conn(dst_peer_id) { + return; + } + + let now = Instant::now(); + if let Some(mut last_seen) = self.recent_have_traffic.get_mut(&dst_peer_id) { + let should_notify = now.saturating_duration_since(*last_seen) > RECENT_HAVE_TRAFFIC_TTL; + *last_seen = now; + if !should_notify { + return; + } + } else { + self.recent_have_traffic.insert(dst_peer_id, now); + } + self.p2p_demand_notify.notify(); + } + + pub fn has(&self, peer_id: PeerId, now: Instant, mut has_directly_connected_conn: F) -> bool + where + F: FnMut(PeerId) -> bool, + { + if has_directly_connected_conn(peer_id) { + return false; + } + + self.recent_have_traffic + .get(&peer_id) + .map(|last_seen| now.saturating_duration_since(*last_seen) <= RECENT_HAVE_TRAFFIC_TTL) + .unwrap_or(false) + } + + pub fn clear(&self, peer_id: PeerId) { + self.recent_have_traffic.remove(&peer_id); + } + + pub fn gc(&self, now: Instant, has_directly_connected_conn: F) + where + F: FnMut(PeerId) -> bool, + { + gc_recent_traffic_entries( + self.recent_have_traffic.as_ref(), + now, + has_directly_connected_conn, + ); + } + + pub fn p2p_demand_notify(&self) -> Arc { + self.p2p_demand_notify.clone() + } +} + +pub(crate) struct PeerMaintenanceTasks { + peer_map: Arc, + relay_peer_map: Arc, + recent_traffic: RecentTrafficTracker, + foreign_network_client: Arc, + peer_session_store: Arc, + context: ArcPeerContext, + traffic_metrics: Arc, +} + +impl PeerMaintenanceTasks { + pub fn new( + peer_map: Arc, + relay_peer_map: Arc, + recent_traffic: RecentTrafficTracker, + foreign_network_client: Arc, + peer_session_store: Arc, + context: ArcPeerContext, + traffic_metrics: Arc, + ) -> Self { + Self { + peer_map, + relay_peer_map, + recent_traffic, + foreign_network_client, + peer_session_store, + context, + traffic_metrics, + } + } + + pub async fn spawn_into(self, tasks: &Mutex>) { + self.spawn_clean_peer_without_conn_routine(tasks).await; + self.spawn_relay_session_gc_routine(tasks).await; + self.spawn_recent_traffic_gc_routine(tasks).await; + self.spawn_peer_session_gc_routine(tasks).await; + self.spawn_credential_gc_routine(tasks).await; + self.spawn_traffic_metrics_gc_routine(tasks).await; + } + + async fn spawn_clean_peer_without_conn_routine(&self, tasks: &Mutex>) { + let peer_map = self.peer_map.clone(); + tasks.lock().await.spawn(async move { + loop { + peer_map.clean_peer_without_conn().await; + crate::foundation::time::sleep(std::time::Duration::from_secs(3)).await; + } + }); + } + + async fn spawn_relay_session_gc_routine(&self, tasks: &Mutex>) { + let relay_peer_map = self.relay_peer_map.clone(); + tasks.lock().await.spawn(async move { + loop { + relay_peer_map.evict_idle_sessions(std::time::Duration::from_secs(60)); + crate::foundation::time::sleep(std::time::Duration::from_secs(30)).await; + } + }); + } + + async fn spawn_recent_traffic_gc_routine(&self, tasks: &Mutex>) { + let recent_traffic = self.recent_traffic.clone(); + let peers = self.peer_map.clone(); + let foreign_network_client = self.foreign_network_client.clone(); + tasks.lock().await.spawn(async move { + loop { + recent_traffic.gc(Instant::now(), |peer_id| { + if let Some(peer) = peers.get_peer_by_id(peer_id) { + peer.has_directly_connected_conn() + } else { + foreign_network_client.get_peer_map().has_peer(peer_id) + } + }); + crate::foundation::time::sleep(std::time::Duration::from_secs(30)).await; + } + }); + } + + async fn spawn_peer_session_gc_routine(&self, tasks: &Mutex>) { + let peer_session_store = self.peer_session_store.clone(); + tasks.lock().await.spawn(async move { + loop { + crate::foundation::time::sleep(std::time::Duration::from_secs(60)).await; + peer_session_store.evict_unused_sessions(); + } + }); + } + + async fn spawn_credential_gc_routine(&self, tasks: &Mutex>) { + let context = self.context.clone(); + let peer_map = self.peer_map.clone(); + tasks.lock().await.spawn(async move { + loop { + if context.network_identity().network_secret.is_some() { + if context.remove_expired_credentials() { + context.issue_credential_changed(); + } + + let network_name = context.network_name(); + close_untrusted_credential_peers( + peer_map.as_ref(), + &network_name, + |pubkey, network_name| context.is_pubkey_trusted(pubkey, network_name), + ) + .await; + } + crate::foundation::time::sleep(std::time::Duration::from_secs(1)).await; + } + }); + } + + async fn spawn_traffic_metrics_gc_routine(&self, tasks: &Mutex>) { + let Some(mut event_receiver) = self.context.subscribe_peer_events() else { + return; + }; + let context = self.context.clone(); + let traffic_metrics = self.traffic_metrics.clone(); + tasks.lock().await.spawn(async move { + loop { + match event_receiver.recv().await { + Ok(super::context::PeerContextEvent::PeerRemoved(peer_id)) => { + traffic_metrics.remove_peer(peer_id); + } + Ok(_) => {} + Err(tokio::sync::broadcast::error::RecvError::Lagged(skipped)) => { + tracing::warn!( + skipped, + "traffic metrics GC receiver lagged; clearing peer cache to avoid stale metric attribution" + ); + traffic_metrics.clear_peer_cache(); + let Some(new_receiver) = context.subscribe_peer_events() else { + break; + }; + event_receiver = new_receiver; + } + Err(tokio::sync::broadcast::error::RecvError::Closed) => break, + } + } + }); + } +} + +pub(crate) async fn try_compress_and_encrypt( + compress_algo: CompressorAlgo, + encryptor: &Arc, + msg: &mut ZCPacket, + secure_mode_enabled: bool, +) -> Result<(), Error> { + let compressor = DefaultCompressor {}; + compressor + .compress(msg, compress_algo) + .await + .with_context(|| "compress failed")?; + if !secure_mode_enabled { + encryptor.encrypt(msg).with_context(|| "encrypt failed")?; + } + Ok(()) +} + +struct PeerOutboundPacketRouterCounters { + self_tx_packets: CounterHandle, + self_tx_bytes: CounterHandle, + compress_tx_bytes_before: CounterHandle, + compress_tx_bytes_after: CounterHandle, +} + +pub(crate) struct PeerOutboundPacketRouter { + my_peer_id: PeerId, + context: ArcPeerContext, + host_routing: HostRoutingPolicy, + peers: Arc, + route: ArcRoute, + foreign_network_client: Arc, + relay_peer_map: Arc, + nic_packet_process_pipeline: Arc>>>, + encryptor: Arc, + data_compress_algo: CompressorAlgo, + exit_nodes: Arc>>, + recent_traffic: RecentTrafficTracker, + traffic_metrics: Arc, + acl_filter: Arc, + is_secure_mode_enabled: bool, + counters: PeerOutboundPacketRouterCounters, +} + +impl PeerOutboundPacketRouter { + #[allow(clippy::too_many_arguments)] + fn new( + my_peer_id: PeerId, + context: ArcPeerContext, + peers: Arc, + route: ArcRoute, + foreign_network_client: Arc, + relay_peer_map: Arc, + nic_packet_process_pipeline: Arc>>>, + encryptor: Arc, + data_compress_algo: CompressorAlgo, + exit_nodes: Arc>>, + recent_traffic: RecentTrafficTracker, + traffic_metrics: Arc, + acl_filter: Arc, + is_secure_mode_enabled: bool, + self_tx_packets: CounterHandle, + self_tx_bytes: CounterHandle, + compress_tx_bytes_before: CounterHandle, + compress_tx_bytes_after: CounterHandle, + ) -> Self { + let host_routing = context.host_routing_policy(); + Self { + my_peer_id, + context, + host_routing, + peers, + route, + foreign_network_client, + relay_peer_map, + nic_packet_process_pipeline, + encryptor, + data_compress_algo, + exit_nodes, + recent_traffic, + traffic_metrics, + acl_filter, + is_secure_mode_enabled, + counters: PeerOutboundPacketRouterCounters { + self_tx_packets, + self_tx_bytes, + compress_tx_bytes_before, + compress_tx_bytes_after, + }, + } + } + + fn has_directly_connected_conn(&self, peer_id: PeerId) -> bool { + if let Some(peer) = self.peers.get_peer_by_id(peer_id) { + peer.has_directly_connected_conn() + } else { + self.foreign_network_client.get_peer_map().has_peer(peer_id) + } + } + + fn mark_recent_traffic(&self, dst_peer_id: PeerId) { + let flags = self.context.flags(); + self.recent_traffic + .mark(dst_peer_id, flags.disable_p2p, flags.lazy_p2p, |peer_id| { + self.has_directly_connected_conn(peer_id) + }); + } + + async fn run_nic_packet_process_pipeline(&self, data: &mut ZCPacket) -> bool { + // Enforce ACL for outbound (NIC-originated) packets. If ACL denies, stop processing. + if !self.acl_filter.process_packet_with_acl( + data, + false, + None, + |_| false, + self.route.as_ref(), + ) { + return false; + } + + for pipeline in self.nic_packet_process_pipeline.read().await.iter().rev() { + if !pipeline.active.load(Ordering::Acquire) { + continue; + } + let filter = pipeline.filter.read().clone(); + if let Some(filter) = filter { + let _ = filter.try_process_packet_from_nic(data).await; + } + } + + true + } + + fn check_p2p_only_before_send(&self, dst_peer_id: PeerId) -> Result<(), Error> { + if self.context.p2p_only() && !self.peers.has_peer(dst_peer_id) { + return Err(Error::RouteError(None)); + } + Ok(()) + } + + async fn check_allow_wrapped_proxy_to_dst( + &self, + dst_ip: &IpAddr, + dst_allows_input: impl Fn(crate::proto::common::PeerFeatureFlag) -> bool, + next_hop_disables_relay: impl Fn(crate::proto::common::PeerFeatureFlag) -> bool, + ) -> bool { + let Some(dst_peer_id) = self.route.get_peer_id_by_ip(dst_ip).await else { + return false; + }; + let Some(peer_info) = self.route.get_peer_info(dst_peer_id).await else { + return false; + }; + + if !peer_info + .feature_flag + .map(dst_allows_input) + .unwrap_or(false) + { + return false; + } + + let next_hop_policy = get_next_hop_policy(self.context.flags().latency_first); + let Some(next_hop_id) = self + .route + .get_next_hop_with_policy(dst_peer_id, next_hop_policy) + .await + else { + return false; + }; + + if next_hop_id == dst_peer_id { + return true; + } + + let Some(next_hop_info) = self.route.get_peer_info(next_hop_id).await else { + return false; + }; + + !next_hop_info + .feature_flag + .map(next_hop_disables_relay) + .unwrap_or(false) + } + + pub async fn check_allow_kcp_to_dst(&self, dst_ip: &IpAddr) -> bool { + self.check_allow_wrapped_proxy_to_dst( + dst_ip, + |feature_flag| feature_flag.kcp_input, + |feature_flag| feature_flag.no_relay_kcp, + ) + .await + } + + pub async fn check_allow_quic_to_dst(&self, dst_ip: &IpAddr) -> bool { + self.check_allow_wrapped_proxy_to_dst( + dst_ip, + |feature_flag| feature_flag.quic_input, + |feature_flag| feature_flag.no_relay_quic, + ) + .await + } + + pub async fn send_msg_for_proxy( + &self, + mut msg: ZCPacket, + dst_peer_id: PeerId, + ) -> Result<(), Error> { + self.mark_recent_traffic(dst_peer_id); + self.check_p2p_only_before_send(dst_peer_id)?; + + self.counters + .compress_tx_bytes_before + .add(msg.buf_len() as u64); + + try_compress_and_encrypt( + self.data_compress_algo, + &self.encryptor, + &mut msg, + self.is_secure_mode_enabled, + ) + .await?; + + self.counters + .compress_tx_bytes_after + .add(msg.buf_len() as u64); + + let msg_len = msg.buf_len() as u64; + let result = send_msg_internal( + self.peers.as_ref(), + &self.foreign_network_client, + &self.relay_peer_map, + Some(&self.traffic_metrics), + msg, + dst_peer_id, + ) + .await; + if result.is_ok() { + self.counters.self_tx_bytes.add(msg_len); + self.counters.self_tx_packets.inc(); + } + result + } + + pub async fn get_msg_dst_peer(&self, addr: &IpAddr) -> (Vec, bool) { + match addr { + IpAddr::V4(ipv4_addr) => self.get_msg_dst_peer_ipv4(ipv4_addr).await, + IpAddr::V6(ipv6_addr) => self.get_msg_dst_peer_ipv6(ipv6_addr).await, + } + } + + fn is_all_peers_broadcast_ipv4(&self, ipv4_addr: &Ipv4Addr) -> bool { + let network_length = self + .context + .ipv4() + .map(|x| x.network_length()) + .unwrap_or(24); + let ipv4_inet = cidr::Ipv4Inet::new(*ipv4_addr, network_length).unwrap(); + ipv4_addr.is_broadcast() + || ipv4_addr.is_multicast() + || *ipv4_addr == ipv4_inet.last_address() + } + + fn is_all_peers_broadcast_ipv6(&self, ipv6_addr: &Ipv6Addr) -> bool { + let network_length = self + .context + .ipv6() + .map(|x| x.network_length()) + .unwrap_or(64); + let ipv6_inet = cidr::Ipv6Inet::new(*ipv6_addr, network_length).unwrap(); + ipv6_addr.is_multicast() || *ipv6_addr == ipv6_inet.last_address() + } + + fn select_ipv4_broadcast_peers<'a>( + routes: impl IntoIterator, + my_peer_id: PeerId, + ) -> Vec { + routes + .into_iter() + .filter_map(|route| { + (route.peer_id != my_peer_id && route.ipv4_addr.is_some()).then_some(route.peer_id) + }) + .collect() + } + + pub async fn get_msg_dst_peer_ipv4(&self, ipv4_addr: &Ipv4Addr) -> (Vec, bool) { + let mut is_exit_node = false; + let mut dst_peers = vec![]; + if self.is_all_peers_broadcast_ipv4(ipv4_addr) { + dst_peers.extend(Self::select_ipv4_broadcast_peers( + &self.peers.list_route_infos().await, + self.my_peer_id, + )); + } else if let Some(peer_id) = self.peers.get_peer_id_by_ipv4(ipv4_addr).await { + dst_peers.push(peer_id); + } else if !self + .context + .is_ip_in_same_network(&std::net::IpAddr::V4(*ipv4_addr)) + { + for exit_node in self.exit_nodes.read().await.iter() { + let IpAddr::V4(exit_node) = exit_node else { + continue; + }; + if let Some(peer_id) = self.peers.get_peer_id_by_ipv4(exit_node).await { + dst_peers.push(peer_id); + is_exit_node = true; + break; + } + } + } + if self.host_routing.local_exit_node_fallback + && dst_peers.is_empty() + && !self + .context + .is_ip_in_same_network(&std::net::IpAddr::V4(*ipv4_addr)) + { + tracing::trace!( + %ipv4_addr, + "no peer route for external IPv4; use local exit-node fallback" + ); + dst_peers.push(self.my_peer_id); + is_exit_node = true; + } + (dst_peers, is_exit_node) + } + + pub async fn get_msg_dst_peer_ipv6(&self, ipv6_addr: &Ipv6Addr) -> (Vec, bool) { + let mut is_exit_node = false; + let mut dst_peers = vec![]; + if self.is_all_peers_broadcast_ipv6(ipv6_addr) { + dst_peers.extend(self.peers.list_routes().await.iter().map(|x| *x.key())); + } else if let Some(peer_id) = self.peers.get_peer_id_by_ipv6(ipv6_addr).await { + dst_peers.push(peer_id); + } else if !ipv6_addr.is_unicast_link_local() + && let Some(peer_id) = self.route.get_public_ipv6_gateway_peer_id().await + { + dst_peers.push(peer_id); + } else if !ipv6_addr.is_unicast_link_local() { + // NOTE: never route link local address to exit node. + for exit_node in self.exit_nodes.read().await.iter() { + let IpAddr::V6(exit_node) = exit_node else { + continue; + }; + if let Some(peer_id) = self.peers.get_peer_id_by_ipv6(exit_node).await { + dst_peers.push(peer_id); + is_exit_node = true; + break; + } + } + } + + (dst_peers, is_exit_node) + } + + pub async fn send_msg_by_ip( + &self, + mut msg: ZCPacket, + ip_addr: IpAddr, + not_send_to_self: bool, + ) -> Result<(), Error> { + tracing::trace!( + "do send_msg in peer manager, msg: {:?}, ip_addr: {}", + msg, + ip_addr + ); + + msg.fill_peer_manager_hdr(self.my_peer_id, 0, PacketType::Data as u8); + if !self.run_nic_packet_process_pipeline(&mut msg).await { + return Ok(()); + } + let cur_to_peer_id = msg.peer_manager_header().unwrap().to_peer_id.into(); + if cur_to_peer_id != 0 { + self.mark_recent_traffic(cur_to_peer_id); + return send_msg_internal( + self.peers.as_ref(), + &self.foreign_network_client, + &self.relay_peer_map, + Some(&self.traffic_metrics), + msg, + cur_to_peer_id, + ) + .await; + } + + let (dst_peers, is_exit_node) = match ip_addr { + IpAddr::V4(ipv4_addr) => self.get_msg_dst_peer_ipv4(&ipv4_addr).await, + IpAddr::V6(ipv6_addr) => self.get_msg_dst_peer_ipv6(&ipv6_addr).await, + }; + + if dst_peers.is_empty() { + tracing::info!("no peer id for ip: {}", ip_addr); + return Ok(()); + } + + self.counters + .compress_tx_bytes_before + .add(msg.buf_len() as u64); + + try_compress_and_encrypt( + self.data_compress_algo, + &self.encryptor, + &mut msg, + self.is_secure_mode_enabled, + ) + .await?; + + self.counters + .compress_tx_bytes_after + .add(msg.buf_len() as u64); + + let is_latency_first = self.context.latency_first(); + msg.mut_peer_manager_header() + .unwrap() + .set_latency_first(is_latency_first) + .set_exit_node(is_exit_node); + + let mut errs: Vec = vec![]; + let mut msg = Some(msg); + let total_dst_peers = dst_peers.len(); + let should_mark_recent_traffic = should_mark_recent_traffic_for_fanout(total_dst_peers); + for (i, peer_id) in dst_peers.iter().enumerate() { + if should_mark_recent_traffic { + self.mark_recent_traffic(*peer_id); + } + if let Err(e) = self.check_p2p_only_before_send(*peer_id) { + errs.push(e); + continue; + } + + let mut msg = if i == total_dst_peers - 1 { + msg.take().unwrap() + } else { + msg.clone().unwrap() + }; + + let hdr = msg.mut_peer_manager_header().unwrap(); + hdr.to_peer_id.set(*peer_id); + + if !self.host_routing.local_exit_node_fallback + && not_send_to_self + && *peer_id == self.my_peer_id + && !self.context.is_ip_local_virtual_ip(&ip_addr) + { + // Keep the loop-prevention flags for proxy-induced self-delivery where + // the destination is not this node's own EasyTier-managed IP. + hdr.set_not_send_to_tun(true); + hdr.set_no_proxy(true); + } + + self.counters.self_tx_bytes.add(msg.buf_len() as u64); + self.counters.self_tx_packets.inc(); + + if let Err(e) = send_msg_internal( + self.peers.as_ref(), + &self.foreign_network_client, + &self.relay_peer_map, + Some(&self.traffic_metrics), + msg, + *peer_id, + ) + .await + { + errs.push(e); + } + } + + tracing::trace!(?dst_peers, "do send_msg in peer manager done"); + + if errs.is_empty() { + Ok(()) + } else { + tracing::error!(?errs, "send_msg has error"); + Err(anyhow::anyhow!("send_msg has error: {:?}", errs).into()) + } + } +} + +struct PeerPacketRouterCounters { + self_tx_packets: CounterHandle, + self_tx_bytes: CounterHandle, + self_rx_packets: CounterHandle, + self_rx_bytes: CounterHandle, + forward_data_tx_packets: CounterHandle, + forward_data_tx_bytes: CounterHandle, + forward_control_tx_packets: CounterHandle, + forward_control_tx_bytes: CounterHandle, + compress_tx_bytes_before: CounterHandle, + compress_tx_bytes_after: CounterHandle, + compress_rx_bytes_before: CounterHandle, + compress_rx_bytes_after: CounterHandle, +} + +pub(crate) struct PeerPacketRouter { + packet_recv: PacketRecvChanReceiver, + my_peer_id: PeerId, + peers: Arc, + peer_packet_process_pipeline: Arc>>>, + foreign_client: Arc, + relay_peer_map: Arc, + foreign_network_manager: Arc, + encryptor: Arc, + compress_algo: CompressorAlgo, + acl_filter: Arc, + context: ArcPeerContext, + secure_mode_enabled: bool, + route: ArcRoute, + is_credential_node: bool, + traffic_metrics: Arc, + stats_mgr: Arc, + counters: PeerPacketRouterCounters, +} + +impl PeerPacketRouter { + #[allow(clippy::too_many_arguments)] + fn new( + packet_recv: PacketRecvChanReceiver, + my_peer_id: PeerId, + peers: Arc, + peer_packet_process_pipeline: Arc>>>, + foreign_client: Arc, + relay_peer_map: Arc, + foreign_network_manager: Arc, + encryptor: Arc, + compress_algo: CompressorAlgo, + acl_filter: Arc, + context: ArcPeerContext, + secure_mode_enabled: bool, + route: ArcRoute, + is_credential_node: bool, + traffic_metrics: Arc, + stats_mgr: Arc, + network_name: String, + self_tx_packets: CounterHandle, + self_tx_bytes: CounterHandle, + compress_tx_bytes_before: CounterHandle, + compress_tx_bytes_after: CounterHandle, + ) -> Self { + let label_set = LabelSet::new().with_label_type(LabelType::NetworkName(network_name)); + Self { + packet_recv, + my_peer_id, + peers, + peer_packet_process_pipeline, + foreign_client, + relay_peer_map, + foreign_network_manager, + encryptor, + compress_algo, + acl_filter, + context, + secure_mode_enabled, + route, + is_credential_node, + traffic_metrics, + stats_mgr: stats_mgr.clone(), + counters: PeerPacketRouterCounters { + self_tx_packets, + self_tx_bytes, + self_rx_bytes: stats_mgr + .get_counter(MetricName::TrafficBytesSelfRx, label_set.clone()), + self_rx_packets: stats_mgr + .get_counter(MetricName::TrafficPacketsSelfRx, label_set.clone()), + forward_data_tx_bytes: stats_mgr + .get_counter(MetricName::TrafficBytesForwarded, label_set.clone()), + forward_data_tx_packets: stats_mgr + .get_counter(MetricName::TrafficPacketsForwarded, label_set.clone()), + forward_control_tx_bytes: stats_mgr + .get_counter(MetricName::TrafficControlBytesForwarded, label_set.clone()), + forward_control_tx_packets: stats_mgr.get_counter( + MetricName::TrafficControlPacketsForwarded, + label_set.clone(), + ), + compress_tx_bytes_before, + compress_tx_bytes_after, + compress_rx_bytes_before: stats_mgr + .get_counter(MetricName::CompressionBytesRxBefore, label_set.clone()), + compress_rx_bytes_after: stats_mgr + .get_counter(MetricName::CompressionBytesRxAfter, label_set), + }, + } + } + + pub async fn run(mut self) { + tracing::trace!("start_peer_recv"); + while let Ok(ret) = recv_packet_from_chan(&mut self.packet_recv).await { + let disable_relay_data = self.context.disable_relay_data(); + let Err(ret) = try_handle_foreign_network_packet( + ret, + self.my_peer_id, + &self.peers, + &self.foreign_network_manager, + self.stats_mgr.as_ref(), + disable_relay_data, + ) + .await + else { + continue; + }; + + self.handle_packet(ret, disable_relay_data).await; + } + panic!("done_peer_recv"); + } + + async fn handle_packet(&self, mut ret: ZCPacket, disable_relay_data: bool) { + let buf_len = ret.buf_len(); + let is_relay_data_packet = is_relay_data_zc_packet(&ret); + let Some(hdr) = ret.mut_peer_manager_header() else { + tracing::warn!(?ret, "invalid packet, skip"); + return; + }; + + tracing::trace!(?hdr, "peer recv a packet..."); + let from_peer_id = hdr.from_peer_id.get(); + let to_peer_id = hdr.to_peer_id.get(); + let packet_type = hdr.packet_type; + let is_encrypted = hdr.is_encrypted(); + if to_peer_id != self.my_peer_id { + if disable_relay_data && is_relay_data_packet { + tracing::debug!( + ?from_peer_id, + ?to_peer_id, + packet_type, + "drop forwarded relay data while relay data is disabled" + ); + return; + } + + if hdr.forward_counter > 7 { + tracing::warn!(?hdr, "forward counter exceed, drop packet"); + return; + } + + // Step 10b: credential nodes don't forward handshake packets + if self.is_credential_node + && (packet_type == PacketType::HandShake as u8 + || packet_type == PacketType::NoiseHandshakeMsg1 as u8 + || packet_type == PacketType::NoiseHandshakeMsg2 as u8 + || packet_type == PacketType::NoiseHandshakeMsg3 as u8) + { + tracing::debug!("credential node dropping forwarded handshake packet"); + return; + } + + if hdr.forward_counter > 2 && hdr.is_latency_first() { + tracing::trace!(?hdr, "set_latency_first false because too many hop"); + hdr.set_latency_first(false); + } + + hdr.forward_counter += 1; + + if from_peer_id == self.my_peer_id { + self.counters.compress_tx_bytes_before.add(buf_len as u64); + + if packet_type == PacketType::Data as u8 + || packet_type == PacketType::KcpSrc as u8 + || packet_type == PacketType::KcpDst as u8 + { + let _ = try_compress_and_encrypt( + self.compress_algo, + &self.encryptor, + &mut ret, + self.secure_mode_enabled, + ) + .await; + } + + self.counters + .compress_tx_bytes_after + .add(ret.buf_len() as u64); + self.counters.self_tx_bytes.add(ret.buf_len() as u64); + self.counters.self_tx_packets.inc(); + } else { + match traffic_kind(packet_type) { + TrafficKind::Data => { + self.counters.forward_data_tx_bytes.add(buf_len as u64); + self.counters.forward_data_tx_packets.inc(); + } + TrafficKind::Control => { + self.counters.forward_control_tx_bytes.add(buf_len as u64); + self.counters.forward_control_tx_packets.inc(); + } + } + } + + tracing::trace!(?to_peer_id, my_peer_id = ?self.my_peer_id, "need forward"); + let tx_metrics = if from_peer_id == self.my_peer_id { + Some(&self.traffic_metrics) + } else { + None + }; + let ret = send_msg_internal( + self.peers.as_ref(), + &self.foreign_client, + &self.relay_peer_map, + tx_metrics, + ret, + to_peer_id, + ) + .await; + if ret.is_err() { + tracing::error!(?ret, ?to_peer_id, ?from_peer_id, "forward packet error"); + } + } else { + if packet_type == PacketType::RelayHandshake as u8 + || packet_type == PacketType::RelayHandshakeAck as u8 + { + let _ = self.relay_peer_map.handle_handshake_packet(ret).await; + return; + } + if !self.secure_mode_enabled { + if let Err(e) = self.encryptor.decrypt(&mut ret) { + tracing::error!(?e, "decrypt failed"); + return; + } + } else if is_encrypted { + match self.relay_peer_map.decrypt_if_needed(&mut ret).await { + Ok(true) => {} + Ok(false) => { + tracing::error!("secure session not found"); + return; + } + Err(e) => { + tracing::error!(?e, "secure decrypt failed"); + return; + } + } + } + + self.counters.self_rx_bytes.add(buf_len as u64); + self.counters.self_rx_packets.inc(); + self.traffic_metrics + .record_rx(from_peer_id, packet_type, buf_len as u64) + .await; + self.counters.compress_rx_bytes_before.add(buf_len as u64); + + let compressor = DefaultCompressor {}; + if let Err(e) = compressor.decompress(&mut ret).await { + tracing::error!(?e, "decompress failed"); + return; + } + + self.counters + .compress_rx_bytes_after + .add(ret.buf_len() as u64); + + if !self.acl_filter.process_packet_with_acl( + &ret, + true, + self.context.ipv4().map(|x| x.address()), + |dst| self.context.is_ip_local_ipv6(&dst), + self.route.as_ref(), + ) { + return; + } + + let mut processed = false; + let mut zc_packet = Some(ret); + tracing::trace!(?zc_packet, "try_process_packet_from_peer"); + for pipeline in self.peer_packet_process_pipeline.read().await.iter().rev() { + if !pipeline.active.load(Ordering::Acquire) { + continue; + } + let filter = pipeline.filter.read().clone(); + if let Some(filter) = filter { + zc_packet = filter + .try_process_packet_from_peer(zc_packet.unwrap()) + .await; + } + if zc_packet.is_none() { + processed = true; + break; + } + } + if !processed { + tracing::error!(?zc_packet, "unhandled packet"); + } + } + } +} + +pub(crate) fn is_relay_data_packet(packet_type: u8) -> bool { + super::traffic_metrics::is_relay_data_packet_type(packet_type) +} + +pub(crate) fn is_relay_data_zc_packet(packet: &ZCPacket) -> bool { + let Some(hdr) = packet.peer_manager_header() else { + return false; + }; + + if hdr.packet_type == PacketType::ForeignNetworkPacket as u8 { + let inner_packet_type = packet.foreign_network_inner_packet_type(); + if inner_packet_type.is_none() { + tracing::warn!( + ?hdr, + "foreign network packet has unparseable inner peer manager header" + ); + } + return inner_packet_type.is_none_or(is_relay_data_packet); + } + + is_relay_data_packet(hdr.packet_type) +} + +pub(crate) async fn try_handle_foreign_network_packet( + mut packet: ZCPacket, + my_peer_id: PeerId, + peer_map: &PeerMap, + foreign_network_manager: &ForeignNetworkManager, + stats_manager: &StatsManager, + disable_relay_data: bool, +) -> Result<(), ZCPacket> { + let pm_header = packet.peer_manager_header().unwrap(); + if pm_header.packet_type != PacketType::ForeignNetworkPacket as u8 { + return Err(packet); + } + + let from_peer_id = pm_header.from_peer_id.get(); + let to_peer_id = pm_header.to_peer_id.get(); + + if disable_relay_data && is_relay_data_zc_packet(&packet) { + tracing::debug!( + ?from_peer_id, + ?to_peer_id, + inner_packet_type = ?packet.foreign_network_inner_packet_type(), + "drop foreign network relay data while relay data is disabled" + ); + return Ok(()); + } + + let foreign_hdr = packet.foreign_network_hdr().unwrap(); + let foreign_network_name = foreign_hdr.get_network_name(packet.payload()); + let foreign_peer_id = foreign_hdr.get_dst_peer_id(); + + let foreign_network_my_peer_id = + foreign_network_manager.get_network_peer_id(&foreign_network_name); + + let buf_len = packet.buf_len(); + let label_set = + LabelSet::new().with_label_type(LabelType::NetworkName(foreign_network_name.clone())); + let add_counter = move |bytes_metric, packets_metric| { + stats_manager + .get_counter(bytes_metric, label_set.clone()) + .add(buf_len as u64); + stats_manager.get_counter(packets_metric, label_set).inc(); + }; + + // NOTICE: the to peer id is modified by the src from foreign network my peer id to the origin my peer id + if to_peer_id == my_peer_id { + // packet sent from other peer to me, extract the inner packet and forward it + add_counter( + MetricName::TrafficBytesForeignForwardRx, + MetricName::TrafficPacketsForeignForwardRx, + ); + if let Err(e) = foreign_network_manager + .forward_foreign_network_packet( + &foreign_network_name, + foreign_peer_id, + packet.foreign_network_packet(), + ) + .await + { + tracing::debug!( + ?e, + ?foreign_network_name, + ?foreign_peer_id, + "foreign network mgr send_msg_to_peer failed" + ); + } + Ok(()) + } else if Some(from_peer_id) == foreign_network_my_peer_id { + // to_peer_id is my peer id for the foreign network, need to convert to the origin my_peer_id of dst + let Some(to_peer_id) = peer_map + .get_origin_my_peer_id(&foreign_network_name, to_peer_id) + .await + else { + tracing::debug!( + ?foreign_network_name, + ?to_peer_id, + "cannot find origin my peer id for foreign network." + ); + return Err(packet); + }; + + add_counter( + MetricName::TrafficBytesForeignForwardTx, + MetricName::TrafficPacketsForeignForwardTx, + ); + + // modify the to_peer id from foreign network my peer id to the origin my peer id + packet + .mut_peer_manager_header() + .unwrap() + .to_peer_id + .set(to_peer_id); + + // packet is generated from foreign network mgr and should be forward to other peer + if let Err(e) = peer_map + .send_msg(packet, to_peer_id, NextHopPolicy::LeastHop) + .await + { + tracing::debug!( + ?e, + ?to_peer_id, + "send_msg_directly failed when forward local generated foreign network packet" + ); + } + Ok(()) + } else { + // target is not me, forward it. try get origin peer id + add_counter( + MetricName::TrafficBytesForeignForwardForwarded, + MetricName::TrafficPacketsForeignForwardForwarded, + ); + Err(packet) + } +} + +struct PeerManagerRouteInterface { + my_peer_id: PeerId, + peers: Weak, + foreign_network_client: Weak, + foreign_network_manager: Weak, +} + +#[async_trait::async_trait] +impl RouteInterface for PeerManagerRouteInterface { + async fn list_peers(&self) -> Vec { + let Some(foreign_client) = self.foreign_network_client.upgrade() else { + return vec![]; + }; + + let Some(peer_map) = self.peers.upgrade() else { + return vec![]; + }; + + let mut peers = foreign_client.list_public_peers().await; + peers.extend(peer_map.list_peers_with_conn().await); + peers + } + + fn my_peer_id(&self) -> PeerId { + self.my_peer_id + } + + async fn close_peer(&self, peer_id: PeerId) { + if let Some(peer_map) = self.peers.upgrade() { + let _ = peer_map.close_peer(peer_id).await; + } + + if let Some(foreign_client) = self.foreign_network_client.upgrade() { + let _ = foreign_client.get_peer_map().close_peer(peer_id).await; + } + } + + async fn get_peer_public_key(&self, peer_id: PeerId) -> Option> { + let peer_map = self.peers.upgrade()?; + peer_map.get_peer_public_key(peer_id) + } + + async fn get_peer_identity_type(&self, peer_id: PeerId) -> Option { + let peer_map = self.peers.upgrade()?; + peer_map.get_peer_identity_type(peer_id) + } + + async fn list_foreign_networks(&self) -> ForeignNetworkRouteInfoMap { + let ret = ForeignNetworkRouteInfoMap::new(); + let Some(foreign_network_manager) = self.foreign_network_manager.upgrade() else { + return ret; + }; + + let networks = foreign_network_manager + .list_foreign_network_route_infos() + .await; + for info in networks { + if info.peer_ids.is_empty() { + continue; + } + + let last_update = foreign_network_manager + .get_foreign_network_last_update(&info.network_name) + .unwrap_or(SystemTime::now()); + ret.insert( + ForeignNetworkRouteInfoKey { + peer_id: self.my_peer_id, + network_name: info.network_name, + }, + ForeignNetworkRouteInfoEntry { + foreign_peer_ids: info.peer_ids, + last_update: Some(last_update.into()), + version: 0, + network_secret_digest: info.network_secret_digest, + my_peer_id_for_this_network: info.my_peer_id_for_this_network, + }, + ); + } + ret + } +} + +pub(crate) async fn send_msg_internal( + peers: &PeerMap, + foreign_network_client: &Arc, + relay_peer_map: &Arc, + direct_tx_metrics: Option<&Arc>, + msg: ZCPacket, + dst_peer_id: PeerId, +) -> Result<(), Error> { + let policy = get_next_hop_policy(msg.peer_manager_header().unwrap().is_latency_first()); + let is_latency_first = msg.peer_manager_header().unwrap().is_latency_first(); + let packet_type = msg.peer_manager_header().unwrap().packet_type; + let msg_len = msg.buf_len() as u64; + let latency_first_gateway = if is_latency_first { + peers + .get_gateway_peer_id(dst_peer_id, policy.clone()) + .await + .filter(|gateway| *gateway != dst_peer_id) + } else { + None + }; + let send_result = if let Some(gateway) = latency_first_gateway + && (peers.has_peer(gateway) || foreign_network_client.has_next_hop(gateway)) + { + relay_peer_map.send_msg(msg, dst_peer_id, policy).await + } else if peers.has_peer(dst_peer_id) { + peers.send_msg_directly(msg, dst_peer_id).await + } else if foreign_network_client.has_next_hop(dst_peer_id) { + foreign_network_client.send_msg(msg, dst_peer_id).await + } else if let Some(gateway) = peers.get_gateway_peer_id(dst_peer_id, policy.clone()).await { + if peers.has_peer(gateway) || foreign_network_client.has_next_hop(gateway) { + relay_peer_map.send_msg(msg, dst_peer_id, policy).await + } else { + tracing::warn!( + ?gateway, + ?dst_peer_id, + "cannot send msg to peer through gateway" + ); + Err(Error::RouteError(None)) + } + } else if foreign_network_client.has_next_hop(dst_peer_id) { + // check foreign network again. so in happy path we can avoid extra check + foreign_network_client.send_msg(msg, dst_peer_id).await + } else { + tracing::debug!(?dst_peer_id, "no gateway for peer"); + Err(Error::RouteError(None)) + }; + + if send_result.is_ok() + && let Some(metrics) = direct_tx_metrics + { + metrics.record_tx(dst_peer_id, packet_type, msg_len).await; + } + + send_result +} + +#[cfg(test)] +mod tests { + use std::{ + sync::atomic::{AtomicUsize, Ordering}, + time::Duration, + }; + + use base64::{Engine as _, engine::general_purpose::STANDARD as BASE64_STANDARD}; + use dashmap::DashMap; + use quanta::Instant; + use x25519_dalek::{PublicKey, StaticSecret}; + + use super::*; + use crate::{ + config::runtime::CoreRuntimeConfig, + config::{CoreConfig, IpPrefix, NetworkIdentity, NodeConfig, ProxyNetworkConfig}, + peers::{ + context::{PeerContext, PeerEvent}, + create_packet_recv_chan, + }, + proto::common::{PeerFeatureFlag, StunInfo}, + }; + + impl PeerManagerCore { + pub(crate) fn new_portable_for_test( + config: PortablePeerManagerConfig, + nic_channel: PacketRecvChan, + ) -> anyhow::Result { + let runtime_config = CoreRuntimeConfigStore::new( + CoreRuntimeConfig::default(), + Arc::new(config.snapshot.clone()), + ); + let public_ipv6_runtime = + CorePublicIpv6Runtime::new(runtime_config.clone(), Arc::new(()), Arc::new(())); + let stun_info_source = Arc::new(RuntimeConfigStunInfoSource(runtime_config.clone())); + Self::new( + config, + runtime_config, + stun_info_source, + nic_channel, + public_ipv6_runtime, + Arc::new(()), + None, + Arc::new(()), + ) + } + } + + struct RuntimeConfigStunInfoSource(CoreRuntimeConfigStore); + + impl PeerStunInfoSource for RuntimeConfigStunInfoSource { + fn stun_info(&self) -> StunInfo { + self.0.snapshot().peer.runtime.stun_info.clone() + } + } + + struct SameNetworkContext { + contains_every_address: bool, + } + + impl PeerContext for SameNetworkContext { + fn network_identity(&self) -> NetworkIdentity { + NetworkIdentity { + network_name: "test".to_string(), + network_secret: None, + network_secret_digest: None, + } + } + + fn is_ip_in_same_network(&self, ip: &IpAddr) -> bool { + self.contains_every_address + || matches!(ip, IpAddr::V4(ip) if ip.octets()[0..2] == [10, 144]) + } + } + + #[derive(Default)] + struct CountingPeerEventSink(AtomicUsize); + + impl crate::events::CoreEventSink for CountingPeerEventSink { + fn emit(&self, _event: crate::events::CoreEvent) { + self.0.fetch_add(1, Ordering::Relaxed); + } + } + + struct DropCountingNicFilter(Arc); + + impl Drop for DropCountingNicFilter { + fn drop(&mut self) { + self.0.fetch_add(1, Ordering::Relaxed); + } + } + + #[async_trait::async_trait] + impl super::super::NicPacketFilter for DropCountingNicFilter { + async fn try_process_packet_from_nic(&self, _data: &mut ZCPacket) -> bool { + true + } + } + + #[tokio::test] + async fn managed_nic_pipeline_removal_waits_for_readers_and_drops_filter() { + let drops = Arc::new(AtomicUsize::new(0)); + let (entry, registration) = + managed_nic_pipeline_entry(Box::new(DropCountingNicFilter(drops.clone()))); + let pipeline = Arc::new(RwLock::new(vec![entry])); + let reader = pipeline.read().await; + let active_filter = reader[0].filter.read().clone().unwrap(); + let remove_pipeline = pipeline.clone(); + let remove_registration = registration.clone(); + let removal = tokio::spawn(async move { + remove_managed_nic_pipeline_entry(&remove_pipeline, &remove_registration).await; + }); + + tokio::task::yield_now().await; + assert!(!removal.is_finished()); + assert_eq!(drops.load(Ordering::Relaxed), 0); + + drop(reader); + drop(active_filter); + removal.await.unwrap(); + assert!(pipeline.read().await.is_empty()); + assert_eq!(drops.load(Ordering::Relaxed), 1); + } + + #[test] + fn managed_pipeline_guard_releases_filter_without_a_runtime() { + let drops = Arc::new(AtomicUsize::new(0)); + let (entry, registration) = + managed_nic_pipeline_entry(Box::new(DropCountingNicFilter(drops.clone()))); + + drop(registration); + + assert!(entry.filter.read().is_none()); + assert_eq!(drops.load(Ordering::Relaxed), 1); + } + + fn portable_runtime_config(network_name: &str) -> PeerRuntimeConfig { + PeerRuntimeConfig { + core: CoreConfig { + node: NodeConfig { + peer_id: None, + network_name: network_name.to_owned(), + ..Default::default() + }, + ..Default::default() + }, + network_identity: NetworkIdentity { + network_name: network_name.to_owned(), + network_secret: Some("secret".to_owned()), + network_secret_digest: None, + }, + stun_info: StunInfo::default(), + feature_flags: PeerFeatureFlag::default(), + secure_mode: None, + host_routing: HostRoutingPolicy::default(), + } + } + + fn credential_secure_mode() -> crate::proto::common::SecureModeConfig { + let private = StaticSecret::from([7; 32]); + let public = PublicKey::from(&private); + crate::proto::common::SecureModeConfig { + enabled: true, + local_private_key: Some(BASE64_STANDARD.encode(private.as_bytes())), + local_public_key: Some(BASE64_STANDARD.encode(public.as_bytes())), + } + } + + fn build_portable_for_test(runtime: PeerRuntimeConfig) -> anyhow::Result { + build_portable_config_for_test(PortablePeerManagerConfig::new(runtime)) + } + + fn build_portable_config_for_test( + config: PortablePeerManagerConfig, + ) -> anyhow::Result { + let (packet_tx, _packet_rx) = create_packet_recv_chan(); + PeerManagerCore::new_portable_for_test(config, packet_tx) + } + + #[tokio::test] + async fn portable_peer_manager_builds_and_stops_from_normalized_config() { + let runtime = portable_runtime_config("portable-net"); + let core = build_portable_for_test(runtime).unwrap(); + + assert_eq!(core.context.network_name(), "portable-net"); + assert_ne!(core.context.instance_id(), uuid::Uuid::nil()); + assert_eq!(core.data_compress_algo, CompressorAlgo::None); + assert!(core.list_foreign_network_infos(false).await.is_empty()); + + core.run().await.unwrap(); + let route = core.route_algo_inst.ospf_route().unwrap(); + assert!(route.task_count() > 0); + assert!(!core.stats_manager().cleanup_task_is_stopped()); + assert!(!core.acl_filter.cleanup_task_is_stopped()); + core.clear_resources().await; + assert_eq!(route.task_count(), 0); + assert!(core.stats_manager().cleanup_task_is_stopped()); + assert!(core.acl_filter.cleanup_task_is_stopped()); + assert!(core.foreign_network_manager.is_stopped_for_test().await); + assert!( + !core + .foreign_network_manager + .admission_is_open_for_test() + .await + ); + } + + #[cfg(not(feature = "zstd"))] + #[tokio::test] + async fn portable_peer_manager_rejects_requested_unavailable_zstd() { + let mut config = PortablePeerManagerConfig::new(portable_runtime_config("portable-net")); + config.snapshot.flags.data_compress_algo = + crate::proto::common::CompressionAlgoPb::Zstd.into(); + + let error = build_portable_config_for_test(config).err().unwrap(); + + assert_eq!( + error.to_string(), + "compression algorithm is unavailable in this build: ZstdDefault" + ); + } + + #[cfg(not(feature = "aes-gcm"))] + #[tokio::test] + async fn portable_peer_manager_rejects_requested_unavailable_aes() { + let mut config = PortablePeerManagerConfig::new(portable_runtime_config("portable-net")); + config.snapshot.flags.enable_encryption = true; + config.snapshot.flags.encryption_algorithm = "aes-gcm".to_owned(); + + let error = build_portable_config_for_test(config).err().unwrap(); + + assert_eq!( + error.to_string(), + "encryption algorithm is unavailable in this build: aes-gcm" + ); + } + + #[tokio::test] + async fn portable_peer_manager_rejects_unknown_encryption_algorithm() { + let mut config = PortablePeerManagerConfig::new(portable_runtime_config("portable-net")); + config.snapshot.flags.enable_encryption = true; + config.snapshot.flags.encryption_algorithm = "rot13".to_owned(); + + let error = build_portable_config_for_test(config).err().unwrap(); + + assert_eq!(error.to_string(), "invalid encryption algorithm: rot13"); + } + + #[tokio::test] + async fn unknown_ipv6_has_no_peer_destination() { + let core = build_portable_for_test(portable_runtime_config("ipv6-net")).unwrap(); + let unknown = "fd00::2".parse().unwrap(); + + let (peers, is_self) = core.get_msg_dst_peer_ipv6(&unknown).await; + + assert!(peers.is_empty()); + assert!(!is_self); + } + + #[tokio::test] + async fn foreign_network_stop_waits_for_inflight_admission() { + let core = build_portable_for_test(portable_runtime_config("portable-net")).unwrap(); + let manager = core.foreign_network_manager.clone(); + let entered = Arc::new(tokio::sync::Notify::new()); + let release = Arc::new(tokio::sync::Notify::new()); + let admission_manager = manager.clone(); + let admission_entered = entered.clone(); + let admission_release = release.clone(); + let admission = tokio::spawn(async move { + admission_manager + .hold_admission_for_test(admission_entered, admission_release) + .await + }); + entered.notified().await; + + let stop_manager = manager.clone(); + let stop = tokio::spawn(async move { stop_manager.stop().await }); + tokio::task::yield_now().await; + assert!(!stop.is_finished()); + + release.notify_waiters(); + admission.await.unwrap().unwrap(); + stop.await.unwrap(); + assert!(manager.is_stopped_for_test().await); + assert!(!manager.admission_is_open_for_test().await); + + core.clear_resources().await; + } + + #[tokio::test] + async fn portable_peer_manager_uses_host_context_adapters() { + let config = PortablePeerManagerConfig::new(portable_runtime_config("portable-net")); + let runtime_config = CoreRuntimeConfigStore::new( + CoreRuntimeConfig::default(), + Arc::new(config.snapshot.clone()), + ); + let public_ipv6_runtime = + CorePublicIpv6Runtime::new(runtime_config.clone(), Arc::new(()), Arc::new(())); + let events = Arc::new(CountingPeerEventSink::default()); + let (packet_tx, _packet_rx) = create_packet_recv_chan(); + + let core = PeerManagerCore::new( + config, + runtime_config, + Arc::new(()), + packet_tx, + public_ipv6_runtime, + events.clone(), + None, + Arc::new(()), + ) + .unwrap(); + + core.context.issue_event(PeerEvent::PeerAdded(99)); + assert_eq!(events.0.load(Ordering::Relaxed), 1); + core.clear_resources().await; + } + + #[tokio::test] + async fn portable_peer_assembly_preserves_submitted_acl_groups() { + let mut config = PortablePeerManagerConfig::new(portable_runtime_config("portable-net")); + let acl = crate::proto::acl::Acl { + acl_v1: Some(crate::proto::acl::AclV1 { + chains: Vec::new(), + group: Some(crate::proto::acl::GroupInfo { + declares: vec![crate::proto::acl::GroupIdentity { + group_name: "ops".to_owned(), + group_secret: "ops-secret".to_owned(), + }], + members: vec!["ops".to_owned()], + }), + }), + }; + config.snapshot.set_acl_groups(Some(&acl)); + let (packet_tx, _packet_rx) = create_packet_recv_chan(); + + let core = PeerManagerCore::new_portable_for_test(config, packet_tx).unwrap(); + + let groups = core.context.peer_groups(86); + assert_eq!(groups.len(), 1); + assert_eq!(groups[0].group_name, "ops"); + assert!(groups[0].verify("ops-secret", 86)); + assert_eq!(core.context.acl_group_declarations()[0].group_name, "ops"); + core.clear_resources().await; + } + + #[tokio::test] + async fn node_snapshot_exposes_normalized_runtime_state() { + let instance_id = uuid::Uuid::from_u128(0x00112233445566778899aabbccddeeff); + let mut runtime = portable_runtime_config("portable-net"); + runtime.core.node.instance_id = Some(*instance_id.as_bytes()); + runtime.core.node.hostname = Some("portable-node".to_owned()); + runtime.core.routes.ipv4 = Some(IpPrefix::new("10.20.0.91".parse().unwrap(), 16).unwrap()); + runtime.core.routes.proxy_networks = vec![ProxyNetworkConfig { + real: IpPrefix::new("10.40.0.0".parse().unwrap(), 16).unwrap(), + mapped: Some(IpPrefix::new("10.50.0.0".parse().unwrap(), 16).unwrap()), + }]; + runtime.stun_info.public_ip = vec!["192.0.2.91".to_owned()]; + let core = build_portable_for_test(runtime).unwrap(); + let listener = Url::parse("tcp://0.0.0.0:11010").unwrap(); + + let snapshot = core.node_snapshot(vec![listener.clone()]).await; + + assert_eq!(snapshot.peer_id, core.my_peer_id()); + assert_eq!(snapshot.instance_id, instance_id); + assert_eq!(snapshot.hostname, "portable-node"); + assert_eq!(snapshot.ipv4_addr, Some("10.20.0.91/16".parse().unwrap())); + assert_eq!(snapshot.proxy_networks.len(), 1); + assert_eq!(snapshot.listeners, vec![listener]); + assert_eq!(snapshot.stun_info.public_ip, vec!["192.0.2.91"]); + assert_eq!(snapshot.version, env!("CARGO_PKG_VERSION")); + assert!(snapshot.public_ipv6_addr.is_none()); + assert!(snapshot.ipv6_public_addr_prefix.is_none()); + + let dns_identity = core.dns_route_identity(); + assert_eq!( + dns_identity, + ( + "portable-node".to_owned(), + Some("10.20.0.91/16".parse::().unwrap().into()), + String::new(), + ) + ); + } + + #[tokio::test] + async fn portable_peer_manager_auth_uses_managed_credentials() { + let admin_a = build_portable_for_test(portable_runtime_config("portable-net")).unwrap(); + let admin_b = build_portable_for_test(portable_runtime_config("portable-net")).unwrap(); + let generated = admin_a.credential_manager().generate_credential( + vec!["guest".to_owned()], + false, + Vec::new(), + Duration::from_secs(3600), + ); + let private_bytes: [u8; 32] = BASE64_STANDARD + .decode(generated.secret) + .unwrap() + .try_into() + .unwrap(); + let public_key = PublicKey::from(&StaticSecret::from(private_bytes)); + + assert!( + admin_a + .context + .is_pubkey_trusted(public_key.as_bytes(), "portable-net") + ); + assert!( + !admin_a + .context + .is_pubkey_trusted(public_key.as_bytes(), "other") + ); + let trusted = admin_a.context.trusted_credential_pubkeys("secret"); + assert_eq!(trusted.len(), 1); + + let propagated_key = trusted[0].credential.as_ref().unwrap().pubkey.clone(); + admin_b.context.update_trusted_keys( + std::collections::HashMap::from([( + propagated_key.clone(), + crate::peers::context::TrustedKeyMetadata { + source: crate::peers::context::TrustedKeySource::OspfCredential, + expiry_unix: None, + }, + )]), + "portable-net", + ); + assert!( + admin_b + .context + .is_pubkey_trusted(&propagated_key, "portable-net") + ); + assert!(admin_b.context.is_pubkey_trusted_with_source( + &propagated_key, + "portable-net", + crate::peers::context::TrustedKeySource::OspfCredential, + )); + assert!(!admin_b.context.is_pubkey_trusted_with_source( + &propagated_key, + "portable-net", + crate::peers::context::TrustedKeySource::OspfNode, + )); + + assert!( + admin_a + .credential_manager() + .revoke_credential(&generated.credential_id) + ); + admin_b + .context + .update_trusted_keys(std::collections::HashMap::new(), "portable-net"); + assert!( + !admin_b + .context + .is_pubkey_trusted(&propagated_key, "portable-net") + ); + } + + #[tokio::test] + async fn portable_host_policy_controls_local_exit_node_fallback() { + let external_ipv4 = Ipv4Addr::new(203, 0, 113, 10); + let default_core = + build_portable_for_test(portable_runtime_config("portable-net")).unwrap(); + assert_eq!( + default_core.get_msg_dst_peer_ipv4(&external_ipv4).await, + (Vec::new(), false) + ); + + let mut runtime = portable_runtime_config("portable-net"); + runtime.host_routing.local_exit_node_fallback = true; + let fallback_core = build_portable_for_test(runtime).unwrap(); + assert_eq!( + fallback_core.get_msg_dst_peer_ipv4(&external_ipv4).await, + (vec![fallback_core.my_peer_id()], true) + ); + } + + #[tokio::test] + async fn portable_peer_manager_rejects_inconsistent_network_names() { + let mut runtime = portable_runtime_config("identity-net"); + runtime.core.node.network_name = "node-net".to_owned(); + let (packet_tx, _packet_rx) = create_packet_recv_chan(); + + let result = PeerManagerCore::new_portable_for_test( + PortablePeerManagerConfig::new(runtime), + packet_tx, + ); + + assert!(result.is_err()); + } + + #[tokio::test] + async fn portable_peer_manager_rejects_unavailable_config_capabilities() { + let mut digest_mismatch = portable_runtime_config("portable-net"); + digest_mismatch.network_identity.network_secret_digest = Some([1; 32]); + assert!(build_portable_for_test(digest_mismatch).is_err()); + + let mut secure_without_keys = portable_runtime_config("portable-net"); + secure_without_keys.secure_mode = Some(crate::proto::common::SecureModeConfig { + enabled: true, + ..Default::default() + }); + assert!(build_portable_for_test(secure_without_keys).is_err()); + + let mut mismatched_keys = portable_runtime_config("portable-net"); + mismatched_keys.secure_mode = Some(credential_secure_mode()); + mismatched_keys + .secure_mode + .as_mut() + .unwrap() + .local_public_key = Some(BASE64_STANDARD.encode([9; 32])); + assert!(build_portable_for_test(mismatched_keys).is_err()); + + let mut credential_without_secure_mode = portable_runtime_config("portable-net"); + credential_without_secure_mode + .network_identity + .network_secret = None; + credential_without_secure_mode + .network_identity + .network_secret_digest = None; + assert!(build_portable_for_test(credential_without_secure_mode).is_err()); + } + + #[tokio::test] + async fn portable_peer_manager_accepts_credential_client_config() { + let mut runtime = portable_runtime_config("portable-net"); + runtime.network_identity.network_secret = None; + runtime.network_identity.network_secret_digest = None; + runtime.secure_mode = Some(credential_secure_mode()); + + let core = build_portable_for_test(runtime).unwrap(); + assert!(core.context.feature_flags().is_credential_peer); + assert!(core.context.network_identity().network_secret.is_none()); + assert!(core.is_secure_mode_enabled); + core.clear_resources().await; + } + + #[tokio::test] + async fn portable_peer_manager_accepts_legacy_unlimited_limits() { + let runtime = portable_runtime_config("portable-net"); + let mut flags = PortablePeerManagerConfig::new(runtime.clone()) + .snapshot + .flags; + flags.instance_recv_bps_limit = u64::MAX; + flags.foreign_relay_bps_limit = u64::MAX; + let mut config = PortablePeerManagerConfig::new(runtime.clone()); + config.snapshot = PeerRuntimeSnapshot::new(runtime, flags); + + let core = build_portable_config_for_test(config).unwrap(); + assert!(core.context.recv_limiter("portable-net", false).is_none()); + assert!(core.context.recv_limiter("foreign-net", true).is_none()); + core.clear_resources().await; + } + + #[tokio::test] + async fn portable_peer_manager_builds_configured_recv_limiters() { + let mut runtime = portable_runtime_config("portable-net"); + runtime.core.traffic.instance_recv_bps_limit = Some(1024); + runtime.core.traffic.foreign_relay_bps_limit = Some(2048); + let core = build_portable_for_test(runtime).unwrap(); + + let instance_a = core.context.recv_limiter("portable-net", false).unwrap(); + let instance_b = core.context.recv_limiter("other-net", false).unwrap(); + assert!(Arc::ptr_eq(&instance_a, &instance_b)); + + let foreign_a = core.context.recv_limiter("foreign-a", true).unwrap(); + let foreign_a_again = core.context.recv_limiter("foreign-a", true).unwrap(); + let foreign_b = core.context.recv_limiter("foreign-b", true).unwrap(); + let foreign_named_instance = core.context.recv_limiter("instance", true).unwrap(); + assert!(Arc::ptr_eq(&foreign_a, &foreign_a_again)); + assert!(!Arc::ptr_eq(&foreign_a, &foreign_b)); + assert!(!Arc::ptr_eq(&instance_a, &foreign_a)); + assert!(!Arc::ptr_eq(&instance_a, &foreign_named_instance)); + + core.clear_resources().await; + assert!(core.context.recv_limiter("portable-net", false).is_none()); + assert!(core.context.recv_limiter("foreign-a", true).is_none()); + } + + #[tokio::test] + async fn portable_peer_manager_rejects_invalid_identity_and_prefixes() { + let mut digest_only = portable_runtime_config("portable-net"); + digest_only.network_identity.network_secret = None; + digest_only.network_identity.network_secret_digest = Some([1; 32]); + assert!(build_portable_for_test(digest_only).is_err()); + + let mut wrong_family = portable_runtime_config("portable-net"); + wrong_family.core.routes.ipv4 = Some(IpPrefix { + address: "2001:db8::1".parse().unwrap(), + prefix_len: 64, + }); + assert!(build_portable_for_test(wrong_family).is_err()); + + let mut proxy_host_bits = portable_runtime_config("portable-net"); + proxy_host_bits.core.routes.proxy_networks = vec![crate::config::ProxyNetworkConfig { + real: IpPrefix::new("10.50.0.7".parse().unwrap(), 16).unwrap(), + mapped: None, + }]; + assert!(build_portable_for_test(proxy_host_bits).is_err()); + + let mut wrong_proxy_family = portable_runtime_config("portable-net"); + wrong_proxy_family.core.routes.proxy_networks = vec![crate::config::ProxyNetworkConfig { + real: IpPrefix::new("10.50.0.0".parse().unwrap(), 16).unwrap(), + mapped: Some(IpPrefix::new("2001:db8::".parse().unwrap(), 64).unwrap()), + }]; + assert!(build_portable_for_test(wrong_proxy_family).is_err()); + + let mut advertised = portable_runtime_config("portable-net"); + advertised + .core + .routes + .advertised_routes + .push(IpPrefix::new("10.60.0.0".parse().unwrap(), 16).unwrap()); + assert!(build_portable_for_test(advertised).is_err()); + + let mut foreign = portable_runtime_config("portable-net"); + foreign + .core + .routes + .foreign_networks + .push(crate::config::ForeignNetworkConfig { + name: "other-net".to_owned(), + cidrs: Vec::new(), + }); + assert!(build_portable_for_test(foreign).is_err()); + } + + #[test] + fn portable_peer_manager_reports_missing_tokio_runtime() { + let result = build_portable_for_test(portable_runtime_config("portable-net")); + + let Err(error) = result else { + panic!("construction outside Tokio must fail"); + }; + assert!(error.to_string().contains("entered Tokio runtime")); + } + + #[test] + fn recent_traffic_fanout_policy_only_marks_single_peer() { + assert!(should_mark_recent_traffic_for_fanout(0)); + assert!(should_mark_recent_traffic_for_fanout(1)); + assert!(!should_mark_recent_traffic_for_fanout(2)); + } + + fn resolved_remote_addr_from_url(addr: Option<&str>) -> Option { + addr.map(|addr| Url::parse(addr).unwrap().into()) + } + + #[test] + fn resolved_remote_addr_check_rejects_virtual_network_ip() { + let context: ArcPeerContext = Arc::new(SameNetworkContext { + contains_every_address: false, + }); + let resolved_remote_addr = resolved_remote_addr_from_url(Some("tcp://10.144.0.2:1234")); + + let err = + check_resolved_remote_addr_not_from_virtual_network(&context, resolved_remote_addr); + + assert!(matches!(err, Err(Error::Other(_)))); + } + + #[test] + fn resolved_remote_addr_check_allows_external_or_non_ip_sources() { + let context: ArcPeerContext = Arc::new(SameNetworkContext { + contains_every_address: false, + }); + for resolved_remote_addr_url in [ + Some("tcp://192.0.2.10:1234"), + Some("tcp://example.test:1234"), + Some("ring://peer"), + Some("unix:///tmp/easytier.sock"), + None, + ] { + assert!( + check_resolved_remote_addr_not_from_virtual_network( + &context, + resolved_remote_addr_from_url(resolved_remote_addr_url), + ) + .is_ok() + ); + } + } + + #[test] + fn resolved_remote_addr_check_allows_loopback_inside_virtual_network() { + let context: ArcPeerContext = Arc::new(SameNetworkContext { + contains_every_address: true, + }); + let resolved_remote_addr = resolved_remote_addr_from_url(Some("tcp://127.0.0.1:1234")); + + let ret = + check_resolved_remote_addr_not_from_virtual_network(&context, resolved_remote_addr); + + assert!(ret.is_ok()); + } + + #[test] + fn disable_relay_data_classifies_data_plane_packets_only() { + for packet_type in [ + PacketType::Data, + PacketType::KcpSrc, + PacketType::KcpDst, + PacketType::QuicSrc, + PacketType::QuicDst, + PacketType::DataWithKcpSrcModified, + PacketType::DataWithQuicSrcModified, + PacketType::ForeignNetworkPacket, + ] { + assert!(is_relay_data_packet(packet_type as u8)); + } + + for packet_type in [ + PacketType::RpcReq, + PacketType::RpcResp, + PacketType::Ping, + PacketType::Pong, + PacketType::HandShake, + PacketType::NoiseHandshakeMsg1, + PacketType::NoiseHandshakeMsg2, + PacketType::NoiseHandshakeMsg3, + PacketType::RelayHandshake, + PacketType::RelayHandshakeAck, + ] { + assert!(!is_relay_data_packet(packet_type as u8)); + } + } + + #[test] + fn disable_relay_data_inspects_foreign_network_inner_packet_type() { + let network_name = "net1".to_string(); + + let mut rpc_packet = ZCPacket::new_with_payload(b"rpc"); + rpc_packet.fill_peer_manager_hdr(1, 2, PacketType::RpcReq as u8); + let mut foreign_rpc_packet = + ZCPacket::new_for_foreign_network(&network_name, 2, &rpc_packet); + foreign_rpc_packet.fill_peer_manager_hdr(10, 20, PacketType::ForeignNetworkPacket as u8); + + assert_eq!( + foreign_rpc_packet.foreign_network_inner_packet_type(), + Some(PacketType::RpcReq as u8) + ); + assert!(!is_relay_data_zc_packet(&foreign_rpc_packet)); + + let mut data_packet = ZCPacket::new_with_payload(b"data"); + data_packet.fill_peer_manager_hdr(1, 2, PacketType::Data as u8); + let mut foreign_data_packet = + ZCPacket::new_for_foreign_network(&network_name, 2, &data_packet); + foreign_data_packet.fill_peer_manager_hdr(10, 20, PacketType::ForeignNetworkPacket as u8); + + assert_eq!( + foreign_data_packet.foreign_network_inner_packet_type(), + Some(PacketType::Data as u8) + ); + assert!(is_relay_data_zc_packet(&foreign_data_packet)); + } + + fn route_with_ipv4( + peer_id: u32, + ipv4_addr: Option, + ) -> crate::proto::core_peer::peer::Route { + crate::proto::core_peer::peer::Route { + peer_id, + ipv4_addr: ipv4_addr.map(|addr| cidr::Ipv4Inet::new(addr, 24).unwrap().into()), + ..Default::default() + } + } + + #[test] + fn ipv4_broadcast_peer_selection_skips_peers_without_ipv4() { + let routes = vec![ + route_with_ipv4(1, Some(std::net::Ipv4Addr::new(10, 126, 126, 1))), + route_with_ipv4(2, None), + route_with_ipv4(3, Some(std::net::Ipv4Addr::new(10, 126, 126, 3))), + route_with_ipv4(4, None), + ]; + + assert_eq!( + PeerOutboundPacketRouter::select_ipv4_broadcast_peers(&routes, 3), + vec![1] + ); + } + + #[test] + fn gc_recent_traffic_removes_expired_and_connected_entries() { + let stale_peer = 1; + let direct_peer = 2; + let active_peer = 3; + let recent_have_traffic = DashMap::new(); + + recent_have_traffic.insert( + stale_peer, + Instant::now() - RECENT_HAVE_TRAFFIC_TTL - Duration::from_millis(1), + ); + recent_have_traffic.insert(direct_peer, Instant::now()); + recent_have_traffic.insert(active_peer, Instant::now()); + + let future_peer = 4; + recent_have_traffic.insert(future_peer, Instant::now() + Duration::from_secs(1)); + + gc_recent_traffic_entries(&recent_have_traffic, Instant::now(), |peer_id| { + peer_id == direct_peer + }); + + assert!(!recent_have_traffic.contains_key(&stale_peer)); + assert!(!recent_have_traffic.contains_key(&direct_peer)); + assert!(recent_have_traffic.contains_key(&active_peer)); + assert!(recent_have_traffic.contains_key(&future_peer)); + } + + #[test] + fn recent_traffic_notifies_only_when_demand_becomes_active() { + let tracker = RecentTrafficTracker::new(1); + let peer_id = 2; + let signal = tracker.p2p_demand_notify(); + + let initial_version = signal.version(); + tracker.mark(peer_id, false, true, |_| false); + assert_eq!(signal.version(), initial_version + 1); + + let first_seen = *tracker.recent_have_traffic.get(&peer_id).unwrap(); + std::thread::sleep(Duration::from_millis(5)); + tracker.mark(peer_id, false, true, |_| false); + assert_eq!( + signal.version(), + initial_version + 1, + "fresh demand should not wake all p2p workers again" + ); + let refreshed_seen = *tracker.recent_have_traffic.get(&peer_id).unwrap(); + assert!(refreshed_seen > first_seen); + + if let Some(mut last_seen) = tracker.recent_have_traffic.get_mut(&peer_id) { + *last_seen = Instant::now() - RECENT_HAVE_TRAFFIC_TTL - Duration::from_millis(1); + } + tracker.mark(peer_id, false, true, |_| false); + assert_eq!(signal.version(), initial_version + 2); + } + + #[test] + fn recent_traffic_tolerates_future_timestamps() { + let tracker = RecentTrafficTracker::new(1); + let peer_id = 2; + tracker + .recent_have_traffic + .insert(peer_id, Instant::now() + Duration::from_secs(1)); + + assert!(tracker.has(peer_id, Instant::now(), |_| false)); + tracker.mark(peer_id, false, true, |_| false); + } +} diff --git a/easytier-core/src/peers/peer_rpc.rs b/easytier-core/src/peers/peer_rpc.rs new file mode 100644 index 00000000..274857f8 --- /dev/null +++ b/easytier-core/src/peers/peer_rpc.rs @@ -0,0 +1,113 @@ +use std::sync::{Arc, Mutex}; + +use futures::{SinkExt as _, StreamExt}; +use tokio::task::JoinSet; + +use crate::{ + config::PeerId, + foundation::stats::{ArcRpcMetrics, RpcMetricsProvider}, + packet::ZCPacket, + rpc::{self, bidirect::BidirectRpcManager}, +}; + +#[async_trait::async_trait] +#[auto_impl::auto_impl(Arc)] +pub trait PeerRpcManagerTransport: Send + Sync + 'static { + fn my_peer_id(&self) -> PeerId; + async fn send(&self, msg: ZCPacket, dst_peer_id: PeerId) -> anyhow::Result<()>; + async fn recv(&self) -> anyhow::Result; +} + +pub struct PeerRpcManager { + tspt: Arc>, + bidirect_rpc: BidirectRpcManager, + tasks: Mutex>, +} + +impl std::fmt::Debug for PeerRpcManager { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("PeerRpcManager") + .field("node_id", &self.tspt.my_peer_id()) + .finish() + } +} + +impl PeerRpcManager { + pub fn new(tspt: impl PeerRpcManagerTransport) -> Self { + Self { + tspt: Arc::new(Box::new(tspt)), + bidirect_rpc: BidirectRpcManager::new(), + tasks: Mutex::new(JoinSet::new()), + } + } + + pub fn new_with_stats_manager(tspt: impl PeerRpcManagerTransport, stats_manager: T) -> Self + where + T: Clone + RpcMetricsProvider, + { + Self { + tspt: Arc::new(Box::new(tspt)), + bidirect_rpc: BidirectRpcManager::new_with_stats_manager(stats_manager), + tasks: Mutex::new(JoinSet::new()), + } + } + + pub fn new_with_metrics(tspt: impl PeerRpcManagerTransport, metrics: ArcRpcMetrics) -> Self { + Self { + tspt: Arc::new(Box::new(tspt)), + bidirect_rpc: BidirectRpcManager::new_with_metrics(metrics), + tasks: Mutex::new(JoinSet::new()), + } + } + + pub fn run(&self) { + let ret = self.bidirect_rpc.run_and_create_tunnel(); + let (mut rx, mut tx) = ret.split(); + let tspt = self.tspt.clone(); + self.tasks.lock().unwrap().spawn(async move { + while let Some(Ok(packet)) = rx.next().await { + let dst_peer_id = packet.peer_manager_header().unwrap().to_peer_id.into(); + if let Err(e) = tspt.send(packet, dst_peer_id).await { + tracing::error!("send to rpc tspt error: {:?}", e); + } + } + }); + + let tspt = self.tspt.clone(); + self.tasks.lock().unwrap().spawn(async move { + while let Ok(packet) = tspt.recv().await { + if let Err(e) = tx.send(packet).await { + tracing::error!("send to rpc tspt error: {:?}", e); + } + } + }); + } + + pub async fn stop(&self) { + self.bidirect_rpc.stop().await; + let mut tasks = { + let mut task_slot = self.tasks.lock().unwrap(); + std::mem::replace(&mut *task_slot, JoinSet::new()) + }; + tasks.abort_all(); + while tasks.join_next().await.is_some() {} + } + + pub fn rpc_client(&self) -> &rpc::client::Client { + self.bidirect_rpc.rpc_client() + } + + pub fn rpc_server(&self) -> &rpc::server::Server { + self.bidirect_rpc.rpc_server() + } + + pub fn my_peer_id(&self) -> PeerId { + self.tspt.my_peer_id() + } +} + +impl Drop for PeerRpcManager { + fn drop(&mut self) { + tracing::debug!("PeerRpcManager drop, my_peer_id: {:?}", self.my_peer_id()); + } +} diff --git a/easytier-core/src/peers/public_ipv6/mod.rs b/easytier-core/src/peers/public_ipv6/mod.rs new file mode 100644 index 00000000..7e9883f2 --- /dev/null +++ b/easytier-core/src/peers/public_ipv6/mod.rs @@ -0,0 +1,738 @@ +pub mod provider; +pub(crate) mod service; + +pub(crate) use service::PublicIpv6Service; + +use std::{collections::HashSet, net::Ipv6Addr, sync::Arc}; + +use cidr::{Ipv6Cidr, Ipv6Inet}; + +use crate::{ + config::PeerId, + config::peers::PublicIpv6ProviderConfig, + config::runtime::CoreRuntimeConfigStore, + events::{CoreEvent, CoreEventSink}, + peers::context::PeerPublicIpv6State, +}; + +impl PublicIpv6ProviderConfig { + pub fn validate(self) -> Result<(), PublicIpv6ProviderConfigError> { + if !self.provider_enabled { + return Ok(()); + } + if !self.provider_supported { + return Err(PublicIpv6ProviderConfigError::UnsupportedProvider); + } + if let Some(prefix) = self.configured_prefix + && !is_global_routable_public_ipv6_prefix(prefix) + { + return Err(PublicIpv6ProviderConfigError::InvalidPrefix(prefix)); + } + Ok(()) + } +} + +#[derive(Debug, thiserror::Error, PartialEq, Eq)] +pub enum PublicIpv6ProviderConfigError { + #[error( + "the provider feature requires Linux; run without --ipv6-public-addr-provider on this node, or move the provider role to a Linux node. client mode (--ipv6-public-addr-auto) works on all platforms" + )] + UnsupportedProvider, + #[error( + "the prefix {0} is not a valid global unicast IPv6 prefix; it must be a routable address range, not a private, link-local, or multicast address" + )] + InvalidPrefix(Ipv6Cidr), +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub(crate) enum PublicIpv6ProviderResolution { + Disabled, + Pending(String), + Active(Ipv6Cidr), +} + +pub(crate) fn resolve_public_ipv6_provider( + config: PublicIpv6ProviderConfig, + detected_prefix: Result, String>, +) -> PublicIpv6ProviderResolution { + if !config.provider_enabled { + return PublicIpv6ProviderResolution::Disabled; + } + if !config.provider_supported { + return PublicIpv6ProviderResolution::Pending( + PublicIpv6ProviderConfigError::UnsupportedProvider.to_string(), + ); + } + + if let Some(prefix) = config.configured_prefix { + return if is_global_routable_public_ipv6_prefix(prefix) { + PublicIpv6ProviderResolution::Active(prefix) + } else { + PublicIpv6ProviderResolution::Pending(format!( + "the configured prefix {prefix} is not a valid global unicast IPv6 prefix" + )) + }; + } + + match detected_prefix { + Ok(Some(prefix)) if is_global_routable_public_ipv6_prefix(prefix) => { + PublicIpv6ProviderResolution::Active(prefix) + } + Ok(Some(prefix)) => PublicIpv6ProviderResolution::Pending(format!( + "the detected prefix {prefix} is not a valid global unicast IPv6 prefix" + )), + Ok(None) => PublicIpv6ProviderResolution::Pending( + "no public IPv6 prefix found on this system; set --ipv6-public-addr-prefix manually, or check that your ISP has delegated an IPv6 prefix and a default-from route exists in the kernel routing table".to_owned(), + ), + Err(error) => PublicIpv6ProviderResolution::Pending(error), + } +} + +pub fn is_global_routable_public_ipv6_prefix(prefix: Ipv6Cidr) -> bool { + let addr = prefix.first_address(); + !addr.is_loopback() + && !addr.is_multicast() + && !addr.is_unicast_link_local() + && !addr.is_unique_local() + && !addr.is_unspecified() +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub(crate) struct PublicIpv6PeerRouteInfo { + pub peer_id: PeerId, + pub inst_id: Option, + pub is_provider: bool, + pub prefix: Option, + pub lease: Option, + pub reachable: bool, +} + +pub(crate) trait PublicIpv6RouteControl: Send + Sync { + fn my_peer_id(&self) -> PeerId; + fn peer_route_snapshot(&self) -> Vec; + fn publish_self_public_ipv6_lease(&self, lease: Option) -> bool; +} + +pub(crate) trait PublicIpv6SyncTrigger: Send + Sync { + fn sync_now(&self, reason: &str); +} + +#[async_trait::async_trait] +pub trait PublicIpv6Host: Send + Sync { + async fn collect_reserved_public_ipv6_addrs(&self, prefix: Ipv6Cidr) -> HashSet; +} + +#[async_trait::async_trait] +impl PublicIpv6Host for () { + async fn collect_reserved_public_ipv6_addrs(&self, _prefix: Ipv6Cidr) -> HashSet { + HashSet::new() + } +} + +#[async_trait::async_trait] +#[auto_impl::auto_impl(Arc)] +pub(crate) trait PublicIpv6Runtime: Send + Sync { + fn ipv6_public_addr_auto(&self) -> bool; + fn ipv6_public_addr_provider(&self) -> bool; + fn instance_id(&self) -> uuid::Uuid; + fn network_name(&self) -> String; + async fn collect_reserved_public_ipv6_addrs(&self, prefix: Ipv6Cidr) -> HashSet; + fn public_ipv6_lease_changed(&self, old: Option, new: Option); + fn public_ipv6_routes_changed(&self, added: Vec, removed: Vec); +} + +pub struct CorePublicIpv6Runtime { + config: CoreRuntimeConfigStore, + host: Arc, + events: Arc, + provider_prefix: std::sync::Mutex>, + lease: std::sync::Mutex>, +} + +impl CorePublicIpv6Runtime { + pub fn new( + config: CoreRuntimeConfigStore, + host: Arc, + events: Arc, + ) -> Arc { + Arc::new(Self { + config, + host, + events, + provider_prefix: std::sync::Mutex::new(None), + lease: std::sync::Mutex::new(None), + }) + } + + pub fn set_provider_prefix(&self, prefix: Option) -> bool { + let mut current = self.provider_prefix.lock().unwrap(); + if *current == prefix { + return false; + } + *current = prefix; + true + } +} + +impl PeerPublicIpv6State for CorePublicIpv6Runtime { + fn public_ipv6_lease_contains(&self, ip: &Ipv6Addr) -> bool { + self.lease + .lock() + .unwrap() + .is_some_and(|lease| lease.address() == *ip) + } + + fn public_ipv6_provider_enabled(&self) -> bool { + self.provider_prefix.lock().unwrap().is_some() + } + + fn advertised_ipv6_public_addr_prefix(&self) -> Option { + *self.provider_prefix.lock().unwrap() + } +} + +#[async_trait::async_trait] +impl PublicIpv6Runtime for CorePublicIpv6Runtime { + fn ipv6_public_addr_auto(&self) -> bool { + self.config.snapshot().services.public_ipv6_auto + } + + fn ipv6_public_addr_provider(&self) -> bool { + self.config + .snapshot() + .services + .public_ipv6_provider + .provider_enabled + } + + fn instance_id(&self) -> uuid::Uuid { + self.config + .snapshot() + .peer + .runtime + .core + .node + .instance_id + .map(uuid::Uuid::from_bytes) + .expect("core peer identity must be finalized before public IPv6 starts") + } + + fn network_name(&self) -> String { + self.config + .snapshot() + .peer + .runtime + .network_identity + .network_name + .clone() + } + + async fn collect_reserved_public_ipv6_addrs(&self, prefix: Ipv6Cidr) -> HashSet { + self.host.collect_reserved_public_ipv6_addrs(prefix).await + } + + fn public_ipv6_lease_changed(&self, old: Option, new: Option) { + *self.lease.lock().unwrap() = new; + self.events + .emit(CoreEvent::PublicIpv6LeaseChanged { old, new }); + } + + fn public_ipv6_routes_changed(&self, added: Vec, removed: Vec) { + self.events + .emit(CoreEvent::PublicIpv6RoutesChanged { added, removed }); + } +} + +pub(super) struct DisabledPublicIpv6Runtime { + instance_id: uuid::Uuid, + network_name: String, +} + +impl DisabledPublicIpv6Runtime { + pub(super) fn new(instance_id: uuid::Uuid, network_name: String) -> Self { + Self { + instance_id, + network_name, + } + } +} + +#[async_trait::async_trait] +impl PublicIpv6Runtime for DisabledPublicIpv6Runtime { + fn ipv6_public_addr_auto(&self) -> bool { + false + } + + fn ipv6_public_addr_provider(&self) -> bool { + false + } + + fn instance_id(&self) -> uuid::Uuid { + self.instance_id + } + + fn network_name(&self) -> String { + self.network_name.clone() + } + + async fn collect_reserved_public_ipv6_addrs(&self, _prefix: Ipv6Cidr) -> HashSet { + HashSet::new() + } + + fn public_ipv6_lease_changed(&self, _old: Option, _new: Option) {} + + fn public_ipv6_routes_changed(&self, _added: Vec, _removed: Vec) {} +} + +#[cfg(test)] +mod tests { + use std::net::Ipv6Addr; + use std::{ + collections::{HashMap, HashSet}, + sync::{Arc, Mutex}, + }; + + use cidr::{Ipv6Cidr, Ipv6Inet}; + + use crate::{ + config::PeerId, + config::runtime::{CoreRuntimeConfig, CoreRuntimeConfigStore}, + events::{CoreEvent, CoreEventSink}, + peers::{context::PeerPublicIpv6State, peer_rpc::PeerRpcManager}, + }; + + use super::{ + CorePublicIpv6Runtime, PublicIpv6Host, PublicIpv6PeerRouteInfo, PublicIpv6ProviderConfig, + PublicIpv6ProviderConfigError, PublicIpv6ProviderResolution, PublicIpv6RouteControl, + PublicIpv6Runtime, PublicIpv6Service, PublicIpv6SyncTrigger, resolve_public_ipv6_provider, + service::allocate_public_ipv6_leases, + }; + + struct TestRouteControl { + my_peer_id: PeerId, + peers: Mutex>, + } + + impl PublicIpv6RouteControl for TestRouteControl { + fn my_peer_id(&self) -> PeerId { + self.my_peer_id + } + + fn peer_route_snapshot(&self) -> Vec { + self.peers.lock().unwrap().clone() + } + + fn publish_self_public_ipv6_lease(&self, _lease: Option) -> bool { + false + } + } + + struct TestSyncTrigger; + + impl PublicIpv6SyncTrigger for TestSyncTrigger { + fn sync_now(&self, _reason: &str) {} + } + + struct TestRuntime { + auto: bool, + provider: bool, + inst_id: uuid::Uuid, + network_name: String, + reserved: Mutex>, + lease: Mutex>, + } + + impl TestRuntime { + fn new(auto: bool) -> Self { + Self { + auto, + provider: false, + inst_id: uuid::Uuid::from_u128(1), + network_name: "default".to_string(), + reserved: Mutex::new(HashSet::new()), + lease: Mutex::new(None), + } + } + } + + #[async_trait::async_trait] + impl PublicIpv6Runtime for TestRuntime { + fn ipv6_public_addr_auto(&self) -> bool { + self.auto + } + + fn ipv6_public_addr_provider(&self) -> bool { + self.provider + } + + fn instance_id(&self) -> uuid::Uuid { + self.inst_id + } + + fn network_name(&self) -> String { + self.network_name.clone() + } + + async fn collect_reserved_public_ipv6_addrs(&self, prefix: Ipv6Cidr) -> HashSet { + self.reserved + .lock() + .unwrap() + .iter() + .copied() + .filter(|addr| prefix.contains(addr)) + .collect() + } + + fn public_ipv6_lease_changed(&self, _old: Option, new: Option) { + *self.lease.lock().unwrap() = new; + } + + fn public_ipv6_routes_changed(&self, _added: Vec, _removed: Vec) {} + } + + #[derive(Default)] + struct RecordingPublicIpv6Host { + reserved: Mutex>, + } + + #[derive(Default)] + struct RecordingPublicIpv6Events { + leases: Mutex, Option)>>, + route_deltas: Mutex, Vec)>>, + } + + #[async_trait::async_trait] + impl PublicIpv6Host for RecordingPublicIpv6Host { + async fn collect_reserved_public_ipv6_addrs(&self, prefix: Ipv6Cidr) -> HashSet { + self.reserved + .lock() + .unwrap() + .iter() + .copied() + .filter(|addr| prefix.contains(addr)) + .collect() + } + } + + impl CoreEventSink for RecordingPublicIpv6Events { + fn emit(&self, event: CoreEvent) { + match event { + CoreEvent::PublicIpv6LeaseChanged { old, new } => { + self.leases.lock().unwrap().push((old, new)); + } + CoreEvent::PublicIpv6RoutesChanged { added, removed } => { + self.route_deltas.lock().unwrap().push((added, removed)); + } + _ => {} + } + } + } + + #[tokio::test] + async fn core_runtime_owns_public_ipv6_state_and_projects_only_host_effects() { + let instance_id = uuid::Uuid::from_u128(42); + let mut peer = crate::config::peers::PeerRuntimeSnapshot::default(); + peer.runtime.core.node.instance_id = Some(*instance_id.as_bytes()); + peer.runtime.network_identity.network_name = "owned-by-core".to_owned(); + let config = CoreRuntimeConfigStore::new( + CoreRuntimeConfig { + public_ipv6_auto: true, + public_ipv6_provider: PublicIpv6ProviderConfig { + provider_enabled: true, + configured_prefix: None, + provider_supported: true, + }, + ..Default::default() + }, + Arc::new(peer), + ); + let host = Arc::new(RecordingPublicIpv6Host::default()); + let events = Arc::new(RecordingPublicIpv6Events::default()); + let reserved = "2001:db8::10".parse().unwrap(); + host.reserved.lock().unwrap().insert(reserved); + let runtime = CorePublicIpv6Runtime::new(config.clone(), host.clone(), events.clone()); + let prefix = "2001:db8::/64".parse().unwrap(); + let lease = "2001:db8::20/64".parse().unwrap(); + let route = "2001:db8::30/128".parse().unwrap(); + + assert!(runtime.ipv6_public_addr_auto()); + assert!(runtime.ipv6_public_addr_provider()); + assert_eq!(runtime.instance_id(), instance_id); + assert_eq!(runtime.network_name(), "owned-by-core"); + assert_eq!( + runtime.collect_reserved_public_ipv6_addrs(prefix).await, + HashSet::from([reserved]) + ); + assert!(runtime.set_provider_prefix(Some(prefix))); + assert!(!runtime.set_provider_prefix(Some(prefix))); + assert_eq!(runtime.advertised_ipv6_public_addr_prefix(), Some(prefix)); + + runtime.public_ipv6_lease_changed(None, Some(lease)); + assert!(runtime.public_ipv6_lease_contains(&lease.address())); + runtime.public_ipv6_routes_changed(vec![route], Vec::new()); + assert_eq!( + events.leases.lock().unwrap().as_slice(), + &[(None, Some(lease))] + ); + assert_eq!( + events.route_deltas.lock().unwrap().as_slice(), + &[(vec![route], Vec::new())] + ); + + config.update_services(|services| { + services.public_ipv6_auto = false; + services.public_ipv6_provider.provider_enabled = false; + }); + assert!(!runtime.ipv6_public_addr_auto()); + assert!(!runtime.ipv6_public_addr_provider()); + } + + #[test] + fn public_ipv6_lease_allocator_keeps_stable_addresses() { + let prefix = "2001:db8::/124".parse::().unwrap(); + let first = uuid::Uuid::from_u128(1); + let second = uuid::Uuid::from_u128(2); + + let leases = + allocate_public_ipv6_leases(prefix, &[first, second], &HashSet::new(), &HashMap::new()); + assert_eq!(leases.len(), 2); + assert_ne!(leases[0].addr, leases[1].addr); + + let initial_map = HashMap::from([(first, leases[0].addr)]); + let next = allocate_public_ipv6_leases(prefix, &[first], &HashSet::new(), &initial_map); + assert_eq!(next.len(), 1); + assert_eq!(next[0].addr, leases[0].addr); + assert!(next[0].reused); + } + + #[test] + fn public_ipv6_provider_prefers_smallest_instance_id() { + let info_a = PublicIpv6PeerRouteInfo { + peer_id: 2, + inst_id: Some(uuid::Uuid::from_u128(2)), + is_provider: true, + prefix: Some("2001:db8:1::/120".parse().unwrap()), + lease: None, + reachable: true, + }; + let info_b = PublicIpv6PeerRouteInfo { + peer_id: 1, + inst_id: Some(uuid::Uuid::from_u128(1)), + is_provider: true, + prefix: Some("2001:db8:2::/120".parse().unwrap()), + lease: None, + reachable: true, + }; + + let selected = + PublicIpv6Service::selected_provider_from_snapshot(&[info_a, info_b]).unwrap(); + assert_eq!(selected.peer_id, 1); + } + + #[test] + fn public_ipv6_provider_prefers_reachable_provider() { + let unreachable_lower_id = PublicIpv6PeerRouteInfo { + peer_id: 1, + inst_id: Some(uuid::Uuid::from_u128(1)), + is_provider: true, + prefix: Some("2001:db8:1::/120".parse().unwrap()), + lease: None, + reachable: false, + }; + let reachable_higher_id = PublicIpv6PeerRouteInfo { + peer_id: 2, + inst_id: Some(uuid::Uuid::from_u128(2)), + is_provider: true, + prefix: Some("2001:db8:2::/120".parse().unwrap()), + lease: None, + reachable: true, + }; + + let selected = PublicIpv6Service::selected_provider_from_snapshot(&[ + unreachable_lower_id, + reachable_higher_id, + ]) + .unwrap(); + assert_eq!(selected.peer_id, 2); + } + + #[test] + fn public_ipv6_lease_allocator_stops_when_only_network_offset_is_left() { + let prefix = "2001:db8::/126".parse::().unwrap(); + let network = prefix.first_address(); + let reserved = HashSet::from([ + Ipv6Addr::from(u128::from(network) + 1), + Ipv6Addr::from(u128::from(network) + 2), + Ipv6Addr::from(u128::from(network) + 3), + ]); + + let leases = allocate_public_ipv6_leases( + prefix, + &[uuid::Uuid::from_u128(42)], + &reserved, + &HashMap::new(), + ); + + assert!(leases.is_empty()); + } + + #[tokio::test] + async fn reconcile_runtime_clears_public_ipv6_lease_when_auto_is_disabled() { + let stale_addr = "2001:db8::123/64".parse().unwrap(); + let runtime = Arc::new(TestRuntime::new(false)); + *runtime.lease.lock().unwrap() = Some(stale_addr); + + let service = Arc::new(PublicIpv6Service::new( + runtime.clone(), + std::sync::Weak::::new(), + Arc::new(TestRouteControl { + my_peer_id: 1, + peers: Mutex::new(Vec::new()), + }), + Arc::new(TestSyncTrigger), + )); + *service.my_addr_cache.lock().unwrap() = Some(stale_addr); + + service.reconcile_runtime_from_snapshot(&[]); + + assert_eq!(*service.my_addr_cache.lock().unwrap(), None); + assert_eq!(*runtime.lease.lock().unwrap(), None); + } + + #[tokio::test] + async fn reconcile_runtime_updates_public_lease_when_auto_enabled() { + let public_addr = "2001:db8::123/64".parse().unwrap(); + let runtime = Arc::new(TestRuntime::new(true)); + + let service = Arc::new(PublicIpv6Service::new( + runtime.clone(), + std::sync::Weak::::new(), + Arc::new(TestRouteControl { + my_peer_id: 1, + peers: Mutex::new(vec![PublicIpv6PeerRouteInfo { + peer_id: 1, + inst_id: Some(uuid::Uuid::from_u128(1)), + is_provider: false, + prefix: None, + lease: Some(public_addr), + reachable: true, + }]), + }), + Arc::new(TestSyncTrigger), + )); + + service.reconcile_runtime(); + + assert_eq!(*runtime.lease.lock().unwrap(), Some(public_addr)); + } + + #[test] + fn provider_config_uses_explicit_host_capability() { + let unsupported = PublicIpv6ProviderConfig { + provider_enabled: true, + configured_prefix: None, + provider_supported: false, + }; + assert_eq!( + unsupported.validate(), + Err(PublicIpv6ProviderConfigError::UnsupportedProvider) + ); + + let disabled = PublicIpv6ProviderConfig { + provider_enabled: false, + ..unsupported + }; + assert!(disabled.validate().is_ok()); + assert!(!disabled.should_run_reconcile()); + } + + #[test] + fn provider_config_rejects_non_global_prefixes() { + for prefix in ["::1/128", "fe80::/64", "fd00::/48", "ff00::/8", "::/0"] { + let config = PublicIpv6ProviderConfig { + provider_enabled: true, + configured_prefix: Some(prefix.parse().unwrap()), + provider_supported: true, + }; + assert!(matches!( + config.validate(), + Err(PublicIpv6ProviderConfigError::InvalidPrefix(_)) + )); + } + } + + #[test] + fn provider_config_accepts_global_prefix() { + let config = PublicIpv6ProviderConfig { + provider_enabled: true, + configured_prefix: Some("2001:db8::/48".parse().unwrap()), + provider_supported: true, + }; + assert!(config.validate().is_ok()); + assert!(config.should_run_reconcile()); + } + + #[test] + fn provider_resolution_prefers_configured_prefix_without_detection() { + let prefix = "2001:db8::/48".parse().unwrap(); + let config = PublicIpv6ProviderConfig { + provider_enabled: true, + configured_prefix: Some(prefix), + provider_supported: true, + }; + assert_eq!( + resolve_public_ipv6_provider(config, Err("detection must be ignored".to_owned())), + PublicIpv6ProviderResolution::Active(prefix) + ); + } + + #[test] + fn provider_resolution_normalizes_auto_detection_results() { + let config = PublicIpv6ProviderConfig { + provider_enabled: true, + configured_prefix: None, + provider_supported: true, + }; + let prefix = "2001:db8:1::/56".parse().unwrap(); + assert_eq!( + resolve_public_ipv6_provider(config, Ok(Some(prefix))), + PublicIpv6ProviderResolution::Active(prefix) + ); + assert!(matches!( + resolve_public_ipv6_provider(config, Ok(None)), + PublicIpv6ProviderResolution::Pending(message) + if message.contains("ipv6-public-addr-prefix") + )); + assert_eq!( + resolve_public_ipv6_provider(config, Err("host detection failed".to_owned())), + PublicIpv6ProviderResolution::Pending("host detection failed".to_owned()) + ); + } + + #[test] + fn provider_resolution_rejects_invalid_configured_and_detected_prefixes() { + let configured = PublicIpv6ProviderConfig { + provider_enabled: true, + configured_prefix: Some("fd00::/48".parse().unwrap()), + provider_supported: true, + }; + assert!(matches!( + resolve_public_ipv6_provider(configured, Ok(None)), + PublicIpv6ProviderResolution::Pending(message) + if message.contains("configured prefix") + )); + + let detected = PublicIpv6ProviderConfig { + configured_prefix: None, + ..configured + }; + assert!(matches!( + resolve_public_ipv6_provider( + detected, + Ok(Some("fe80::/64".parse().unwrap())) + ), + PublicIpv6ProviderResolution::Pending(message) + if message.contains("detected prefix") + )); + } +} diff --git a/easytier-core/src/peers/public_ipv6/provider.rs b/easytier-core/src/peers/public_ipv6/provider.rs new file mode 100644 index 00000000..8277bb11 --- /dev/null +++ b/easytier-core/src/peers/public_ipv6/provider.rs @@ -0,0 +1,601 @@ +use std::{ + sync::{ + Arc, Weak, + atomic::{AtomicBool, Ordering}, + }, + time::Duration, +}; + +use async_trait::async_trait; +use cidr::Ipv6Cidr; +use tokio::sync::Mutex; +use tokio_util::sync::CancellationToken; + +use crate::{ + config::peers::PublicIpv6ProviderConfig, + config::runtime::CoreRuntimeConfigStore, + peers::public_ipv6::{ + CorePublicIpv6Runtime, PublicIpv6ProviderResolution, resolve_public_ipv6_provider, + }, +}; + +const DEFAULT_RECONCILE_INTERVAL: Duration = Duration::from_secs(5); +const MAX_CONFIG_RETRIES: usize = 3; + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct PublicIpv6NdpTarget { + pub wan_interface: String, +} + +#[derive(Debug, Clone, Default, PartialEq, Eq)] +pub struct PublicIpv6PlatformObservation { + pub detected_prefix: Option, + pub ndp_target: Option, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct PublicIpv6NdpDesired { + pub prefix: Ipv6Cidr, + pub target: PublicIpv6NdpTarget, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +enum PublicIpv6ProviderState { + Disabled, + Pending(String), + Active { + prefix: Ipv6Cidr, + ndp_target: Option, + }, +} + +impl PublicIpv6ProviderState { + fn from_resolution( + resolution: PublicIpv6ProviderResolution, + ndp_target: Option, + ) -> Self { + match resolution { + PublicIpv6ProviderResolution::Disabled => Self::Disabled, + PublicIpv6ProviderResolution::Pending(error) => Self::Pending(error), + PublicIpv6ProviderResolution::Active(prefix) => Self::Active { prefix, ndp_target }, + } + } + + fn advertised_prefix(&self) -> Option { + match self { + Self::Active { prefix, .. } => Some(*prefix), + Self::Disabled | Self::Pending(_) => None, + } + } + + fn ndp_desired(&self) -> Option { + match self { + Self::Active { + prefix, + ndp_target: Some(target), + } => Some(PublicIpv6NdpDesired { + prefix: *prefix, + target: target.clone(), + }), + Self::Disabled | Self::Pending(_) | Self::Active { .. } => None, + } + } +} + +#[derive(Debug, Clone, thiserror::Error, PartialEq, Eq)] +pub enum PublicIpv6PlatformError { + #[error("public IPv6 platform adapter is unavailable")] + Unavailable, + #[error("{0}")] + Failed(String), +} + +#[async_trait] +pub trait PublicIpv6ProviderPlatform: Send + Sync + 'static { + fn inspect( + &self, + config: PublicIpv6ProviderConfig, + ) -> Result; + + fn sync_ndp( + &self, + desired: Option, + ) -> Result<(), PublicIpv6PlatformError>; + + /// Waits for a platform state change that requires an immediate retry. + /// Returns `false` when the event source has closed. + async fn wait_for_change(&self) -> bool; +} + +struct PublicIpv6ProviderTask { + cancel: CancellationToken, + handle: tokio::task::JoinHandle<()>, +} + +pub struct PublicIpv6ProviderService { + platform: Arc, + runtime_config: CoreRuntimeConfigStore, + runtime: Arc, + reconcile_interval: Duration, + reconcile: Mutex<()>, + last_state: std::sync::Mutex>, + task: Mutex>, + closing: AtomicBool, +} + +#[cfg(feature = "public-ipv6-provider")] +pub(crate) struct PublicIpv6ProviderRuntime { + service: Option>, + runtime_config: CoreRuntimeConfigStore, +} + +#[cfg(feature = "public-ipv6-provider")] +impl PublicIpv6ProviderRuntime { + pub(crate) fn new( + platform: Option>, + runtime_config: CoreRuntimeConfigStore, + runtime: Arc, + ) -> Self { + let service = platform.map(|platform| { + PublicIpv6ProviderService::new(platform, runtime_config.clone(), runtime) + }); + Self { + service, + runtime_config, + } + } + + pub(crate) async fn validate_before_start(&self) -> anyhow::Result<()> { + let config = self.runtime_config.snapshot().services.public_ipv6_provider; + config.validate().map_err(anyhow::Error::new)?; + if config.provider_enabled && self.service.is_none() { + anyhow::bail!("public IPv6 provider is enabled but no host adapter was provided"); + } + if let Some(service) = &self.service { + service.apply_config().await; + } + Ok(()) + } + + pub(crate) async fn start(&self) { + if let Some(service) = &self.service { + service.start().await; + } + } + + pub(crate) async fn stop(&self) { + if let Some(service) = &self.service { + service.stop().await; + } + } + + pub(crate) async fn reconcile(&self) -> bool { + let Some(service) = &self.service else { + return false; + }; + let applied = service.apply_config().await; + service.start().await; + applied + } +} + +impl PublicIpv6ProviderService { + pub fn new( + platform: Arc, + runtime_config: CoreRuntimeConfigStore, + runtime: Arc, + ) -> Arc { + Self::new_with_interval( + platform, + runtime_config, + runtime, + DEFAULT_RECONCILE_INTERVAL, + ) + } + + fn new_with_interval( + platform: Arc, + runtime_config: CoreRuntimeConfigStore, + runtime: Arc, + reconcile_interval: Duration, + ) -> Arc { + Arc::new(Self { + platform, + runtime_config, + runtime, + reconcile_interval, + reconcile: Mutex::new(()), + last_state: std::sync::Mutex::new(None), + task: Mutex::new(None), + closing: AtomicBool::new(false), + }) + } + + pub async fn reconcile_now(&self) -> bool { + let _reconcile = self.reconcile.lock().await; + if self.closing.load(Ordering::Acquire) { + return false; + } + for attempt in 0..MAX_CONFIG_RETRIES { + let config = self.runtime_config.snapshot().services.public_ipv6_provider; + let observation = if config.provider_enabled && config.provider_supported { + match self.platform.inspect(config) { + Ok(observation) => Ok(observation), + Err(PublicIpv6PlatformError::Unavailable) => return false, + Err(PublicIpv6PlatformError::Failed(error)) => Err(error), + } + } else { + Ok(PublicIpv6PlatformObservation::default()) + }; + let (detected_prefix, ndp_target) = match observation { + Ok(observation) => (Ok(observation.detected_prefix), observation.ndp_target), + Err(error) => (Err(error), None), + }; + let next_state = PublicIpv6ProviderState::from_resolution( + resolve_public_ipv6_provider(config, detected_prefix), + ndp_target, + ); + + if self.runtime_config.snapshot().services.public_ipv6_provider != config { + tracing::debug!( + attempt = attempt + 1, + max_retries = MAX_CONFIG_RETRIES, + "public IPv6 provider config changed during reconcile, retrying" + ); + continue; + } + + let changed = self + .runtime + .set_provider_prefix(next_state.advertised_prefix()); + if let Err(error) = self.platform.sync_ndp(next_state.ndp_desired()) { + match error { + PublicIpv6PlatformError::Unavailable => return false, + PublicIpv6PlatformError::Failed(error) => { + tracing::warn!(%error, "failed to synchronize public IPv6 NDP state"); + } + } + } + self.log_state_change(&next_state, changed); + *self.last_state.lock().unwrap() = Some(next_state); + return true; + } + + tracing::warn!( + max_retries = MAX_CONFIG_RETRIES, + "skipping public IPv6 provider reconcile because config kept changing" + ); + true + } + + pub async fn apply_config(&self) -> bool { + self.reconcile_now().await + } + + pub async fn start(self: &Arc) { + let mut task = self.task.lock().await; + let config = self.runtime_config.snapshot().services.public_ipv6_provider; + if self.closing.load(Ordering::Acquire) || task.is_some() || !config.should_run_reconcile() + { + return; + } + + let cancel = CancellationToken::new(); + let task_cancel = cancel.clone(); + let service = Arc::downgrade(self); + let platform = self.platform.clone(); + let handle = tokio::spawn(async move { + Self::run(service, platform, task_cancel).await; + }); + task.replace(PublicIpv6ProviderTask { cancel, handle }); + } + + pub async fn stop(&self) { + self.closing.store(true, Ordering::Release); + let task = self.task.lock().await.take(); + if let Some(task) = task { + task.cancel.cancel(); + if let Err(error) = task.handle.await { + tracing::warn!(?error, "public IPv6 provider task failed during shutdown"); + } + } + let _reconcile = self.reconcile.lock().await; + if let Err(error) = self.platform.sync_ndp(None) + && !matches!(error, PublicIpv6PlatformError::Unavailable) + { + tracing::warn!(%error, "failed to clean public IPv6 NDP state during shutdown"); + } + } + + fn log_state_change(&self, next_state: &PublicIpv6ProviderState, changed: bool) { + let last_state = self.last_state.lock().unwrap(); + if last_state.as_ref() != Some(next_state) { + match next_state { + PublicIpv6ProviderState::Disabled if last_state.is_some() => { + tracing::info!("public IPv6 provider disabled"); + } + PublicIpv6ProviderState::Disabled => {} + PublicIpv6ProviderState::Pending(reason) => { + tracing::warn!(%reason, "public IPv6 provider not ready"); + } + PublicIpv6ProviderState::Active { prefix, ndp_target } => { + if let Some(target) = ndp_target { + tracing::info!( + %prefix, + wan_interface = %target.wan_interface, + "public IPv6 provider is active with NDP proxy" + ); + } else { + tracing::info!(%prefix, "public IPv6 provider is active"); + } + } + } + } else if changed { + tracing::info!("public IPv6 provider runtime state changed"); + } + } + + async fn run( + service: Weak, + platform: Arc, + cancel: CancellationToken, + ) { + loop { + let Some(service) = service.upgrade() else { + let _ = platform.sync_ndp(None); + return; + }; + if !service.reconcile_now().await { + let _ = service.platform.sync_ndp(None); + return; + } + + let interval = service.reconcile_interval; + drop(service); + let should_continue = tokio::select! { + _ = cancel.cancelled() => false, + _ = crate::foundation::time::sleep(interval) => true, + changed = platform.wait_for_change() => changed, + }; + if !should_continue { + let _ = platform.sync_ndp(None); + return; + } + } + } +} + +#[cfg(test)] +mod tests { + use std::sync::{ + Mutex as StdMutex, + atomic::{AtomicUsize, Ordering}, + }; + + use tokio::sync::Notify; + + use super::*; + use crate::{ + config::peers::PeerRuntimeSnapshot, + config::runtime::{CoreRuntimeConfig, CoreRuntimeConfigStore}, + peers::context::PeerPublicIpv6State, + }; + + struct RecordingHost { + observation: StdMutex>, + inspect_calls: AtomicUsize, + ndp_desired: StdMutex>>, + change: Notify, + } + + impl Default for RecordingHost { + fn default() -> Self { + Self { + observation: StdMutex::new(Ok(PublicIpv6PlatformObservation::default())), + inspect_calls: AtomicUsize::new(0), + ndp_desired: StdMutex::new(Vec::new()), + change: Notify::new(), + } + } + } + + #[async_trait] + impl PublicIpv6ProviderPlatform for RecordingHost { + fn inspect( + &self, + _config: PublicIpv6ProviderConfig, + ) -> Result { + self.inspect_calls.fetch_add(1, Ordering::AcqRel); + self.observation.lock().unwrap().clone() + } + + async fn wait_for_change(&self) -> bool { + self.change.notified().await; + true + } + + fn sync_ndp( + &self, + desired: Option, + ) -> Result<(), PublicIpv6PlatformError> { + self.ndp_desired.lock().unwrap().push(desired); + Ok(()) + } + } + + fn provider_config(enabled: bool, prefix: Option) -> PublicIpv6ProviderConfig { + PublicIpv6ProviderConfig { + provider_enabled: enabled, + configured_prefix: prefix, + provider_supported: true, + } + } + + fn runtime_config(config: PublicIpv6ProviderConfig) -> CoreRuntimeConfigStore { + let services = CoreRuntimeConfig { + public_ipv6_provider: config, + ..Default::default() + }; + CoreRuntimeConfigStore::new(services, Arc::new(PeerRuntimeSnapshot::default())) + } + + fn runtime( + config: PublicIpv6ProviderConfig, + ) -> (CoreRuntimeConfigStore, Arc) { + let config = runtime_config(config); + let runtime = CorePublicIpv6Runtime::new(config.clone(), Arc::new(()), Arc::new(())); + (config, runtime) + } + + async fn wait_for_calls(counter: &AtomicUsize, expected: usize) { + crate::foundation::time::timeout(Duration::from_secs(1), async { + while counter.load(Ordering::Acquire) < expected { + tokio::task::yield_now().await; + } + }) + .await + .unwrap(); + } + + #[tokio::test] + async fn starts_only_when_enabled_and_reacts_to_host_changes() { + let host = Arc::new(RecordingHost::default()); + let (runtime_config, runtime) = runtime(provider_config(false, None)); + let service = PublicIpv6ProviderService::new_with_interval( + host.clone(), + runtime_config.clone(), + runtime, + Duration::from_secs(60), + ); + + service.start().await; + assert_eq!(host.inspect_calls.load(Ordering::Acquire), 0); + + runtime_config.update_services(|services| { + services.public_ipv6_provider = + provider_config(true, Some("2001:db8::/48".parse().unwrap())); + }); + service.start().await; + wait_for_calls(&host.inspect_calls, 1).await; + host.change.notify_one(); + wait_for_calls(&host.inspect_calls, 2).await; + + service.stop().await; + assert_eq!(host.ndp_desired.lock().unwrap().last(), Some(&None)); + } + + #[tokio::test] + async fn does_not_reconcile_or_restart_after_stop() { + let host = Arc::new(RecordingHost::default()); + let (runtime_config, runtime) = runtime(provider_config( + true, + Some("2001:db8::/48".parse().unwrap()), + )); + let service = PublicIpv6ProviderService::new_with_interval( + host.clone(), + runtime_config, + runtime, + Duration::from_secs(60), + ); + + service.stop().await; + assert!(!service.reconcile_now().await); + service.start().await; + + assert_eq!(host.inspect_calls.load(Ordering::Acquire), 0); + assert_eq!(host.ndp_desired.lock().unwrap().as_slice(), &[None]); + } + + #[tokio::test] + async fn resolves_observation_and_publishes_ndp_desired_state() { + let host = Arc::new(RecordingHost::default()); + let prefix = "2001:db8::/48".parse().unwrap(); + let target = PublicIpv6NdpTarget { + wan_interface: "wan0".to_owned(), + }; + *host.observation.lock().unwrap() = Ok(PublicIpv6PlatformObservation { + detected_prefix: Some(prefix), + ndp_target: Some(target.clone()), + }); + let (runtime_config, runtime) = runtime(provider_config(true, None)); + let service = PublicIpv6ProviderService::new(host.clone(), runtime_config, runtime.clone()); + + assert!(service.reconcile_now().await); + + assert_eq!(runtime.advertised_ipv6_public_addr_prefix(), Some(prefix)); + assert!(runtime.public_ipv6_provider_enabled()); + assert_eq!( + host.ndp_desired.lock().unwrap().as_slice(), + &[Some(PublicIpv6NdpDesired { prefix, target })] + ); + } + + #[tokio::test] + async fn turns_platform_failure_into_pending_provider_state() { + let host = Arc::new(RecordingHost::default()); + *host.observation.lock().unwrap() = Err(PublicIpv6PlatformError::Failed( + "route query failed".to_owned(), + )); + let (runtime_config, runtime) = runtime(provider_config(true, None)); + let service = PublicIpv6ProviderService::new(host.clone(), runtime_config, runtime.clone()); + + assert!(service.reconcile_now().await); + + assert_eq!(runtime.advertised_ipv6_public_addr_prefix(), None); + assert!(!runtime.public_ipv6_provider_enabled()); + assert_eq!(host.ndp_desired.lock().unwrap().as_slice(), &[None]); + } + + struct ReconfiguringHost { + runtime_config: CoreRuntimeConfigStore, + replacement: PublicIpv6ProviderConfig, + inspect_calls: AtomicUsize, + } + + #[async_trait] + impl PublicIpv6ProviderPlatform for ReconfiguringHost { + fn inspect( + &self, + _config: PublicIpv6ProviderConfig, + ) -> Result { + if self.inspect_calls.fetch_add(1, Ordering::AcqRel) == 0 { + self.runtime_config.update_services(|services| { + services.public_ipv6_provider = self.replacement; + }); + } + Ok(PublicIpv6PlatformObservation::default()) + } + + fn sync_ndp( + &self, + _desired: Option, + ) -> Result<(), PublicIpv6PlatformError> { + Ok(()) + } + + async fn wait_for_change(&self) -> bool { + false + } + } + + #[tokio::test] + async fn retries_when_config_changes_during_platform_inspection() { + let first = provider_config(true, Some("2001:db8:1::/48".parse().unwrap())); + let replacement = provider_config(true, Some("2001:db8:2::/48".parse().unwrap())); + let (runtime_config, runtime) = runtime(first); + let host = Arc::new(ReconfiguringHost { + runtime_config: runtime_config.clone(), + replacement, + inspect_calls: AtomicUsize::new(0), + }); + let service = PublicIpv6ProviderService::new(host.clone(), runtime_config, runtime.clone()); + + assert!(service.reconcile_now().await); + + assert_eq!(host.inspect_calls.load(Ordering::Acquire), 2); + assert_eq!( + runtime.advertised_ipv6_public_addr_prefix(), + replacement.configured_prefix + ); + } +} diff --git a/easytier/src/peers/public_ipv6.rs b/easytier-core/src/peers/public_ipv6/service.rs similarity index 72% rename from easytier/src/peers/public_ipv6.rs rename to easytier-core/src/peers/public_ipv6/service.rs index 90432b4e..37cacb2d 100644 --- a/easytier/src/peers/public_ipv6.rs +++ b/easytier-core/src/peers/public_ipv6/service.rs @@ -1,3 +1,7 @@ +//! Lease-driven public IPv6 service: the lease allocator, the per-instance +//! service driving acquisition/renewal, and the RPC server serving lease +//! requests from client peers. + use std::{ collections::{BTreeMap, BTreeSet, HashMap, HashSet}, net::Ipv6Addr, @@ -8,11 +12,10 @@ use std::{ use cidr::{Ipv6Cidr, Ipv6Inet}; use crate::{ - common::{ - PeerId, - global_ctx::{ArcGlobalCtx, GlobalCtxEvent}, - }, + config::PeerId, + peers::peer_rpc::PeerRpcManager, proto::{ + common::Void, peer_rpc::{ AcquireIpv6PublicAddrLeaseRequest, GetIpv6PublicAddrLeaseRequest, Ipv6PublicAddrLeaseReply, PublicIpv6AddrRpc, PublicIpv6AddrRpcClientFactory, @@ -25,7 +28,9 @@ use crate::{ }, }; -use super::peer_rpc::PeerRpcManager; +use super::{ + PublicIpv6PeerRouteInfo, PublicIpv6RouteControl, PublicIpv6Runtime, PublicIpv6SyncTrigger, +}; // Use a longer lease with an early renew window to reduce steady-state RPC // churn while preserving enough margin for transient provider failures. @@ -61,28 +66,8 @@ struct PublicIpv6ClientState { last_error: Option, } -#[derive(Debug, Clone, PartialEq, Eq)] -pub(crate) struct PublicIpv6PeerRouteInfo { - pub peer_id: PeerId, - pub inst_id: Option, - pub is_provider: bool, - pub prefix: Option, - pub lease: Option, - pub reachable: bool, -} - -pub(crate) trait PublicIpv6RouteControl: Send + Sync { - fn my_peer_id(&self) -> PeerId; - fn peer_route_snapshot(&self) -> Vec; - fn publish_self_public_ipv6_lease(&self, lease: Option) -> bool; -} - -pub(crate) trait PublicIpv6SyncTrigger: Send + Sync { - fn sync_now(&self, reason: &str); -} - pub(crate) struct PublicIpv6Service { - global_ctx: ArcGlobalCtx, + runtime: Arc, peer_rpc: Weak, route_control: Arc, sync_trigger: Arc, @@ -90,18 +75,18 @@ pub(crate) struct PublicIpv6Service { provider_state: std::sync::Mutex>, client_state: std::sync::Mutex>, route_cache: std::sync::Mutex>, - my_addr_cache: std::sync::Mutex>, + pub(super) my_addr_cache: std::sync::Mutex>, } impl PublicIpv6Service { - pub(crate) fn new( - global_ctx: ArcGlobalCtx, + pub fn new( + runtime: Arc, peer_rpc: Weak, route_control: Arc, sync_trigger: Arc, ) -> Self { Self { - global_ctx, + runtime, peer_rpc, route_control, sync_trigger, @@ -112,7 +97,7 @@ impl PublicIpv6Service { } } - pub(crate) fn rpc_server(self: &Arc) -> PublicIpv6AddrRpcServerImpl { + pub fn rpc_server(self: &Arc) -> PublicIpv6AddrRpcServerImpl { PublicIpv6AddrRpcServerImpl { service: Arc::downgrade(self), } @@ -152,7 +137,7 @@ impl PublicIpv6Service { true } - fn selected_provider_from_snapshot( + pub(super) fn selected_provider_from_snapshot( peers: &[PublicIpv6PeerRouteInfo], ) -> Option { peers @@ -216,9 +201,9 @@ impl PublicIpv6Service { (my_addr, routes) } - fn reconcile_runtime_from_snapshot(&self, peers: &[PublicIpv6PeerRouteInfo]) { + pub(super) fn reconcile_runtime_from_snapshot(&self, peers: &[PublicIpv6PeerRouteInfo]) { let (mut my_addr, routes) = self.collect_runtime_from_snapshot(peers); - if !self.global_ctx.config.get_ipv6_public_addr_auto() { + if !self.runtime.ipv6_public_addr_auto() { my_addr = None; } @@ -226,9 +211,7 @@ impl PublicIpv6Service { if *cached_my_addr != my_addr { let old = *cached_my_addr; *cached_my_addr = my_addr; - self.global_ctx.set_public_ipv6_lease(my_addr); - self.global_ctx - .issue_event(GlobalCtxEvent::PublicIpv6Changed(old, my_addr)); + self.runtime.public_ipv6_lease_changed(old, my_addr); } drop(cached_my_addr); @@ -243,19 +226,16 @@ impl PublicIpv6Service { .copied() .collect::>(); *cached_routes = routes; - self.global_ctx - .set_public_ipv6_routes(cached_routes.clone()); - self.global_ctx - .issue_event(GlobalCtxEvent::PublicIpv6RoutesUpdated(added, removed)); + self.runtime.public_ipv6_routes_changed(added, removed); } } - fn reconcile_runtime(&self) { + pub(super) fn reconcile_runtime(&self) { let peers = self.route_control.peer_route_snapshot(); self.reconcile_runtime_from_snapshot(&peers); } - pub(crate) fn handle_route_change(&self) -> bool { + pub fn handle_route_change(&self) -> bool { let peers = self.route_control.peer_route_snapshot(); let provider = Self::selected_provider_from_snapshot(&peers); let _provider_changed = self.clear_provider_state_if_provider_changed(provider.as_ref()); @@ -325,23 +305,9 @@ impl PublicIpv6Service { } async fn collect_reserved_addrs(&self, prefix: Ipv6Cidr) -> HashSet { - let mut reserved = HashSet::new(); - let ip_list = self.global_ctx.get_ip_collector().collect_ip_addrs().await; - reserved.extend( - ip_list - .interface_ipv6s - .into_iter() - .map(Ipv6Addr::from) - .filter(|addr| prefix.contains(addr)), - ); - reserved.extend( - ip_list - .public_ipv6 - .into_iter() - .map(Ipv6Addr::from) - .filter(|addr| prefix.contains(addr)), - ); - reserved + self.runtime + .collect_reserved_public_ipv6_addrs(prefix) + .await } fn prune_expired_leases( @@ -463,7 +429,7 @@ impl PublicIpv6Service { Ok((provider, lease.clone())) } - pub(crate) async fn gc_provider_leases(&self) { + pub async fn gc_provider_leases(&self) { let peers = self.route_control.peer_route_snapshot(); let provider = Self::selected_provider_from_snapshot(&peers); self.clear_provider_state_if_provider_changed(provider.as_ref()); @@ -479,8 +445,8 @@ impl PublicIpv6Service { self.set_provider_state(Some(state)); } - pub(crate) async fn sync_client_state(&self) -> bool { - if !self.global_ctx.config.get_ipv6_public_addr_auto() { + pub async fn sync_client_state(&self) -> bool { + if !self.runtime.ipv6_public_addr_auto() { return self .clear_client_lease_state(self.clear_client_state_if_provider_changed(None)); } @@ -525,10 +491,10 @@ impl PublicIpv6Service { .scoped_client::>( self.my_peer_id(), provider.peer_id, - self.global_ctx.get_network_name(), + self.runtime.network_name(), ); - let inst_id = self.global_ctx.get_id(); + let inst_id = self.runtime.instance_id(); let reply = if let Some(state) = current.as_ref().filter(|state| state.provider == provider) { match rpc_stub @@ -615,39 +581,39 @@ impl PublicIpv6Service { peer_info_changed } - pub(crate) async fn provider_gc_routine(self: Arc) { - if !self.global_ctx.config.get_ipv6_public_addr_provider() { + pub async fn provider_gc_routine(self: Arc) { + if !self.runtime.ipv6_public_addr_provider() { return; } loop { - tokio::time::sleep(Duration::from_secs(15)).await; + crate::foundation::time::sleep(Duration::from_secs(15)).await; self.gc_provider_leases().await; } } - pub(crate) async fn client_routine(self: Arc) { + pub async fn client_routine(self: Arc) { loop { if self.sync_client_state().await { self.sync_trigger.sync_now("sync_public_ipv6_client_state"); } - tokio::time::sleep(Duration::from_secs(5)).await; + crate::foundation::time::sleep(Duration::from_secs(5)).await; } } - pub(crate) fn list_routes(&self) -> BTreeSet { + pub fn list_routes(&self) -> BTreeSet { self.route_cache.lock().unwrap().clone() } - pub(crate) fn my_addr(&self) -> Option { + pub fn my_addr(&self) -> Option { *self.my_addr_cache.lock().unwrap() } - pub(crate) fn provider_peer_id_for_client(&self) -> Option { + pub fn provider_peer_id_for_client(&self) -> Option { self.current_client_state() .map(|state| state.provider.peer_id) } - pub(crate) fn local_provider_state( + pub fn local_provider_state( &self, ) -> Option<(PublicIpv6Provider, Vec)> { let provider = self.selected_provider()?; @@ -749,7 +715,7 @@ impl PublicIpv6AddrRpc for PublicIpv6AddrRpcServerImpl { &self, _: BaseController, request: ReleaseIpv6PublicAddrLeaseRequest, - ) -> rpc_types::error::Result { + ) -> rpc_types::error::Result { let Some(service) = self.service.upgrade() else { return Err(anyhow::anyhow!("public ipv6 service stopped").into()); }; @@ -784,7 +750,7 @@ impl PublicIpv6AddrRpc for PublicIpv6AddrRpcServerImpl { } } -fn allocate_public_ipv6_leases( +pub(super) fn allocate_public_ipv6_leases( prefix: Ipv6Cidr, auto_peer_ids: &[uuid::Uuid], reserved: &HashSet, @@ -846,198 +812,3 @@ fn allocate_public_ipv6_leases( leases } - -#[cfg(test)] -mod tests { - use std::net::Ipv6Addr; - use std::{ - collections::{HashMap, HashSet}, - sync::{Arc, Mutex}, - }; - - use cidr::{Ipv6Cidr, Ipv6Inet}; - - use crate::{ - common::{PeerId, global_ctx::tests::get_mock_global_ctx}, - peers::peer_rpc::PeerRpcManager, - }; - - use super::{ - PublicIpv6PeerRouteInfo, PublicIpv6RouteControl, PublicIpv6Service, PublicIpv6SyncTrigger, - allocate_public_ipv6_leases, - }; - - struct TestRouteControl { - my_peer_id: PeerId, - peers: Mutex>, - } - - impl PublicIpv6RouteControl for TestRouteControl { - fn my_peer_id(&self) -> PeerId { - self.my_peer_id - } - - fn peer_route_snapshot(&self) -> Vec { - self.peers.lock().unwrap().clone() - } - - fn publish_self_public_ipv6_lease(&self, _lease: Option) -> bool { - false - } - } - - struct TestSyncTrigger; - - impl PublicIpv6SyncTrigger for TestSyncTrigger { - fn sync_now(&self, _reason: &str) {} - } - - #[test] - fn public_ipv6_lease_allocator_keeps_stable_addresses() { - let prefix = "2001:db8::/124".parse::().unwrap(); - let first = uuid::Uuid::from_u128(1); - let second = uuid::Uuid::from_u128(2); - - let leases = - allocate_public_ipv6_leases(prefix, &[first, second], &HashSet::new(), &HashMap::new()); - assert_eq!(leases.len(), 2); - assert_ne!(leases[0].addr, leases[1].addr); - - let initial_map = HashMap::from([(first, leases[0].addr)]); - let next = allocate_public_ipv6_leases(prefix, &[first], &HashSet::new(), &initial_map); - assert_eq!(next.len(), 1); - assert_eq!(next[0].addr, leases[0].addr); - assert!(next[0].reused); - } - - #[test] - fn public_ipv6_provider_prefers_smallest_instance_id() { - let info_a = PublicIpv6PeerRouteInfo { - peer_id: 2, - inst_id: Some(uuid::Uuid::from_u128(2)), - is_provider: true, - prefix: Some("2001:db8:1::/120".parse().unwrap()), - lease: None, - reachable: true, - }; - let info_b = PublicIpv6PeerRouteInfo { - peer_id: 1, - inst_id: Some(uuid::Uuid::from_u128(1)), - is_provider: true, - prefix: Some("2001:db8:2::/120".parse().unwrap()), - lease: None, - reachable: true, - }; - - let selected = - PublicIpv6Service::selected_provider_from_snapshot(&[info_a, info_b]).unwrap(); - assert_eq!(selected.peer_id, 1); - } - - #[test] - fn public_ipv6_provider_prefers_reachable_provider() { - let unreachable_lower_id = PublicIpv6PeerRouteInfo { - peer_id: 1, - inst_id: Some(uuid::Uuid::from_u128(1)), - is_provider: true, - prefix: Some("2001:db8:1::/120".parse().unwrap()), - lease: None, - reachable: false, - }; - let reachable_higher_id = PublicIpv6PeerRouteInfo { - peer_id: 2, - inst_id: Some(uuid::Uuid::from_u128(2)), - is_provider: true, - prefix: Some("2001:db8:2::/120".parse().unwrap()), - lease: None, - reachable: true, - }; - - let selected = PublicIpv6Service::selected_provider_from_snapshot(&[ - unreachable_lower_id, - reachable_higher_id, - ]) - .unwrap(); - assert_eq!(selected.peer_id, 2); - } - - #[test] - fn public_ipv6_lease_allocator_stops_when_only_network_offset_is_left() { - let prefix = "2001:db8::/126".parse::().unwrap(); - let network = prefix.first_address(); - let reserved = HashSet::from([ - Ipv6Addr::from(u128::from(network) + 1), - Ipv6Addr::from(u128::from(network) + 2), - Ipv6Addr::from(u128::from(network) + 3), - ]); - - let leases = allocate_public_ipv6_leases( - prefix, - &[uuid::Uuid::from_u128(42)], - &reserved, - &HashMap::new(), - ); - - assert!(leases.is_empty()); - } - - #[tokio::test] - async fn reconcile_runtime_clears_public_ipv6_lease_when_auto_is_disabled() { - let global_ctx = get_mock_global_ctx(); - global_ctx.config.set_ipv6_public_addr_auto(false); - - let virtual_addr = "fd00::1/64".parse().unwrap(); - let stale_addr = "2001:db8::123/64".parse().unwrap(); - global_ctx.set_ipv6(Some(virtual_addr)); - global_ctx.set_public_ipv6_lease(Some(stale_addr)); - - let service = Arc::new(PublicIpv6Service::new( - global_ctx.clone(), - std::sync::Weak::::new(), - Arc::new(TestRouteControl { - my_peer_id: 1, - peers: Mutex::new(Vec::new()), - }), - Arc::new(TestSyncTrigger), - )); - *service.my_addr_cache.lock().unwrap() = Some(stale_addr); - - service.reconcile_runtime_from_snapshot(&[]); - - assert_eq!(*service.my_addr_cache.lock().unwrap(), None); - assert_eq!(global_ctx.get_ipv6(), Some(virtual_addr)); - assert_eq!(global_ctx.get_public_ipv6_lease(), None); - } - - #[tokio::test] - async fn reconcile_runtime_keeps_virtual_ipv6_when_public_lease_changes() { - let global_ctx = get_mock_global_ctx(); - global_ctx.config.set_ipv6_public_addr_auto(true); - - let virtual_addr = "fd00::1/64".parse().unwrap(); - let public_addr = "2001:db8::123/64".parse().unwrap(); - global_ctx.set_ipv6(Some(virtual_addr)); - - let service = Arc::new(PublicIpv6Service::new( - global_ctx.clone(), - std::sync::Weak::::new(), - Arc::new(TestRouteControl { - my_peer_id: 1, - peers: Mutex::new(vec![PublicIpv6PeerRouteInfo { - peer_id: 1, - inst_id: Some(uuid::Uuid::from_u128(1)), - is_provider: false, - prefix: None, - lease: Some(public_addr), - reachable: true, - }]), - }), - Arc::new(TestSyncTrigger), - )); - - service.reconcile_runtime(); - - assert_eq!(global_ctx.get_ipv6(), Some(virtual_addr)); - assert_eq!(global_ctx.get_public_ipv6_lease(), Some(public_addr)); - } -} diff --git a/easytier/src/peers/relay_peer_map.rs b/easytier-core/src/peers/relay_peer_map.rs similarity index 87% rename from easytier/src/peers/relay_peer_map.rs rename to easytier-core/src/peers/relay_peer_map.rs index ec5562c3..763141cf 100644 --- a/easytier/src/peers/relay_peer_map.rs +++ b/easytier-core/src/peers/relay_peer_map.rs @@ -5,18 +5,24 @@ use prost::Message; use quanta::Instant; use snow::params::NoiseParams; use tokio::sync::{Mutex, OwnedMutexGuard, oneshot}; -use tokio::time::{Duration, timeout}; -use crate::peers::foreign_network_client::ForeignNetworkClient; use crate::{ - common::error::Error, - common::{PeerId, global_ctx::ArcGlobalCtx, shrink_dashmap}, - peers::peer_map::PeerMap, - peers::peer_session::{PeerSession, PeerSessionAction, PeerSessionStore, SessionKey}, - peers::route_trait::NextHopPolicy, - peers::traffic_metrics::AggregateTrafficMetrics, + config::PeerId, + foundation::time::{Duration, timeout}, + packet::{PacketType, ZCPacket}, + peers::{ + conn::{ + peer_map::PeerMap, + peer_session::{PeerSession, PeerSessionAction, PeerSessionStore, SessionKey}, + }, + context::ArcPeerContext, + error::Error, + foreign_network::client::ForeignNetworkClient, + route::NextHopPolicy, + util::shrink_dashmap, + }, + proto::peer_rpc::RoutePeerInfo, proto::peer_rpc::{PeerConnSessionActionPb, RelayNoiseMsg1Pb, RelayNoiseMsg2Pb}, - tunnel::packet_def::{PacketType, ZCPacket}, }; const RELAY_NOISE_VERSION: u32 = 1; @@ -44,9 +50,9 @@ impl Default for RelayPeerState { } pub struct RelayPeerMap { - peer_map: Arc, - foreign_network_client: Option>, - global_ctx: ArcGlobalCtx, + route_transport: Arc, + context: ArcPeerContext, + metric_network_name: String, my_peer_id: PeerId, peer_session_store: Arc, states: DashMap, @@ -55,30 +61,89 @@ pub struct RelayPeerMap { pub(crate) pending_packets: DashMap>, is_secure_mode_enabled: bool, - control_metrics: AggregateTrafficMetrics, +} + +#[async_trait::async_trait] +#[auto_impl::auto_impl(Arc)] +pub trait RelayRouteTransport: Send + Sync { + async fn get_route_peer_info(&self, peer_id: PeerId) -> Option; + + async fn send_msg_to_next_hop( + &self, + msg: ZCPacket, + dst_peer_id: PeerId, + policy: NextHopPolicy, + ) -> Result<(), Error>; +} + +pub struct PeerMapRelayRouteTransport { + peer_map: Arc, + foreign_network_client: Option>, +} + +#[async_trait::async_trait] +impl RelayRouteTransport for PeerMapRelayRouteTransport { + async fn get_route_peer_info(&self, peer_id: PeerId) -> Option { + self.peer_map.get_route_peer_info(peer_id).await + } + + async fn send_msg_to_next_hop( + &self, + msg: ZCPacket, + dst_peer_id: PeerId, + policy: NextHopPolicy, + ) -> Result<(), Error> { + let Some(next_hop) = self.peer_map.get_gateway_peer_id(dst_peer_id, policy).await else { + return Err(Error::RouteError(Some(format!( + "next hop not found in route for peer {dst_peer_id:?}" + )))); + }; + if self.peer_map.has_peer(next_hop) { + self.peer_map.send_msg_directly(msg, next_hop).await + } else if let Some(foreign_network_client) = &self.foreign_network_client { + foreign_network_client.send_msg(msg, next_hop).await + } else { + Err(Error::RouteError(Some(format!( + "next hop not found in direct peer map: {next_hop:?}" + )))) + } + } +} + +pub fn new_relay_peer_map( + peer_map: Arc, + foreign_network_client: Option>, + context: ArcPeerContext, + my_peer_id: PeerId, + peer_session_store: Arc, +) -> Arc { + RelayPeerMap::new( + Arc::new(PeerMapRelayRouteTransport { + peer_map, + foreign_network_client, + }), + context, + my_peer_id, + peer_session_store, + ) } impl RelayPeerMap { - pub fn new( - peer_map: Arc, - foreign_network_client: Option>, - global_ctx: ArcGlobalCtx, + pub(crate) fn new( + route_transport: Arc, + context: ArcPeerContext, my_peer_id: PeerId, peer_session_store: Arc, ) -> Arc { - let is_secure_mode_enabled = global_ctx - .config - .get_secure_mode() + let is_secure_mode_enabled = context + .secure_mode() .map(|cfg| cfg.enabled) .unwrap_or(false); + let metric_network_name = context.network_name(); Arc::new(Self { - control_metrics: AggregateTrafficMetrics::control( - global_ctx.stats_manager().clone(), - global_ctx.get_network_name(), - ), - peer_map, - foreign_network_client, - global_ctx, + route_transport, + context, + metric_network_name, my_peer_id, peer_session_store, states: DashMap::new(), @@ -95,9 +160,8 @@ impl RelayPeerMap { fn get_local_keypair(&self) -> Result<(Vec, Vec), Error> { let cfg = self - .global_ctx - .config - .get_secure_mode() + .context + .secure_mode() .ok_or_else(|| Error::RouteError(Some("secure mode config not set".to_string())))?; let private = cfg .private_key() @@ -110,7 +174,7 @@ impl RelayPeerMap { async fn get_remote_static_pubkey(&self, peer_id: PeerId) -> Result, Error> { let info = self - .peer_map + .route_transport .get_route_peer_info(peer_id) .await .ok_or_else(|| Error::RouteError(Some("route peer info not found".to_string())))?; @@ -140,7 +204,8 @@ impl RelayPeerMap { pkt.fill_peer_manager_hdr(self.my_peer_id, dst_peer_id, packet_type as u8); let pkt_len = pkt.buf_len() as u64; self.send_via_next_hop(pkt, dst_peer_id, policy).await?; - self.control_metrics.record_tx(pkt_len); + self.context + .record_control_tx(&self.metric_network_name, pkt_len); Ok(()) } @@ -150,20 +215,9 @@ impl RelayPeerMap { dst_peer_id: PeerId, policy: NextHopPolicy, ) -> Result<(), Error> { - let Some(next_hop) = self.peer_map.get_gateway_peer_id(dst_peer_id, policy).await else { - return Err(Error::RouteError(Some(format!( - "next hop not found in route for peer {dst_peer_id:?}" - )))); - }; - if self.peer_map.has_peer(next_hop) { - self.peer_map.send_msg_directly(msg, next_hop).await - } else if let Some(foreign_network_client) = &self.foreign_network_client { - foreign_network_client.send_msg(msg, next_hop).await - } else { - Err(Error::RouteError(Some(format!( - "next hop not found in direct peer map: {next_hop:?}" - )))) - } + self.route_transport + .send_msg_to_next_hop(msg, dst_peer_id, policy) + .await } pub async fn send_msg( @@ -230,7 +284,7 @@ impl RelayPeerMap { pub fn has_session(&self, dst_peer_id: PeerId) -> bool { self.peer_session_store .get(&SessionKey::new( - self.global_ctx.get_network_identity().network_name.clone(), + self.context.network_identity().network_name, dst_peer_id, )) .is_some() @@ -241,7 +295,7 @@ impl RelayPeerMap { dst_peer_id: PeerId, policy: NextHopPolicy, ) -> Result, Error> { - let network = self.global_ctx.get_network_identity(); + let network = self.context.network_identity(); let key = SessionKey::new(network.network_name.clone(), dst_peer_id); if let Some(session) = self.peer_session_store.get(&key) { return Ok(session); @@ -268,7 +322,7 @@ impl RelayPeerMap { policy: NextHopPolicy, _lock_guard: Option>, ) -> Result<(), Error> { - let network = self.global_ctx.get_network_identity(); + let network = self.context.network_identity(); let key = SessionKey::new(network.network_name.clone(), dst_peer_id); if let Some(session) = self.peer_session_store.get(&key) { self.flush_pending_packets(dst_peer_id, session).await; @@ -300,7 +354,7 @@ impl RelayPeerMap { self.register_handshake_failure(dst_peer_id, attempt); if attempt + 1 < HANDSHAKE_MAX_ATTEMPTS { let backoff = HANDSHAKE_RETRY_BASE_MS.saturating_mul(1 << attempt); - tokio::time::sleep(Duration::from_millis(backoff)).await; + crate::foundation::time::sleep(Duration::from_millis(backoff)).await; } } } @@ -319,7 +373,7 @@ impl RelayPeerMap { dst_peer_id: PeerId, policy: NextHopPolicy, ) -> Result, Error> { - let network = self.global_ctx.get_network_identity(); + let network = self.context.network_identity(); let session_key = SessionKey::new(network.network_name.clone(), dst_peer_id); let (local_private_key, _local_public_key) = self.get_local_keypair()?; let remote_static = self.get_remote_static_pubkey(dst_peer_id).await?; @@ -347,7 +401,7 @@ impl RelayPeerMap { version: RELAY_NOISE_VERSION, a_session_generation, a_conn_id: Some(a_conn_id.into()), - client_encryption_algorithm: self.global_ctx.get_flags().encryption_algorithm.clone(), + client_encryption_algorithm: self.context.flags().encryption_algorithm, }; let payload = msg1_pb.encode_to_vec(); let mut out = vec![0u8; 4096]; @@ -420,7 +474,7 @@ impl RelayPeerMap { key_bytes.copy_from_slice(v); key_bytes }); - let algo = self.global_ctx.get_flags().encryption_algorithm.clone(); + let algo = self.context.flags().encryption_algorithm; let session = self .peer_session_store .apply_initiator_action( @@ -477,7 +531,8 @@ impl RelayPeerMap { .peer_manager_header() .ok_or_else(|| Error::RouteError(Some("packet without header".to_string())))?; let src_peer_id = hdr.from_peer_id.get(); - self.control_metrics.record_rx(packet.buf_len() as u64); + self.context + .record_control_rx(&self.metric_network_name, packet.buf_len() as u64); match hdr.packet_type { x if x == PacketType::RelayHandshake as u8 => { tracing::debug!("handle_relay_msg1 from {:?}", src_peer_id); @@ -564,8 +619,8 @@ impl RelayPeerMap { )))); } - let server_network_name = self.global_ctx.get_network_name(); - let algo = self.global_ctx.get_flags().encryption_algorithm.clone(); + let server_network_name = self.context.network_name(); + let algo = self.context.flags().encryption_algorithm; let key = SessionKey::new(server_network_name.clone(), remote_peer_id); let upsert = self .peer_session_store @@ -621,7 +676,7 @@ impl RelayPeerMap { .peer_manager_header() .ok_or_else(|| Error::RouteError(Some("packet without header".to_string())))?; let from_peer_id = hdr.from_peer_id.get(); - let network = self.global_ctx.get_network_identity(); + let network = self.context.network_identity(); let key = SessionKey::new(network.network_name.clone(), from_peer_id); let Some(session) = self.peer_session_store.get(&key) else { tracing::debug!( @@ -663,17 +718,6 @@ impl RelayPeerMap { self.states.contains_key(&peer_id) } - pub fn failure_count(&self, peer_id: PeerId) -> Option { - self.states.get(&peer_id).map(|v| v.failure_count) - } - - pub fn is_backoff_active(&self, peer_id: PeerId) -> bool { - self.states - .get(&peer_id) - .and_then(|v| v.next_retry_at) - .is_some_and(|ts| Instant::now() < ts) - } - /// Remove relay-specific state for a specific peer. /// This does NOT remove the session from PeerSessionStore, because the /// session lifecycle is independent of any particular connection type @@ -692,3 +736,18 @@ impl RelayPeerMap { tracing::debug!(?peer_id, "RelayPeerMap removed peer relay state"); } } + +#[cfg(any(test, feature = "test-utils"))] +mod test_utils { + use super::*; + + impl RelayPeerMap { + #[doc(hidden)] + pub(crate) fn has_session_without_touch(&self, dst_peer_id: PeerId) -> bool { + self.peer_session_store.contains_valid(&SessionKey::new( + self.context.network_identity().network_name, + dst_peer_id, + )) + } + } +} diff --git a/easytier/src/peers/graph_algo.rs b/easytier-core/src/peers/route/graph_algo.rs similarity index 100% rename from easytier/src/peers/graph_algo.rs rename to easytier-core/src/peers/route/graph_algo.rs diff --git a/easytier/src/peers/route_trait.rs b/easytier-core/src/peers/route/mod.rs similarity index 90% rename from easytier/src/peers/route_trait.rs rename to easytier-core/src/peers/route/mod.rs index ca368644..e182195b 100644 --- a/easytier/src/peers/route_trait.rs +++ b/easytier-core/src/peers/route/mod.rs @@ -1,3 +1,10 @@ +//! Route trait surface shared by the peer domain, plus the OSPF route +//! implementation and the graph algorithms backing it. + +pub(crate) mod graph_algo; +pub(crate) mod peer_ospf_route; +mod route_peer_wire; + use cidr::Ipv6Inet; use cidr::{Ipv4Cidr, Ipv6Cidr}; use dashmap::DashMap; @@ -9,9 +16,10 @@ use std::{ }; use crate::{ - common::{PeerId, global_ctx::NetworkIdentity}, + config::PeerId, + peers::context::NetworkIdentity, proto::{ - api::instance::ListPublicIpv6InfoResponse, + core_peer::peer::{ListPublicIpv6InfoResponse, Route as CoreRoute}, peer_rpc::{ ForeignNetworkRouteInfoEntry, ForeignNetworkRouteInfoKey, PeerIdentityType, RouteForeignNetworkInfos, RouteForeignNetworkSummary, RoutePeerInfo, @@ -90,7 +98,7 @@ pub trait Route { self.get_next_hop(peer_id).await } - async fn list_routes(&self) -> Vec; + async fn list_routes(&self) -> Vec; // TODO: rewrite route management, remove this async fn list_proxy_cidrs(&self) -> BTreeSet; @@ -171,24 +179,17 @@ pub trait Route { } } - async fn get_peer_groups_by_ipv4(&self, ipv4: &Ipv4Addr) -> Arc> { - match self.get_peer_id_by_ipv4(ipv4).await { - Some(peer_id) => self.get_peer_groups(peer_id), - None => Arc::new(Vec::new()), - } - } - async fn dump(&self) -> String { "this route implementation does not support dump".to_string() } } -pub type ArcRoute = Arc>; +pub type ArcRoute = Arc; -pub struct MockRoute {} +pub(crate) struct DisabledRoute; #[async_trait::async_trait] -impl Route for MockRoute { +impl Route for DisabledRoute { async fn open(&self, _interface: RouteInterfaceBox) -> Result { panic!("mock route") } @@ -201,7 +202,7 @@ impl Route for MockRoute { panic!("mock route") } - async fn list_routes(&self) -> Vec { + async fn list_routes(&self) -> Vec { panic!("mock route") } diff --git a/easytier/src/peers/peer_ospf_route.rs b/easytier-core/src/peers/route/peer_ospf_route.rs similarity index 53% rename from easytier/src/peers/peer_ospf_route.rs rename to easytier-core/src/peers/route/peer_ospf_route.rs index 8afa660a..a60fa205 100644 --- a/easytier/src/peers/peer_ospf_route.rs +++ b/easytier-core/src/peers/route/peer_ospf_route.rs @@ -10,6 +10,7 @@ use std::{ }; use arc_swap::ArcSwap; +use atomic_shim::AtomicU64; use cidr::{IpCidr, Ipv4Cidr, Ipv6Cidr, Ipv6Inet}; use crossbeam::atomic::AtomicCell; use dashmap::DashMap; @@ -23,27 +24,35 @@ use petgraph::{ }; use prefix_trie::PrefixMap; use prost::Message; -use prost_reflect::{DynamicMessage, ReflectMessage}; use quanta::Instant; use tokio::{ select, - sync::Mutex, + sync::{Mutex, RwLock as AsyncRwLock}, task::{JoinHandle, JoinSet}, }; use crate::{ - common::{ - PeerId, - config::NetworkIdentity, - constants::EASYTIER_VERSION, - global_ctx::{ArcGlobalCtx, GlobalCtxEvent}, - shrink_dashmap, - stun::StunInfoCollectorTrait, + config::PeerId, + config::peers::PeerGroupIdentity, + peers::{ + PeerPacketFilter, + context::{ + ArcPeerContext, NetworkIdentity as CoreNetworkIdentity, PeerContext, PeerContextEvent, + TrustedKeyMetadata, TrustedKeySource, + }, + peer_rpc::PeerRpcManager, + public_ipv6::{ + PublicIpv6PeerRouteInfo, PublicIpv6RouteControl, PublicIpv6Runtime, PublicIpv6Service, + PublicIpv6SyncTrigger, + }, + util::shrink_dashmap, }, - peers::route_trait::{Route, RouteInterfaceBox}, proto::{ - acl::GroupIdentity, - common::{Ipv4Inet, NatType, StunInfo}, + common::{NatType, RuntimeTimestamp as Timestamp, TimestampExt}, + core_peer::peer::{ + ListPublicIpv6InfoResponse as CoreListPublicIpv6InfoResponse, + PublicIpv6LeaseInfo as CorePublicIpv6LeaseInfo, Route as CoreRouteInfo, + }, peer_rpc::{ ForeignNetworkRouteInfoEntry, ForeignNetworkRouteInfoKey, OspfRouteRpc, OspfRouteRpcClientFactory, OspfRouteRpcServer, PeerGroupInfo, PeerIdVersion, @@ -58,44 +67,26 @@ use crate::{ controller::{BaseController, Controller}, }, }, - use_global_var, }; use super::{ - PeerPacketFilter, + DefaultRouteCostCalculator, ForeignNetworkRouteInfoMap, NextHopPolicy, Route, + RouteCostCalculator, RouteCostCalculatorInterface, RouteInterfaceBox, graph_algo::dijkstra_with_first_hop, - peer_rpc::PeerRpcManager, - public_ipv6::{ - PublicIpv6PeerRouteInfo, PublicIpv6RouteControl, PublicIpv6Service, PublicIpv6SyncTrigger, - }, - route_trait::{ - DefaultRouteCostCalculator, ForeignNetworkRouteInfoMap, NextHopPolicy, RouteCostCalculator, - RouteCostCalculatorInterface, + route_peer_wire::{ + self, RawRoutePeerInfo, extract_route_peer_infos, patch_credential_route_peer_info, + raw_credential_bytes, raw_route_peer_info, }, }; -use crate::proto::common::TimestampExt; -use atomic_shim::AtomicU64; -use prost_wkt_types::Timestamp; +pub type Version = u32; -static SERVICE_ID: u32 = 7; -static UPDATE_PEER_INFO_PERIOD: Duration = Duration::from_secs(3600); -static REMOVE_DEAD_PEER_INFO_AFTER: Duration = Duration::from_secs(3660); // the cost (latency between two peers) is i32, i32::MAX is large enough. -static AVOID_RELAY_COST: usize = i32::MAX as usize; -static FORCE_USE_CONN_LIST: AtomicBool = AtomicBool::new(false); - -// if a peer is unreachable for `REMOVE_UNREACHABLE_PEER_INFO_AFTER` time, we can remove it because -// 1. all the ospf sessions between two zone are already destroy, new created session will resend the peer info. -// 2. all the dst_saved_peer_info_version in all sessions already remove the peer info, the peer info will be propagated -// in another zone when two zone restore the conneciton. -static REMOVE_UNREACHABLE_PEER_INFO_AFTER: Duration = Duration::from_secs(90); - -type Version = u32; +const AVOID_RELAY_COST: usize = i32::MAX as usize; /// Check if `child` CIDR is a subset of `parent` CIDR. /// Returns true if `child` is contained within `parent`, or if they are equal. -fn cidr_is_subset(child: &IpCidr, parent: &IpCidr) -> bool { +pub fn cidr_is_subset(child: &IpCidr, parent: &IpCidr) -> bool { match (child, parent) { (IpCidr::V4(c), IpCidr::V4(p)) => { p.first_address() <= c.first_address() && c.last_address() <= p.last_address() @@ -108,7 +99,7 @@ fn cidr_is_subset(child: &IpCidr, parent: &IpCidr) -> bool { } /// Check if `child` CIDR is a subset of `parent` CIDR (both as string representations). -fn cidr_is_subset_str(child: &str, parent: &str) -> bool { +pub fn cidr_is_subset_str(child: &str, parent: &str) -> bool { let Ok(child_cidr) = child.parse::() else { return false; }; @@ -118,57 +109,23 @@ fn cidr_is_subset_str(child: &str, parent: &str) -> bool { cidr_is_subset(&child_cidr, &parent_cidr) } -/// Patch specific fields in a raw DynamicMessage from a decoded RoutePeerInfo, -/// preserving all other fields (including unknown ones). -fn patch_raw_from_info(raw: &mut DynamicMessage, info: &RoutePeerInfo, fields: &[&str]) { - let mut decoded_raw = DynamicMessage::new(RoutePeerInfo::default().descriptor()); - decoded_raw.transcode_from(info).unwrap(); - for field_name in fields { - if let Some(value) = decoded_raw.get_field_by_name(field_name) { - raw.set_field_by_name(field_name, value.into_owned()); - } - } -} - -fn raw_credential_bytes_from_route_info( - raw_route_info: &DynamicMessage, - proof_idx: usize, -) -> Option> { - raw_route_info - .get_field_by_name("trusted_credential_pubkeys")? - .as_list()? - .get(proof_idx)? - .as_message()? - .get_field_by_name("credential")? - .as_message() - .map(|credential| credential.encode_to_vec()) -} - -fn route_peer_inst_id(info: &RoutePeerInfo) -> Option { - info.inst_id.map(Into::into) -} - #[derive(Debug, Clone)] -struct AtomicVersion(Arc); +pub struct AtomicVersion(Arc); impl AtomicVersion { - fn new() -> Self { + pub fn new() -> Self { AtomicVersion(Arc::new(AtomicU32::new(0))) } - fn get(&self) -> Version { + pub fn get(&self) -> Version { self.0.load(Ordering::Relaxed) } - fn set(&self, version: Version) { - self.0.store(version, Ordering::Relaxed); - } - - fn inc(&self) -> Version { + pub fn inc(&self) -> Version { self.0.fetch_add(1, Ordering::Relaxed) + 1 } - fn set_if_larger(&self, version: Version) -> bool { + pub fn set_if_larger(&self, version: Version) -> bool { // return true if the version is set. self.0.fetch_max(version, Ordering::Relaxed) < version } @@ -180,6 +137,599 @@ impl From for AtomicVersion { } } +#[derive(Debug, Clone)] +pub struct OspfPeerInfo { + pub peer_id: PeerId, + pub info: RoutePeerInfo, +} + +#[derive(Debug, Clone)] +pub struct OspfPeerConnInfo { + pub peer_id: PeerId, + pub connected_peers: BTreeSet, +} + +#[derive(Debug, Clone)] +pub struct OspfRouteSnapshot { + pub peer_infos: Vec, + pub conn_map: Vec, + pub suppressed_peer_ids: BTreeSet, + pub version: Version, +} + +impl OspfRouteSnapshot { + fn lookup(&self) -> OspfRouteSnapshotLookup<'_> { + OspfRouteSnapshotLookup { + peer_infos: self + .peer_infos + .iter() + .map(|entry| (entry.peer_id, &entry.info)) + .collect(), + conn_map: self + .conn_map + .iter() + .map(|entry| (entry.peer_id, entry)) + .collect(), + } + } +} + +struct OspfRouteSnapshotLookup<'a> { + peer_infos: HashMap, + conn_map: HashMap, +} + +impl OspfRouteSnapshotLookup<'_> { + fn peer_info(&self, peer_id: PeerId) -> Option<&RoutePeerInfo> { + self.peer_infos.get(&peer_id).copied() + } + + fn connected_peers>(&self, peer_id: PeerId) -> Option { + self.conn_map + .get(&peer_id) + .map(|entry| entry.connected_peers.iter().copied().collect()) + } + + fn get_avoid_relay_data(&self, peer_id: PeerId) -> bool { + // if avoid relay, just set all outgoing edges to a large value: AVOID_RELAY_COST. + self.peer_info(peer_id) + .and_then(|x| x.feature_flag) + .map(|x| x.avoid_relay_data) + .unwrap_or_default() + } +} + +type PeerGraph = Graph; +type PeerIdToNodeIdxMap = DashMap; + +#[derive(Debug, Clone, Copy)] +pub struct OspfNextHopInfo { + pub next_hop_peer_id: PeerId, + pub path_latency: i32, + pub path_len: usize, // path includes src and dst. + pub version: Version, +} + +type NextHopMap = DashMap; + +// computed with SyncedRouteInfo snapshot. used to get next hop. +#[derive(Debug)] +pub struct OspfRouteTable { + peer_infos: DashMap, + next_hop_map: NextHopMap, + suppressed_peer_ids: DashMap, + ipv4_peer_id_map: DashMap, + ipv6_peer_id_map: DashMap, + cidr_peer_id_map: ArcSwap>, + cidr_v6_peer_id_map: ArcSwap>, + next_hop_map_version: AtomicVersion, +} + +impl OspfRouteTable { + pub fn new() -> Self { + OspfRouteTable { + peer_infos: DashMap::new(), + next_hop_map: DashMap::new(), + suppressed_peer_ids: DashMap::new(), + ipv4_peer_id_map: DashMap::new(), + ipv6_peer_id_map: DashMap::new(), + cidr_peer_id_map: ArcSwap::new(Arc::new(PrefixMap::new())), + cidr_v6_peer_id_map: ArcSwap::new(Arc::new(PrefixMap::new())), + next_hop_map_version: AtomicVersion::new(), + } + } + + pub fn get_next_hop(&self, dst_peer_id: PeerId) -> Option { + if self.suppressed_peer_ids.contains_key(&dst_peer_id) { + return None; + } + self.get_topology_next_hop(dst_peer_id) + } + + pub fn get_topology_next_hop(&self, dst_peer_id: PeerId) -> Option { + let cur_version = self.next_hop_map_version.get(); + self.next_hop_map.get(&dst_peer_id).and_then(|x| { + if x.version >= cur_version { + Some(*x) + } else { + None + } + }) + } + + pub fn peer_reachable(&self, peer_id: PeerId) -> bool { + self.get_next_hop(peer_id).is_some() + } + + pub fn topology_peer_reachable(&self, peer_id: PeerId) -> bool { + self.get_topology_next_hop(peer_id).is_some() + } + + pub fn get_udp_nat_type(&self, peer_id: PeerId) -> Option { + self.peer_infos + .get(&peer_id) + .map(|x| NatType::try_from(x.udp_nat_type).unwrap_or_default()) + } + + pub fn get_peer_info(&self, peer_id: PeerId) -> Option { + self.peer_infos.get(&peer_id).map(|x| x.clone()) + } + + pub fn get_peer_id_by_ipv4(&self, ipv4_addr: &Ipv4Addr) -> Option { + self.ipv4_peer_id_map.get(ipv4_addr).map(|p| p.peer_id) + } + + pub fn get_peer_id_by_ipv6(&self, ipv6_addr: &Ipv6Addr) -> Option { + self.ipv6_peer_id_map.get(ipv6_addr).map(|p| p.peer_id) + } + + fn sync_suppressed_peer_ids(&self, snapshot: &OspfRouteSnapshot) { + self.suppressed_peer_ids + .retain(|peer_id, _| snapshot.suppressed_peer_ids.contains(peer_id)); + for peer_id in &snapshot.suppressed_peer_ids { + self.suppressed_peer_ids.insert(*peer_id, ()); + } + } + + // return graph and start node index (node of my peer id). + fn build_peer_graph_from_snapshot( + my_peer_id: PeerId, + snapshot: &OspfRouteSnapshot, + lookup: &OspfRouteSnapshotLookup<'_>, + cost_calc: &T, + ) -> (PeerGraph, NodeIndex) { + let mut graph: PeerGraph = PeerGraph::new(); + + let mut start_node_idx = None; + let peer_id_to_node_index: PeerIdToNodeIdxMap = DashMap::new(); + for entry in &snapshot.peer_infos { + let peer_id = entry.peer_id; + + if entry.info.version == 0 { + continue; + } + + let node_idx = graph.add_node(peer_id); + + peer_id_to_node_index.insert(peer_id, node_idx); + if peer_id == my_peer_id { + start_node_idx = Some(node_idx); + } + } + + if start_node_idx.is_none() { + return (graph, NodeIndex::end()); + } + + for item in peer_id_to_node_index.iter() { + let src_peer_id = *item.key(); + if src_peer_id != my_peer_id && snapshot.suppressed_peer_ids.contains(&src_peer_id) { + continue; + } + let src_node_idx = item.value(); + let connected_peers: BTreeSet<_> = + lookup.connected_peers(src_peer_id).unwrap_or_default(); + + // if avoid relay, just set all outgoing edges to a large value: AVOID_RELAY_COST. + let peer_avoid_relay_data = lookup.get_avoid_relay_data(src_peer_id); + + for dst_peer_id in connected_peers.iter() { + let Some(dst_node_idx) = peer_id_to_node_index.get(dst_peer_id) else { + continue; + }; + + let mut cost = cost_calc.calculate_cost(src_peer_id, *dst_peer_id) as usize; + if peer_avoid_relay_data { + cost += AVOID_RELAY_COST; + } + + graph.add_edge(*src_node_idx, *dst_node_idx, cost); + } + } + + (graph, start_node_idx.unwrap()) + } + + pub fn clean_expired_route_info(&self) { + let cur_version = self.next_hop_map_version.get(); + self.next_hop_map.retain(|_, v| { + // remove next hop map for peers we cannot reach. + v.version >= cur_version + }); + self.peer_infos.retain(|k, _| { + // remove peer info for peers we cannot forward to. + self.peer_reachable(*k) + }); + self.ipv4_peer_id_map.retain(|_, v| { + // remove ipv4 map for peers we cannot forward to. + self.peer_reachable(v.peer_id) + }); + self.ipv6_peer_id_map.retain(|_, v| { + // remove ipv6 map for peers we cannot forward to. + self.peer_reachable(v.peer_id) + }); + + shrink_dashmap(&self.peer_infos, None); + shrink_dashmap(&self.next_hop_map, None); + shrink_dashmap(&self.suppressed_peer_ids, None); + shrink_dashmap(&self.ipv4_peer_id_map, None); + shrink_dashmap(&self.ipv6_peer_id_map, None); + } + + fn gen_next_hop_map_with_least_hop( + &self, + graph: &PeerGraph, + start_node: &NodeIndex, + version: Version, + ) { + if graph.node_weight(*start_node).is_none() { + tracing::warn!( + ?start_node, + version, + "invalid start node for least-hop route rebuild" + ); + return; + } + let normalize_edge_cost = |e: petgraph::graph::EdgeReference| { + if *e.weight() >= AVOID_RELAY_COST { + AVOID_RELAY_COST + 1 + } else { + 1 + } + }; + // Step 1: first Dijkstra, compute shortest hop count. + let path_len_map = dijkstra(graph, *start_node, None, normalize_edge_cost); + + // Step 2: build subgraph containing only shortest-hop and AVOID RELAY edges. + let mut subgraph: PeerGraph = PeerGraph::new(); + let mut start_node_idx = None; + for (node_idx, peer_id) in graph.node_references() { + let new_node_idx = subgraph.add_node(*peer_id); + if node_idx == *start_node { + start_node_idx = Some(new_node_idx); + } + } + + for edge in graph.edge_references() { + let (src, tgt) = graph.edge_endpoints(edge.id()).unwrap(); + let Some(src_path_len) = path_len_map.get(&src) else { + continue; + }; + let Some(tgt_path_len) = path_len_map.get(&tgt) else { + continue; + }; + if *src_path_len + normalize_edge_cost(edge) == *tgt_path_len { + subgraph.add_edge(src, tgt, *edge.weight()); + } + } + + // Step 3: second Dijkstra on subgraph, choose least-cost among shortest-hop paths. + self.gen_next_hop_map_with_least_cost(&subgraph, &start_node_idx.unwrap(), version); + } + + fn gen_next_hop_map_with_least_cost( + &self, + graph: &PeerGraph, + start_node: &NodeIndex, + version: Version, + ) { + if graph.node_weight(*start_node).is_none() { + tracing::warn!( + ?start_node, + version, + "invalid start node for least-cost route rebuild" + ); + return; + } + let (costs, next_hops) = dijkstra_with_first_hop(graph, *start_node, |e| *e.weight()); + + for (dst, (next_hop, path_len)) in next_hops.iter() { + let info = OspfNextHopInfo { + next_hop_peer_id: *graph.node_weight(*next_hop).unwrap(), + path_latency: (*costs.get(dst).unwrap() % AVOID_RELAY_COST) as i32, + path_len: *path_len, + version, + }; + let dst_peer_id = *graph.node_weight(*dst).unwrap(); + self.next_hop_map + .entry(dst_peer_id) + .and_modify(|x| { + if x.version < version { + *x = info; + } + }) + .or_insert(info); + } + + self.next_hop_map_version.set_if_larger(version); + } + + pub fn build_from_snapshot( + &self, + my_peer_id: PeerId, + snapshot: &OspfRouteSnapshot, + policy: NextHopPolicy, + cost_calc: &T, + ) { + let version = snapshot.version; + self.sync_suppressed_peer_ids(snapshot); + let lookup = snapshot.lookup(); + + let local_proxy_cidrs = lookup + .peer_info(my_peer_id) + .into_iter() + .flat_map(|info| &info.proxy_cidrs) + .filter_map(|cidr| cidr.parse::().ok()) + .collect::>(); + + // build next hop map + let (graph, start_node) = + Self::build_peer_graph_from_snapshot(my_peer_id, snapshot, &lookup, cost_calc); + + if graph.node_count() == 0 { + tracing::warn!("no peer in graph, cannot build next hop map"); + self.next_hop_map_version.set_if_larger(version); + self.clean_expired_route_info(); + return; + } + if start_node == NodeIndex::end() { + tracing::warn!( + ?my_peer_id, + version, + "my peer id is missing in graph, skip next-hop rebuild this round" + ); + self.next_hop_map_version.set_if_larger(version); + self.clean_expired_route_info(); + return; + } + + if matches!(policy, NextHopPolicy::LeastHop) { + self.gen_next_hop_map_with_least_hop(&graph, &start_node, version); + } else { + self.gen_next_hop_map_with_least_cost(&graph, &start_node, version); + }; + + let mut new_cidr_prefix_trie = PrefixMap::new(); + let mut new_cidr_v6_prefix_trie = PrefixMap::new(); + + // build peer_infos, ipv4_peer_id_map, cidr_peer_id_map + // only set map for peers we can reach. + for item in self.next_hop_map.iter() { + if item.version < version { + // skip if the next hop entry is outdated. (peer is unreachable) + continue; + } + + let peer_id = item.key(); + if !self.peer_reachable(*peer_id) { + continue; + } + + let Some(info) = lookup.peer_info(*peer_id).cloned() else { + continue; + }; + + self.peer_infos.insert(*peer_id, info.clone()); + + let peer_id_and_version = PeerIdVersion { + peer_id: *peer_id, + version, + }; + + let is_new_peer_better = |old_peer: &PeerIdVersion| -> bool { + if peer_id_and_version.version > old_peer.version { + return true; + } + if peer_id_and_version.peer_id == old_peer.peer_id { + return false; + } + let old_next_hop = self.get_next_hop(old_peer.peer_id); + let new_next_hop = item.value(); + old_next_hop.is_none() || new_next_hop.path_len < old_next_hop.unwrap().path_len + }; + + if let Some(ipv4_addr) = info.ipv4_addr { + self.ipv4_peer_id_map + .entry(ipv4_addr.into()) + .and_modify(|v| { + if is_new_peer_better(v) { + *v = peer_id_and_version; + } + }) + .or_insert(peer_id_and_version); + } + + if let Some(ipv6_addr) = info.ipv6_addr.and_then(|x| x.address) { + self.ipv6_peer_id_map + .entry(ipv6_addr.into()) + .and_modify(|v| { + if is_new_peer_better(v) { + *v = peer_id_and_version; + } + }) + .or_insert(peer_id_and_version); + } + + if let Some(ipv6_addr) = info + .ipv6_public_addr_lease + .as_ref() + .and_then(|addr| addr.address) + { + self.ipv6_peer_id_map + .entry(ipv6_addr.into()) + .and_modify(|v| { + if is_new_peer_better(v) { + *v = peer_id_and_version; + } + }) + .or_insert(peer_id_and_version); + } + + for cidr in info.proxy_cidrs.iter() { + let Ok(cidr) = cidr.parse::() else { + tracing::warn!("invalid proxy cidr: {:?}, from peer: {:?}", cidr, peer_id); + continue; + }; + + if *peer_id != my_peer_id + && local_proxy_cidrs + .iter() + .any(|local_cidr| cidr_is_subset(&cidr, local_cidr)) + { + tracing::debug!( + ?peer_id, + ?my_peer_id, + ?local_proxy_cidrs, + ?cidr, + "skip remote proxy cidr covered by local announced proxy cidr while building route table" + ); + continue; + } + match cidr { + IpCidr::V4(cidr) => { + new_cidr_prefix_trie + .entry(cidr) + .and_modify(|e| { + // if ourself has same cidr, ensure here put my peer id, so we can know deadloop may happen. + if *peer_id == my_peer_id || is_new_peer_better(e) { + *e = peer_id_and_version; + } + }) + .or_insert(peer_id_and_version); + } + + IpCidr::V6(cidr) => { + new_cidr_v6_prefix_trie + .entry(cidr) + .and_modify(|e| { + // if ourself has same cidr, ensure here put my peer id, so we can know deadloop may happen. + if *peer_id == my_peer_id || is_new_peer_better(e) { + *e = peer_id_and_version; + } + }) + .or_insert(peer_id_and_version); + } + } + tracing::debug!( + "add cidr: {:?} to peer: {:?}, my peer id: {:?}", + cidr, + peer_id, + my_peer_id + ); + } + } + + self.cidr_peer_id_map.store(Arc::new(new_cidr_prefix_trie)); + self.cidr_v6_peer_id_map + .store(Arc::new(new_cidr_v6_prefix_trie)); + tracing::trace!( + my_peer_id = my_peer_id, + cidrs = ?self.cidr_peer_id_map.load(), + cidrs_v6 = ?self.cidr_v6_peer_id_map.load(), + "update peer cidr map" + ); + self.clean_expired_route_info(); + } + + pub fn get_peer_id_for_proxy(&self, ip: &IpAddr) -> Option { + match ip { + IpAddr::V4(ipv4) => self + .cidr_peer_id_map + .load() + .get_lpm(&Ipv4Cidr::new(*ipv4, 32).unwrap()) + .map(|x| x.1.peer_id), + IpAddr::V6(ipv6) => self + .cidr_v6_peer_id_map + .load() + .get_lpm(&Ipv6Cidr::new(*ipv6, 128).unwrap()) + .map(|x| x.1.peer_id), + } + } + + pub fn list_routes( + &self, + my_peer_id: PeerId, + route_table_with_cost: &Self, + ) -> Vec { + let mut routes = Vec::new(); + for item in self.peer_infos.iter() { + if *item.key() == my_peer_id { + continue; + } + let Some(next_hop_peer) = self.get_next_hop(*item.key()) else { + continue; + }; + let next_hop_peer_latency_first = route_table_with_cost.get_next_hop(*item.key()); + let mut route: CoreRouteInfo = item.value().clone().into(); + route.next_hop_peer_id = next_hop_peer.next_hop_peer_id; + route.cost = next_hop_peer.path_len as i32; + route.path_latency = next_hop_peer.path_latency; + + route.next_hop_peer_id_latency_first = + next_hop_peer_latency_first.map(|x| x.next_hop_peer_id); + route.cost_latency_first = next_hop_peer_latency_first.map(|x| x.path_len as i32); + route.path_latency_latency_first = next_hop_peer_latency_first.map(|x| x.path_latency); + + route.feature_flag = item.feature_flag; + + routes.push(route); + } + routes + } + + pub fn list_proxy_cidrs_excluding(&self, peer_id: PeerId) -> BTreeSet { + self.cidr_peer_id_map + .load() + .iter() + .filter(|(_, pv)| pv.peer_id != peer_id) + .map(|(cidr, _)| *cidr) + .collect() + } + + pub fn list_proxy_cidrs_v6_excluding(&self, peer_id: PeerId) -> BTreeSet { + self.cidr_v6_peer_id_map + .load() + .iter() + .filter(|(_, pv)| pv.peer_id != peer_id) + .map(|(cidr, _)| *cidr) + .collect() + } +} + +static UPDATE_PEER_INFO_PERIOD: Duration = Duration::from_secs(3600); +static REMOVE_DEAD_PEER_INFO_AFTER: Duration = Duration::from_secs(3660); +static FORCE_USE_CONN_LIST: AtomicBool = AtomicBool::new(false); + +// if a peer is unreachable for `REMOVE_UNREACHABLE_PEER_INFO_AFTER` time, we can remove it because +// 1. all the ospf sessions between two zone are already destroy, new created session will resend the peer info. +// 2. all the dst_saved_peer_info_version in all sessions already remove the peer info, the peer info will be propagated +// in another zone when two zone restore the conneciton. +static REMOVE_UNREACHABLE_PEER_INFO_AFTER: Duration = Duration::from_secs(90); + +fn route_peer_inst_id(info: &RoutePeerInfo) -> Option { + info.inst_id.map(Into::into) +} + fn is_foreign_network_info_newer( next: &ForeignNetworkRouteInfoEntry, prev: &ForeignNetworkRouteInfoEntry, @@ -190,190 +740,100 @@ fn is_foreign_network_info_newer( ) } -impl RoutePeerInfo { - #[allow(deprecated)] - pub fn new() -> Self { - Self { - peer_id: 0, - inst_id: Some(uuid::Uuid::nil().into()), - cost: 0, - ipv4_addr: None, - proxy_cidrs: Vec::new(), - hostname: None, - udp_nat_type: 0, - tcp_nat_type: 0, - // ensure this is updated when the peer_infos/conn_info/foreign_network lock is acquired. - // else we may assign a older timestamp than iterate time. - last_update: None, - version: 0, - easytier_version: EASYTIER_VERSION.to_string(), - feature_flag: None, - peer_route_id: 0, - network_length: 24, - ipv6_addr: None, - groups: Vec::new(), +#[allow(deprecated)] +fn new_route_peer_info_with_version(easytier_version: String) -> RoutePeerInfo { + RoutePeerInfo { + peer_id: 0, + inst_id: Some(uuid::Uuid::nil().into()), + cost: 0, + ipv4_addr: None, + proxy_cidrs: Vec::new(), + hostname: None, + udp_nat_type: 0, + tcp_nat_type: 0, + // ensure this is updated when the peer_infos/conn_info/foreign_network lock is acquired. + // else we may assign a older timestamp than iterate time. + last_update: None, + version: 0, + easytier_version, + feature_flag: None, + peer_route_id: 0, + network_length: 24, + ipv6_addr: None, + groups: Vec::new(), - quic_port: None, - noise_static_pubkey: Vec::new(), - trusted_credential_pubkeys: Vec::new(), - ipv6_public_addr_prefix: None, - ipv6_public_addr_lease: None, - } - } - - /// Creates a new `RoutePeerInfo` instance with updated information from the given context. - /// - /// # Parameters - /// - `my_peer_id`: The unique identifier for the peer. - /// - `peer_route_id`: The route identifier associated with the peer. - /// - `global_ctx`: Reference to the global context containing configuration and state. - /// - /// # Returns - /// A new `RoutePeerInfo` instance initialized with values from the provided context and parameters. - pub fn new_updated_self( - my_peer_id: PeerId, - peer_route_id: u64, - global_ctx: &ArcGlobalCtx, - public_ipv6_addr_lease: Option, - ) -> Self { - let stun_info = global_ctx.get_stun_info_collector().get_stun_info(); - let noise_static_pubkey = global_ctx - .config - .get_secure_mode() - .and_then(|cfg| cfg.public_key().ok()) - .map(|pk| pk.as_bytes().to_vec()) - .unwrap_or_default(); - Self { - peer_id: my_peer_id, - inst_id: Some(global_ctx.get_id().into()), - cost: 0, - ipv4_addr: global_ctx.get_ipv4().map(|x| x.address().into()), - proxy_cidrs: global_ctx - .config - .get_proxy_cidrs() - .iter() - .map(|x| x.mapped_cidr.unwrap_or(x.cidr)) - .chain(global_ctx.get_vpn_portal_cidr()) - .map(|x| x.to_string()) - .collect(), - hostname: Some(global_ctx.get_hostname()), - udp_nat_type: stun_info.udp_nat_type, - tcp_nat_type: stun_info.tcp_nat_type, - - // these two fields should not participate in comparison. - last_update: None, - version: 0, - - easytier_version: EASYTIER_VERSION.to_string(), - feature_flag: Some(global_ctx.get_feature_flags()), - peer_route_id, - network_length: global_ctx - .get_ipv4() - .map(|x| x.network_length() as u32) - .unwrap_or(24), - - ipv6_addr: global_ctx.get_ipv6().map(|x| x.into()), - ipv6_public_addr_prefix: global_ctx.get_advertised_ipv6_public_addr_prefix().map( - |prefix| { - Ipv6Inet::new(prefix.first_address(), prefix.network_length()) - .unwrap() - .into() - }, - ), - ipv6_public_addr_lease: public_ipv6_addr_lease.map(Into::into), - - groups: global_ctx.get_acl_groups(my_peer_id), - - noise_static_pubkey, - - // Only admin nodes (holding network_secret) publish trusted credential pubkeys - trusted_credential_pubkeys: if let Some(network_secret) = - global_ctx.get_network_identity().network_secret - { - global_ctx - .get_credential_manager() - .get_trusted_pubkeys(&network_secret) - } else { - Vec::new() - }, - - ..Default::default() - } - } - - /// Attempts to update the `new` RoutePeerInfo based on the `old` RoutePeerInfo. - /// - /// An update is triggered if any fields in `new` differ from `old`, or if the time since - /// `old.last_update` exceeds the `UPDATE_PEER_INFO_PERIOD`. - /// - /// If an update occurs, `new.last_update` is set to the current time and `new.version` is incremented. - /// Otherwise, `new.last_update` and `new.version` are copied from `old` without modification. - /// - /// Returns `true` if an update was performed (fields changed or periodic update required), - /// or `false` if no update was necessary. - pub fn try_update_new_peer_info(old: &RoutePeerInfo, new: &mut RoutePeerInfo) -> bool { - let need_update_periodically = if let Ok(Ok(d)) = - SystemTime::try_from(old.last_update.unwrap_or_default()).map(|x| x.elapsed()) - { - d > UPDATE_PEER_INFO_PERIOD - } else { - true - }; - - // these two fields should not participate in comparison. - new.version = old.version; - new.last_update = old.last_update; - - if *new != *old || need_update_periodically { - new.version += 1; - true - } else { - false - } + quic_port: None, + noise_static_pubkey: Vec::new(), + trusted_credential_pubkeys: Vec::new(), + ipv6_public_addr_prefix: None, + ipv6_public_addr_lease: None, } } -impl From for crate::proto::api::instance::Route { - fn from(val: RoutePeerInfo) -> Self { - let network_length = if val.network_length == 0 { - 24 +/// Creates a new `RoutePeerInfo` instance with updated information from the given context. +pub fn new_updated_self_route_peer_info( + my_peer_id: PeerId, + peer_route_id: u64, + context: &dyn PeerContext, + public_ipv6_addr_lease: Option, +) -> RoutePeerInfo { + let stun_info = context.stun_info(); + let network_identity = context.network_identity(); + let ipv4 = context.ipv4(); + let noise_static_pubkey = context + .secure_mode() + .and_then(|cfg| cfg.public_key().ok()) + .map(|pk| pk.as_bytes().to_vec()) + .unwrap_or_default(); + RoutePeerInfo { + peer_id: my_peer_id, + inst_id: Some(context.instance_id().into()), + cost: 0, + ipv4_addr: ipv4.as_ref().map(|x| x.address().into()), + proxy_cidrs: context + .proxy_cidrs() + .into_iter() + .chain(context.vpn_portal_cidr()) + .map(|x| x.to_string()) + .collect(), + hostname: Some(context.hostname()), + udp_nat_type: stun_info.udp_nat_type, + tcp_nat_type: stun_info.tcp_nat_type, + + // these two fields should not participate in comparison. + last_update: None, + version: 0, + + easytier_version: context.easytier_version(), + feature_flag: Some(context.feature_flags()), + peer_route_id, + network_length: ipv4 + .as_ref() + .map(|x| x.network_length() as u32) + .unwrap_or(24), + + ipv6_addr: context.ipv6().map(|x| x.into()), + ipv6_public_addr_prefix: context.advertised_ipv6_public_addr_prefix().map(|prefix| { + Ipv6Inet::new(prefix.first_address(), prefix.network_length()) + .unwrap() + .into() + }), + ipv6_public_addr_lease: public_ipv6_addr_lease.map(Into::into), + + groups: context.peer_groups(my_peer_id), + + noise_static_pubkey, + + // Only admin nodes (holding network_secret) publish trusted credential pubkeys + trusted_credential_pubkeys: if let Some(network_secret) = + network_identity.network_secret.as_deref() + { + context.trusted_credential_pubkeys(network_secret) } else { - val.network_length - }; + Vec::new() + }, - crate::proto::api::instance::Route { - peer_id: val.peer_id, - ipv4_addr: val.ipv4_addr.map(|ipv4_addr| Ipv4Inet { - address: Some(ipv4_addr), - network_length, - }), - next_hop_peer_id: 0, // next_hop_peer_id is calculated in RouteTable. - cost: 0, // cost is calculated in RouteTable. - path_latency: 0, // path_latency is calculated in RouteTable. - proxy_cidrs: val.proxy_cidrs.clone(), - hostname: val.hostname.unwrap_or_default(), - stun_info: { - let mut stun_info = StunInfo::default(); - if let Ok(udp_nat_type) = NatType::try_from(val.udp_nat_type) { - stun_info.set_udp_nat_type(udp_nat_type); - } - if let Ok(tcp_nat_type) = NatType::try_from(val.tcp_nat_type) { - stun_info.set_tcp_nat_type(tcp_nat_type); - } - Some(stun_info) - }, - inst_id: val.inst_id.map(|x| x.to_string()).unwrap_or_default(), - version: val.easytier_version, - feature_flag: val.feature_flag, - - next_hop_peer_id_latency_first: None, - cost_latency_first: None, - path_latency_latency_first: None, - - ipv6_addr: val.ipv6_addr, - public_ipv6_addr: val.ipv6_public_addr_lease, - ipv6_public_addr_prefix: val.ipv6_public_addr_prefix, - } + ..Default::default() } } @@ -381,22 +841,25 @@ type RouteConnBitmap = crate::proto::peer_rpc::RouteConnBitmap; type RouteConnPeerList = crate::proto::peer_rpc::RouteConnPeerList; type PeerConnInfo = crate::proto::peer_rpc::route_conn_peer_list::PeerConnInfo; -impl RouteConnBitmap { - fn get_bit(&self, idx: usize) -> bool { - let byte_idx = idx / 8; - let bit_idx = idx % 8; - let byte = self.bitmap[byte_idx]; - (byte >> bit_idx) & 1 == 1 - } +/// Attempts to update the `new` RoutePeerInfo based on the `old` RoutePeerInfo. +fn try_update_new_peer_info(old: &RoutePeerInfo, new: &mut RoutePeerInfo) -> bool { + let need_update_periodically = if let Ok(Ok(d)) = + SystemTime::try_from(old.last_update.unwrap_or_default()).map(|x| x.elapsed()) + { + d > UPDATE_PEER_INFO_PERIOD + } else { + true + }; - fn get_connected_peers(&self, peer_idx: usize) -> BTreeSet { - let mut connected_peers = BTreeSet::new(); - for (idx, peer_id_version) in self.peer_ids.iter().enumerate() { - if self.get_bit(peer_idx * self.peer_ids.len() + idx) { - connected_peers.insert(peer_id_version.peer_id); - } - } - connected_peers + // these two fields should not participate in comparison. + new.version = old.version; + new.last_update = old.last_update; + + if *new != *old || need_update_periodically { + new.version += 1; + true + } else { + false } } @@ -428,9 +891,10 @@ struct InterfacePeerSnapshot { // constructed with all infos synced from all peers. struct SyncedRouteInfo { + default_easytier_version: String, peer_infos: RwLock>, - // prost doesn't support unknown fields, so we use DynamicMessage to store raw infos and propagate them to other peers. - raw_peer_infos: DashMap, + // Keep the original protobuf bytes so unknown fields survive multi-hop propagation. + raw_peer_infos: DashMap, conn_map: RwLock>, foreign_network: DashMap, group_trust_map: DashMap>>, @@ -452,6 +916,7 @@ struct SyncedRouteInfo { impl Debug for SyncedRouteInfo { fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { f.debug_struct("SyncedRouteInfo") + .field("default_easytier_version", &self.default_easytier_version) .field("peer_infos", &self.peer_infos) .field("conn_map", &self.conn_map) .field("foreign_network", &self.foreign_network) @@ -461,6 +926,7 @@ impl Debug for SyncedRouteInfo { } } +#[allow(dead_code)] impl SyncedRouteInfo { fn set_peer_groups(&self, peer_id: PeerId, groups: HashMap>) { if groups.is_empty() { @@ -507,19 +973,22 @@ impl SyncedRouteInfo { fn credential_proof_is_valid( &self, - raw_route_info: Option<&DynamicMessage>, + raw_route_info: Option<&RawRoutePeerInfo>, proof_idx: usize, proof: &TrustedCredentialPubkeyProof, network_secret: Option<&str>, ) -> bool { network_secret .map(|secret| { - raw_route_info - .and_then(|raw| raw_credential_bytes_from_route_info(raw, proof_idx)) - .map(|raw_credential_bytes| { + let Some(raw_route_info) = raw_route_info else { + return proof.verify_credential_hmac(secret); + }; + match raw_credential_bytes(raw_route_info, proof_idx) { + Ok(Some(raw_credential_bytes)) => { proof.verify_credential_hmac_with_bytes(&raw_credential_bytes, secret) - }) - .unwrap_or_else(|| proof.verify_credential_hmac(secret)) + } + Ok(None) | Err(_) => false, + } }) .unwrap_or(true) } @@ -531,10 +1000,8 @@ impl SyncedRouteInfo { now: i64, ) -> ( HashMap, TrustedCredentialPubkey>, - HashMap, crate::common::global_ctx::TrustedKeyMetadata>, + HashMap, TrustedKeyMetadata>, ) { - use crate::common::global_ctx::{TrustedKeyMetadata, TrustedKeySource}; - let mut all_trusted = HashMap::new(); let mut global_trusted_keys = HashMap::new(); @@ -689,11 +1156,6 @@ impl SyncedRouteInfo { true } - fn is_route_suppressed(&self, peer_id: PeerId) -> bool { - self.suppressed_non_reusable_credential_peers - .contains_key(&peer_id) - } - fn update_credential_groups( &self, peer_infos: &OrderedHashMap, @@ -744,6 +1206,36 @@ impl SyncedRouteInfo { .map(|x| x.connected_peers.iter().copied().collect()) } + fn route_snapshot(&self) -> OspfRouteSnapshot { + let version = self.version.get(); + OspfRouteSnapshot { + peer_infos: self + .peer_infos + .read() + .iter() + .map(|(peer_id, info)| OspfPeerInfo { + peer_id: *peer_id, + info: info.clone(), + }) + .collect(), + conn_map: self + .conn_map + .read() + .iter() + .map(|(peer_id, info)| OspfPeerConnInfo { + peer_id: *peer_id, + connected_peers: info.connected_peers.clone(), + }) + .collect(), + suppressed_peer_ids: self + .suppressed_non_reusable_credential_peers + .iter() + .map(|entry| *entry.key()) + .collect(), + version, + } + } + fn remove_peer(&self, peer_id: PeerId) { self.remove_peers([peer_id]); } @@ -791,7 +1283,8 @@ impl SyncedRouteInfo { for peer_id in peer_ids { let guard = self.peer_infos.upgradable_read(); if !guard.contains_key(peer_id) { - let mut peer_info = RoutePeerInfo::new(); + let mut peer_info = + new_route_peer_info_with_version(self.default_easytier_version.clone()); let mut guard = RwLockUpgradableReadGuard::upgrade(guard); peer_info.last_update = Some(Timestamp::now()); guard.insert(*peer_id, peer_info); @@ -872,7 +1365,7 @@ impl SyncedRouteInfo { my_peer_route_id: u64, dst_peer_id: PeerId, peer_infos: &[RoutePeerInfo], - raw_peer_infos: &[DynamicMessage], + raw_peer_infos: &[RawRoutePeerInfo], ) -> Result<(), Error> { let mut need_inc_version = false; for (idx, route_info) in peer_infos.iter().enumerate() { @@ -893,13 +1386,6 @@ impl SyncedRouteInfo { &route_info, )?; - let peer_id_raw = raw_route_info - .get_field_by_name("peer_id") - .unwrap() - .as_u32() - .unwrap(); - assert_eq!(peer_id_raw, route_info.peer_id); - let mut guard = self.peer_infos.write(); // time between peers may not be synchronized, so update last_update to local now. // note only last_update with larger version will be updated to local saved peer info. @@ -1020,20 +1506,20 @@ impl SyncedRouteInfo { &self, my_peer_id: PeerId, my_peer_route_id: u64, - global_ctx: &ArcGlobalCtx, + context: &dyn PeerContext, public_ipv6_addr_lease: Option, ) -> bool { - let mut new = RoutePeerInfo::new_updated_self( + let mut new = new_updated_self_route_peer_info( my_peer_id, my_peer_route_id, - global_ctx, + context, public_ipv6_addr_lease, ); let mut guard = self.peer_infos.upgradable_read(); let old = guard.get(&my_peer_id); let new_version = old.map(|x| x.version).unwrap_or(0) + 1; let need_insert_new = if let Some(old) = old { - RoutePeerInfo::try_update_new_peer_info(old, &mut new) + try_update_new_peer_info(old, &mut new) } else { true }; @@ -1163,7 +1649,7 @@ impl SyncedRouteInfo { fn verify_and_update_group_trusts( &self, peer_infos: &[RoutePeerInfo], - local_group_declarations: &[GroupIdentity], + local_group_declarations: &[PeerGroupIdentity], trust_admin_groups_without_proof: bool, ) { let local_group_declarations = local_group_declarations @@ -1245,14 +1731,11 @@ impl SyncedRouteInfo { /// Collect trusted credential pubkeys from admin nodes (network_secret holders) /// and verify credential peers. Returns set of peer_ids that should be removed. - /// Also returns a HashMap of trusted keys for synchronization to GlobalCtx. + /// Also returns trusted-key metadata for the core trust-state update. fn verify_and_update_credential_trusts( &self, network_secret: Option<&str>, - ) -> ( - Vec, - HashMap, crate::common::global_ctx::TrustedKeyMetadata>, - ) { + ) -> (Vec, HashMap, TrustedKeyMetadata>) { self.verify_and_update_credential_trusts_with_active_peers(network_secret, |_| true) } @@ -1260,10 +1743,7 @@ impl SyncedRouteInfo { &self, network_secret: Option<&str>, is_peer_active: F, - ) -> ( - Vec, - HashMap, crate::common::global_ctx::TrustedKeyMetadata>, - ) + ) -> (Vec, HashMap, TrustedKeyMetadata>) where F: FnMut(PeerId) -> bool, { @@ -1281,11 +1761,7 @@ impl SyncedRouteInfo { network_secret: Option<&str>, is_peer_active: F, protected_peer_id: Option, - ) -> ( - Vec, - HashMap, crate::common::global_ctx::TrustedKeyMetadata>, - bool, - ) + ) -> (Vec, HashMap, TrustedKeyMetadata>, bool) where F: FnMut(PeerId) -> bool, { @@ -1355,462 +1831,6 @@ impl SyncedRouteInfo { } } -type PeerGraph = Graph; -type PeerIdToNodexIdxMap = DashMap; -#[derive(Debug, Clone, Copy)] -struct NextHopInfo { - next_hop_peer_id: PeerId, - path_latency: i32, - path_len: usize, // path includes src and dst. - version: Version, -} -// dst_peer_id -> (next_hop_peer_id, cost, path_len) -type NextHopMap = DashMap; - -// computed with SyncedRouteInfo. used to get next hop. -#[derive(Debug)] -struct RouteTable { - peer_infos: DashMap, - next_hop_map: NextHopMap, - suppressed_peer_ids: DashMap, - ipv4_peer_id_map: DashMap, - ipv6_peer_id_map: DashMap, - cidr_peer_id_map: ArcSwap>, - cidr_v6_peer_id_map: ArcSwap>, - next_hop_map_version: AtomicVersion, -} - -impl RouteTable { - fn new() -> Self { - RouteTable { - peer_infos: DashMap::new(), - next_hop_map: DashMap::new(), - suppressed_peer_ids: DashMap::new(), - ipv4_peer_id_map: DashMap::new(), - ipv6_peer_id_map: DashMap::new(), - cidr_peer_id_map: ArcSwap::new(Arc::new(PrefixMap::new())), - cidr_v6_peer_id_map: ArcSwap::new(Arc::new(PrefixMap::new())), - next_hop_map_version: AtomicVersion::new(), - } - } - - fn get_next_hop(&self, dst_peer_id: PeerId) -> Option { - if self.suppressed_peer_ids.contains_key(&dst_peer_id) { - return None; - } - self.get_topology_next_hop(dst_peer_id) - } - - fn get_topology_next_hop(&self, dst_peer_id: PeerId) -> Option { - let cur_version = self.next_hop_map_version.get(); - self.next_hop_map.get(&dst_peer_id).and_then(|x| { - if x.version >= cur_version { - Some(*x) - } else { - None - } - }) - } - - fn peer_reachable(&self, peer_id: PeerId) -> bool { - self.get_next_hop(peer_id).is_some() - } - - fn topology_peer_reachable(&self, peer_id: PeerId) -> bool { - self.get_topology_next_hop(peer_id).is_some() - } - - fn sync_suppressed_peer_ids(&self, synced_info: &SyncedRouteInfo) { - self.suppressed_peer_ids - .retain(|peer_id, _| synced_info.is_route_suppressed(*peer_id)); - for entry in synced_info.suppressed_non_reusable_credential_peers.iter() { - self.suppressed_peer_ids.insert(*entry.key(), ()); - } - } - - fn get_udp_nat_type(&self, peer_id: PeerId) -> Option { - self.peer_infos - .get(&peer_id) - .map(|x| NatType::try_from(x.udp_nat_type).unwrap_or_default()) - } - - // return graph and start node index (node of my peer id). - fn build_peer_graph_from_synced_info( - my_peer_id: PeerId, - synced_info: &SyncedRouteInfo, - cost_calc: &T, - ) -> (PeerGraph, NodeIndex) { - let mut graph: PeerGraph = PeerGraph::new(); - - let mut start_node_idx = None; - let peer_id_to_node_index: PeerIdToNodexIdxMap = DashMap::new(); - for (peer_id, info) in synced_info.peer_infos.read().iter() { - let peer_id = *peer_id; - - if info.version == 0 { - continue; - } - - let node_idx = graph.add_node(peer_id); - - peer_id_to_node_index.insert(peer_id, node_idx); - if peer_id == my_peer_id { - start_node_idx = Some(node_idx); - } - } - - if start_node_idx.is_none() { - return (graph, NodeIndex::end()); - } - - for item in peer_id_to_node_index.iter() { - let src_peer_id = *item.key(); - if src_peer_id != my_peer_id && synced_info.is_route_suppressed(src_peer_id) { - continue; - } - let src_node_idx = item.value(); - let connected_peers: BTreeSet<_> = synced_info - .get_connected_peers(src_peer_id) - .unwrap_or_default(); - - // if avoid relay, just set all outgoing edges to a large value: AVOID_RELAY_COST. - let peer_avoid_relay_data = synced_info.get_avoid_relay_data(src_peer_id); - - for dst_peer_id in connected_peers.iter() { - let Some(dst_node_idx) = peer_id_to_node_index.get(dst_peer_id) else { - continue; - }; - - let mut cost = cost_calc.calculate_cost(src_peer_id, *dst_peer_id) as usize; - if peer_avoid_relay_data { - cost += AVOID_RELAY_COST; - } - - graph.add_edge(*src_node_idx, *dst_node_idx, cost); - } - } - - (graph, start_node_idx.unwrap()) - } - - fn clean_expired_route_info(&self) { - let cur_version = self.next_hop_map_version.get(); - self.next_hop_map.retain(|_, v| { - // remove next hop map for peers we cannot reach. - v.version >= cur_version - }); - self.peer_infos.retain(|k, _| { - // remove peer info for peers we cannot forward to. - self.peer_reachable(*k) - }); - self.ipv4_peer_id_map.retain(|_, v| { - // remove ipv4 map for peers we cannot forward to. - self.peer_reachable(v.peer_id) - }); - self.ipv6_peer_id_map.retain(|_, v| { - // remove ipv6 map for peers we cannot forward to. - self.peer_reachable(v.peer_id) - }); - - shrink_dashmap(&self.peer_infos, None); - shrink_dashmap(&self.next_hop_map, None); - shrink_dashmap(&self.suppressed_peer_ids, None); - shrink_dashmap(&self.ipv4_peer_id_map, None); - shrink_dashmap(&self.ipv6_peer_id_map, None); - } - - fn gen_next_hop_map_with_least_hop( - &self, - graph: &PeerGraph, - start_node: &NodeIndex, - version: Version, - ) { - if graph.node_weight(*start_node).is_none() { - tracing::warn!( - ?start_node, - version, - "invalid start node for least-hop route rebuild" - ); - return; - } - let normalize_edge_cost = |e: petgraph::graph::EdgeReference| { - if *e.weight() >= AVOID_RELAY_COST { - AVOID_RELAY_COST + 1 - } else { - 1 - } - }; - // Step 1: 第一次 Dijkstra - 计算最短跳数 - let path_len_map = dijkstra(&graph, *start_node, None, normalize_edge_cost); - - // Step 2: 构建最短跳数子图(只保留属于最短路径和 AVOID RELAY 的边) - let mut subgraph: PeerGraph = PeerGraph::new(); - let mut start_node_idx = None; - for (node_idx, peer_id) in graph.node_references() { - let new_node_idx = subgraph.add_node(*peer_id); - if node_idx == *start_node { - start_node_idx = Some(new_node_idx); - } - } - - for edge in graph.edge_references() { - let (src, tgt) = graph.edge_endpoints(edge.id()).unwrap(); - let Some(src_path_len) = path_len_map.get(&src) else { - continue; - }; - let Some(tgt_path_len) = path_len_map.get(&tgt) else { - continue; - }; - if *src_path_len + normalize_edge_cost(edge) == *tgt_path_len { - subgraph.add_edge(src, tgt, *edge.weight()); - } - } - - // Step 3: 第二次 Dijkstra - 在子图上找代价最小的路径 - self.gen_next_hop_map_with_least_cost(&subgraph, &start_node_idx.unwrap(), version); - } - - fn gen_next_hop_map_with_least_cost( - &self, - graph: &PeerGraph, - start_node: &NodeIndex, - version: Version, - ) { - if graph.node_weight(*start_node).is_none() { - tracing::warn!( - ?start_node, - version, - "invalid start node for least-cost route rebuild" - ); - return; - } - let (costs, next_hops) = dijkstra_with_first_hop(&graph, *start_node, |e| *e.weight()); - - for (dst, (next_hop, path_len)) in next_hops.iter() { - let info = NextHopInfo { - next_hop_peer_id: *graph.node_weight(*next_hop).unwrap(), - path_latency: (*costs.get(dst).unwrap() % AVOID_RELAY_COST) as i32, - path_len: { *path_len }, - version, - }; - let dst_peer_id = *graph.node_weight(*dst).unwrap(); - self.next_hop_map - .entry(dst_peer_id) - .and_modify(|x| { - if x.version < version { - *x = info; - } - }) - .or_insert(info); - } - - self.next_hop_map_version.set_if_larger(version); - } - - fn build_from_synced_info( - &self, - my_peer_id: PeerId, - synced_info: &SyncedRouteInfo, - policy: NextHopPolicy, - cost_calc: &T, - ) { - let version = synced_info.version.get(); - self.sync_suppressed_peer_ids(synced_info); - - let local_proxy_cidrs = synced_info - .peer_infos - .read() - .get(&my_peer_id) - .into_iter() - .flat_map(|info| &info.proxy_cidrs) - .filter_map(|cidr| cidr.parse::().ok()) - .collect::>(); - - // build next hop map - let (graph, start_node) = - Self::build_peer_graph_from_synced_info(my_peer_id, synced_info, cost_calc); - - if graph.node_count() == 0 { - tracing::warn!("no peer in graph, cannot build next hop map"); - self.next_hop_map_version.set_if_larger(version); - self.clean_expired_route_info(); - return; - } - if start_node == NodeIndex::end() { - tracing::warn!( - ?my_peer_id, - version, - "my peer id is missing in graph, skip next-hop rebuild this round" - ); - self.next_hop_map_version.set_if_larger(version); - self.clean_expired_route_info(); - return; - } - - if matches!(policy, NextHopPolicy::LeastHop) { - self.gen_next_hop_map_with_least_hop(&graph, &start_node, version); - } else { - self.gen_next_hop_map_with_least_cost(&graph, &start_node, version); - }; - - let mut new_cidr_prefix_trie = PrefixMap::new(); - let mut new_cidr_v6_prefix_trie = PrefixMap::new(); - - // build peer_infos, ipv4_peer_id_map, cidr_peer_id_map - // only set map for peers we can reach. - for item in self.next_hop_map.iter() { - if item.version < version { - // skip if the next hop entry is outdated. (peer is unreachable) - continue; - } - - let peer_id = item.key(); - if !self.peer_reachable(*peer_id) { - continue; - } - - let Some(info) = synced_info.peer_infos.read().get(peer_id).cloned() else { - continue; - }; - - self.peer_infos.insert(*peer_id, info.clone()); - - let peer_id_and_version = PeerIdVersion { - peer_id: *peer_id, - version, - }; - - let is_new_peer_better = |old_peer: &PeerIdVersion| -> bool { - if peer_id_and_version.version > old_peer.version { - return true; - } - if peer_id_and_version.peer_id == old_peer.peer_id { - return false; - } - let old_next_hop = self.get_next_hop(old_peer.peer_id); - let new_next_hop = item.value(); - old_next_hop.is_none() || new_next_hop.path_len < old_next_hop.unwrap().path_len - }; - - if let Some(ipv4_addr) = info.ipv4_addr { - self.ipv4_peer_id_map - .entry(ipv4_addr.into()) - .and_modify(|v| { - if is_new_peer_better(v) { - *v = peer_id_and_version; - } - }) - .or_insert(peer_id_and_version); - } - - if let Some(ipv6_addr) = info.ipv6_addr.and_then(|x| x.address) { - self.ipv6_peer_id_map - .entry(ipv6_addr.into()) - .and_modify(|v| { - if is_new_peer_better(v) { - *v = peer_id_and_version; - } - }) - .or_insert(peer_id_and_version); - } - - if let Some(ipv6_addr) = info - .ipv6_public_addr_lease - .as_ref() - .and_then(|addr| addr.address) - { - self.ipv6_peer_id_map - .entry(ipv6_addr.into()) - .and_modify(|v| { - if is_new_peer_better(v) { - *v = peer_id_and_version; - } - }) - .or_insert(peer_id_and_version); - } - - for cidr in info.proxy_cidrs.iter() { - let Ok(cidr) = cidr.parse::() else { - tracing::warn!("invalid proxy cidr: {:?}, from peer: {:?}", cidr, peer_id); - continue; - }; - - if *peer_id != my_peer_id - && local_proxy_cidrs - .iter() - .any(|local_cidr| cidr_is_subset(&cidr, local_cidr)) - { - tracing::debug!( - ?peer_id, - ?my_peer_id, - ?local_proxy_cidrs, - ?cidr, - "skip remote proxy cidr covered by local announced proxy cidr while building route table" - ); - continue; - } - match cidr { - IpCidr::V4(cidr) => { - new_cidr_prefix_trie - .entry(cidr) - .and_modify(|e| { - // if ourself has same cidr, ensure here put my peer id, so we can know deadloop may happen. - if *peer_id == my_peer_id || is_new_peer_better(e) { - *e = peer_id_and_version; - } - }) - .or_insert(peer_id_and_version); - } - - IpCidr::V6(cidr) => { - new_cidr_v6_prefix_trie - .entry(cidr) - .and_modify(|e| { - // if ourself has same cidr, ensure here put my peer id, so we can know deadloop may happen. - if *peer_id == my_peer_id || is_new_peer_better(e) { - *e = peer_id_and_version; - } - }) - .or_insert(peer_id_and_version); - } - } - tracing::debug!( - "add cidr: {:?} to peer: {:?}, my peer id: {:?}", - cidr, - peer_id, - my_peer_id - ); - } - } - - self.cidr_peer_id_map.store(Arc::new(new_cidr_prefix_trie)); - self.cidr_v6_peer_id_map - .store(Arc::new(new_cidr_v6_prefix_trie)); - tracing::trace!( - my_peer_id = my_peer_id, - cidrs = ?self.cidr_peer_id_map.load(), - cidrs_v6 = ?self.cidr_v6_peer_id_map.load(), - "update peer cidr map" - ); - self.clean_expired_route_info(); - } - - fn get_peer_id_for_proxy(&self, ip: &IpAddr) -> Option { - match ip { - IpAddr::V4(ipv4) => self - .cidr_peer_id_map - .load() - .get_lpm(&Ipv4Cidr::new(*ipv4, 32).unwrap()) - .map(|x| x.1.peer_id), - IpAddr::V6(ipv6) => self - .cidr_v6_peer_id_map - .load() - .get_lpm(&Ipv6Cidr::new(*ipv6, 128).unwrap()) - .map(|x| x.1.peer_id), - } - } -} - type SessionId = u64; type AtomicSessionId = atomic_shim::AtomicU64; @@ -1841,6 +1861,14 @@ impl SessionTask { false } } + + async fn stop(&self) { + let task = self.task.lock().unwrap().take(); + if let Some(task) = task { + task.abort(); + let _ = task.await; + } + } } impl Drop for SessionTask { @@ -1895,6 +1923,7 @@ impl VersionAndTouchTime { // if we need to sync route info with one peer, we create a SyncRouteSession with that peer. #[derive(Debug)] +#[allow(dead_code)] struct SyncRouteSession { my_peer_id: PeerId, dst_peer_id: PeerId, @@ -2150,15 +2179,17 @@ impl Drop for SyncRouteSession { struct PeerRouteServiceImpl { my_peer_id: PeerId, my_peer_route_id: u64, - global_ctx: ArcGlobalCtx, + context: ArcPeerContext, sessions: DashMap>, + stopped: AtomicBool, + session_operation: AsyncRwLock<()>, interface: Mutex>, cost_calculator: std::sync::RwLock>, - route_table: RouteTable, - route_table_with_cost: RouteTable, - foreign_network_owner_map: DashMap>, + route_table: OspfRouteTable, + route_table_with_cost: OspfRouteTable, + foreign_network_owner_map: DashMap>, foreign_network_my_peer_id_map: DashMap<(String, PeerId), PeerId>, synced_route_info: SyncedRouteInfo, public_ipv6_service: std::sync::Mutex>, @@ -2179,7 +2210,7 @@ impl Debug for PeerRouteServiceImpl { f.debug_struct("PeerRouteServiceImpl") .field("my_peer_id", &self.my_peer_id) .field("my_peer_route_id", &self.my_peer_route_id) - .field("network", &self.global_ctx.get_network_identity()) + .field("network", &self.context.network_identity()) .field("sessions", &self.sessions) .field("route_table", &self.route_table) .field("route_table_with_cost", &self.route_table_with_cost) @@ -2197,24 +2228,28 @@ impl Debug for PeerRouteServiceImpl { } } +#[allow(dead_code)] impl PeerRouteServiceImpl { - fn new(my_peer_id: PeerId, global_ctx: ArcGlobalCtx) -> Self { + fn new(my_peer_id: PeerId, context: ArcPeerContext) -> Self { PeerRouteServiceImpl { my_peer_id, my_peer_route_id: rand::random(), - global_ctx, + context: context.clone(), sessions: DashMap::new(), + stopped: AtomicBool::new(false), + session_operation: AsyncRwLock::new(()), interface: Mutex::new(None), cost_calculator: std::sync::RwLock::new(Some(Box::new(DefaultRouteCostCalculator))), - route_table: RouteTable::new(), - route_table_with_cost: RouteTable::new(), + route_table: OspfRouteTable::new(), + route_table_with_cost: OspfRouteTable::new(), foreign_network_owner_map: DashMap::new(), foreign_network_my_peer_id_map: DashMap::new(), synced_route_info: SyncedRouteInfo { + default_easytier_version: context.easytier_version(), peer_infos: RwLock::new(OrderedHashMap::new()), raw_peer_infos: DashMap::new(), conn_map: RwLock::new(OrderedHashMap::new()), @@ -2242,25 +2277,11 @@ impl PeerRouteServiceImpl { } } - fn get_my_secret_digest(&self) -> Option> { - let ni = self.global_ctx.get_network_identity(); - ni.network_secret_digest.map(|d| d.to_vec()) - } - - #[cfg(test)] - fn is_active_non_reusable_credential_peer(&self, peer_id: PeerId) -> bool { - peer_id == self.my_peer_id || self.route_table.topology_peer_reachable(peer_id) - } - fn is_credential_node(&self) -> bool { - self.global_ctx - .get_network_identity() - .network_secret - .is_none() + self.context.network_identity().network_secret.is_none() && self - .global_ctx - .config - .get_secure_mode() + .context + .secure_mode() .map(|c| c.enabled) .unwrap_or(false) } @@ -2348,15 +2369,6 @@ impl PeerRouteServiceImpl { (snapshot.generation, snapshot.peers.clone()) } - async fn list_peers_from_interface>(&self) -> T { - self.interface_peer_snapshot() - .await - .peers - .iter() - .copied() - .collect() - } - async fn get_peer_identity_type_from_interface( &self, peer_id: PeerId, @@ -2389,7 +2401,7 @@ impl PeerRouteServiceImpl { self.synced_route_info.update_my_peer_info( self.my_peer_id, self.my_peer_route_id, - &self.global_ctx, + self.context.as_ref(), *self.self_public_ipv6_addr_lease.lock().unwrap(), ) } @@ -2428,7 +2440,7 @@ impl PeerRouteServiceImpl { let last_time = self.last_update_my_foreign_network.load(); if last_time.is_some() && last_time.unwrap().elapsed().as_secs() - < use_global_var!(OSPF_UPDATE_MY_GLOBAL_FOREIGN_NETWORK_INTERVAL_SEC) + < self.context.ospf_update_my_foreign_network_interval_sec() { return false; } @@ -2460,17 +2472,18 @@ impl PeerRouteServiceImpl { .begin_update(); let calc_locked = self.cost_calculator.read().unwrap(); + let route_snapshot = self.synced_route_info.route_snapshot(); - self.route_table.build_from_synced_info( + self.route_table.build_from_snapshot( self.my_peer_id, - &self.synced_route_info, + &route_snapshot, NextHopPolicy::LeastHop, calc_locked.as_ref().unwrap(), ); - self.route_table_with_cost.build_from_synced_info( + self.route_table_with_cost.build_from_snapshot( self.my_peer_id, - &self.synced_route_info, + &route_snapshot, NextHopPolicy::LeastCost, calc_locked.as_ref().unwrap(), ); @@ -2497,7 +2510,7 @@ impl PeerRouteServiceImpl { { continue; } - let network_identity = NetworkIdentity { + let network_identity = CoreNetworkIdentity { network_name: key.network_name.clone(), network_secret: None, network_secret_digest: Some( @@ -2529,16 +2542,8 @@ impl PeerRouteServiceImpl { .unwrap_or(false) } - fn handle_global_ctx_event(&self, event: &GlobalCtxEvent) { - if matches!( - event, - GlobalCtxEvent::PeerAdded(_) - | GlobalCtxEvent::PeerRemoved(_) - | GlobalCtxEvent::PeerConnAdded(_) - | GlobalCtxEvent::PeerConnRemoved(_) - ) { - self.mark_interface_peers_dirty(); - } + fn handle_peer_context_event(&self, _event: &PeerContextEvent) { + self.mark_interface_peers_dirty(); } fn update_route_table_and_cached_local_conn_bitmap(&self) { @@ -2796,11 +2801,8 @@ impl PeerRouteServiceImpl { async fn refresh_acl_groups(&self) -> bool { let my_peer_info_updated = self.update_my_peer_info(); - let trust_admin_groups_without_proof = self - .global_ctx - .get_network_identity() - .network_secret - .is_none(); + let trust_admin_groups_without_proof = + self.context.network_identity().network_secret.is_none(); let peer_infos: Vec<_> = self .synced_route_info @@ -2811,7 +2813,7 @@ impl PeerRouteServiceImpl { .collect(); self.synced_route_info.verify_and_update_group_trusts( &peer_infos, - &self.global_ctx.get_acl_group_declarations(), + &self.context.acl_group_declarations(), trust_admin_groups_without_proof, ); @@ -2832,7 +2834,7 @@ impl PeerRouteServiceImpl { } fn refresh_credential_trusts(&self) -> Vec { - let network_identity = self.global_ctx.get_network_identity(); + let network_identity = self.context.network_identity(); let (untrusted, global_trusted_keys, _) = self .synced_route_info .verify_and_update_credential_trusts_with_active_peers_protecting( @@ -2840,14 +2842,17 @@ impl PeerRouteServiceImpl { |_| true, Some(self.my_peer_id), ); - self.global_ctx - .update_trusted_keys(global_trusted_keys, &network_identity.network_name); + PeerContext::update_trusted_keys( + self.context.as_ref(), + global_trusted_keys, + &network_identity.network_name, + ); untrusted } fn refresh_credential_trusts_with_current_topology(&self) -> Vec { - let network_identity = self.global_ctx.get_network_identity(); + let network_identity = self.context.network_identity(); // Non-reusable credential owner election depends on reachability, so rebuild the // route table from the latest synced peer/conn state before checking active peers. @@ -2862,8 +2867,11 @@ impl PeerRouteServiceImpl { }, Some(self.my_peer_id), ); - self.global_ctx - .update_trusted_keys(global_trusted_keys, &network_identity.network_name); + PeerContext::update_trusted_keys( + self.context.as_ref(), + global_trusted_keys, + &network_identity.network_name, + ); if !untrusted.is_empty() || suppressed_changed { self.update_route_table_and_cached_local_conn_bitmap(); @@ -2982,33 +2990,25 @@ impl PeerRouteServiceImpl { fn build_sync_route_raw_req( req: &SyncRouteInfoRequest, - raw_peer_infos: &DashMap, - ) -> DynamicMessage { - use prost_reflect::Value; - - let mut req_dynamic_msg = DynamicMessage::new(SyncRouteInfoRequest::default().descriptor()); - req_dynamic_msg.transcode_from(req).unwrap(); - - let peer_infos = req.peer_infos.as_ref().map(|x| &x.items); - if let Some(peer_infos) = peer_infos { - let mut peer_info_raws = Vec::new(); - for peer_info in peer_infos.iter() { - if let Some(info) = raw_peer_infos.get(&peer_info.peer_id) { - peer_info_raws.push(Value::Message(info.clone())); - } else { - let mut p = DynamicMessage::new(RoutePeerInfo::default().descriptor()); - p.transcode_from(peer_info).unwrap(); - peer_info_raws.push(Value::Message(p)); - } - } - - let mut peer_infos = DynamicMessage::new(RoutePeerInfos::default().descriptor()); - peer_infos.set_field_by_name("items", Value::List(peer_info_raws)); - - req_dynamic_msg.set_field_by_name("peer_infos", Value::Message(peer_infos)); - } - - req_dynamic_msg + raw_peer_infos: &DashMap, + ) -> Result, route_peer_wire::WireError> { + let infos = req + .peer_infos + .as_ref() + .map(|peer_infos| { + peer_infos + .items + .iter() + .map(|peer_info| { + raw_peer_infos + .get(&peer_info.peer_id) + .map(|raw| raw.clone()) + .unwrap_or_else(|| raw_route_peer_info(peer_info)) + }) + .collect::>() + }) + .unwrap_or_default(); + route_peer_wire::encode_sync_route_request(req, &infos) } async fn sync_route_with_peer( @@ -3059,7 +3059,7 @@ impl PeerRouteServiceImpl { .scoped_client::>( self.my_peer_id, dst_peer_id, - self.global_ctx.get_network_name(), + self.context.network_name(), ); let sync_route_info_req = SyncRouteInfoRequest { @@ -3073,14 +3073,17 @@ impl PeerRouteServiceImpl { let mut ctrl = BaseController::default(); ctrl.set_timeout_ms(3000); - ctrl.set_raw_input( - Self::build_sync_route_raw_req( - &sync_route_info_req, - &self.synced_route_info.raw_peer_infos, - ) - .encode_to_vec() - .into(), - ); + let raw_request = match Self::build_sync_route_raw_req( + &sync_route_info_req, + &self.synced_route_info.raw_peer_infos, + ) { + Ok(raw_request) => raw_request, + Err(error) => { + tracing::error!(?error, "failed to encode raw OSPF route request"); + return false; + } + }; + ctrl.set_raw_input(raw_request.into()); drop(_session_lock); let ret = rpc_stub @@ -3093,7 +3096,7 @@ impl PeerRouteServiceImpl { ret, sync_route_info_req, session, - self.global_ctx.network, + self.context.network_identity(), next_last_sync_succ_timestamp ); @@ -3113,7 +3116,7 @@ impl PeerRouteServiceImpl { Ok(resp) => { if let Some(err) = resp.error { if err == Error::DuplicatePeerId as i32 { - if !self.global_ctx.get_feature_flags().is_public_server { + if !self.context.feature_flags().is_public_server { panic!("duplicate peer id"); } } else { @@ -3203,24 +3206,6 @@ impl Debug for RouteSessionManager { } } -fn get_raw_peer_infos(req_raw_input: &mut bytes::Bytes) -> Option> { - let sync_req_dynamic_msg = - DynamicMessage::decode(SyncRouteInfoRequest::default().descriptor(), req_raw_input) - .unwrap(); - - let peer_infos = sync_req_dynamic_msg.get_field_by_name("peer_infos")?; - - let infos = peer_infos - .as_message()? - .get_field_by_name("items")? - .as_list()? - .iter() - .map(|x| x.as_message().unwrap().clone()) - .collect(); - - Some(infos) -} - #[async_trait::async_trait] impl OspfRouteRpc for RouteSessionManager { type Controller = BaseController; @@ -3236,8 +3221,23 @@ impl OspfRouteRpc for RouteSessionManager { let conn_info = request.conn_info; let foreign_network = request.foreign_network_infos; let raw_peer_infos = if let Some(peer_infos_ref) = &peer_infos { - let r = get_raw_peer_infos(&mut ctrl.get_raw_input().unwrap()).unwrap(); - assert_eq!(r.len(), peer_infos_ref.len()); + let raw_input = ctrl.get_raw_input().ok_or_else(|| { + rpc_types::error::Error::MalformatRpcPacket( + "OSPF route request is missing raw protobuf input".to_owned(), + ) + })?; + let r = extract_route_peer_infos(&raw_input).map_err(anyhow::Error::new)?; + let raw_matches_decoded = r.len() == peer_infos_ref.len() + && r.iter().zip(peer_infos_ref).all(|(raw, decoded)| { + RoutePeerInfo::decode(raw.clone()) + .map(|raw_decoded| raw_decoded.peer_id == decoded.peer_id) + .unwrap_or(false) + }); + if !raw_matches_decoded { + return Err(rpc_types::error::Error::MalformatRpcPacket( + "raw and decoded OSPF RoutePeerInfo values do not match".to_owned(), + )); + } Some(r) } else { None @@ -3265,6 +3265,7 @@ impl OspfRouteRpc for RouteSessionManager { } } +#[allow(dead_code)] impl RouteSessionManager { fn new(service_impl: Arc, peer_rpc: Arc) -> Self { RouteSessionManager { @@ -3325,14 +3326,14 @@ impl RouteSessionManager { drop(service_impl); drop(peer_rpc); - tokio::time::sleep(Duration::from_millis(retry_delay_ms)).await; + crate::foundation::time::sleep(Duration::from_millis(retry_delay_ms)).await; retry_delay_ms = (retry_delay_ms * 2).min(RETRY_MAX_MS); } sync_now = sync_now.resubscribe(); select! { - _ = tokio::time::sleep(Duration::from_secs(1)) => {} + _ = crate::foundation::time::sleep(Duration::from_secs(1)) => {} ret = sync_now.recv() => if let Err(e) = ret { tracing::debug!(?e, "session_task sync_now recv failed, ospf route may exit"); break; @@ -3365,6 +3366,9 @@ impl RouteSessionManager { let Some(service_impl) = self.service_impl.upgrade() else { return Err(Error::Stopped); }; + if service_impl.stopped.load(Ordering::Acquire) { + return Err(Error::Stopped); + } tracing::info!(?service_impl.my_peer_id, ?peer_id, "start ospf sync session"); @@ -3379,7 +3383,7 @@ impl RouteSessionManager { loop { let mut recv = self.sync_now_broadcast.subscribe(); select! { - _ = tokio::time::sleep(Duration::from_millis(next_sleep_ms)) => {} + _ = crate::foundation::time::sleep(Duration::from_millis(next_sleep_ms)) => {} _ = recv.recv() => {} } @@ -3526,12 +3530,12 @@ impl RouteSessionManager { &self, from_peer_id: PeerId, peer_infos: &[RoutePeerInfo], - raw_peer_infos: &[DynamicMessage], + raw_peer_infos: &[RawRoutePeerInfo], credential: &TrustedCredentialPubkey, - ) -> Option<(RoutePeerInfo, DynamicMessage)> { + ) -> Option<(RoutePeerInfo, RawRoutePeerInfo)> { let info_idx = peer_infos.iter().position(|p| p.peer_id == from_peer_id)?; let mut info = peer_infos[info_idx].clone(); - let mut raw_info = raw_peer_infos[info_idx].clone(); + let raw_info = raw_peer_infos[info_idx].clone(); let allowed_cidrs = &credential.allowed_proxy_cidrs; // Filter proxy_cidrs to only those allowed by credential if !allowed_cidrs.is_empty() { @@ -3545,7 +3549,7 @@ impl RouteSessionManager { info.proxy_cidrs.clear(); } SyncedRouteInfo::mark_credential_peer(&mut info, true); - patch_raw_from_info(&mut raw_info, &info, &["proxy_cidrs", "feature_flag"]); + let raw_info = patch_credential_route_peer_info(&raw_info, &info.proxy_cidrs).ok()?; Some((info, raw_info)) } @@ -3556,13 +3560,17 @@ impl RouteSessionManager { from_session_id: SessionId, is_initiator: bool, peer_infos: Option>, - raw_peer_infos: Option>, + raw_peer_infos: Option>, conn_info: Option, foreign_network: Option, ) -> Result { let Some(service_impl) = self.service_impl.upgrade() else { return Err(Error::Stopped); }; + let _session_operation = service_impl.session_operation.read().await; + if service_impl.stopped.load(Ordering::Acquire) { + return Err(Error::Stopped); + } let my_peer_id = service_impl.my_peer_id; let session = self.get_or_start_session(from_peer_id)?; @@ -3600,7 +3608,6 @@ impl RouteSessionManager { if let Some(peer_infos) = &peer_infos { // Step 9b: credential peers can only propagate their own route info - // patch_raw_from_info(&mut raw, info, &["proxy_cidrs", "feature_flag"]); let (pi, rpi) = if from_is_credential { if let Some(ret) = self.extract_credential_peer_info( from_peer_id, @@ -3617,8 +3624,8 @@ impl RouteSessionManager { }; if !pi.is_empty() { let trust_admin_groups_without_proof = service_impl - .global_ctx - .get_network_identity() + .context + .network_identity() .network_secret .is_none(); service_impl.synced_route_info.update_peer_infos( @@ -3632,7 +3639,7 @@ impl RouteSessionManager { .synced_route_info .verify_and_update_group_trusts( pi, - &service_impl.global_ctx.get_acl_group_declarations(), + &service_impl.context.acl_group_declarations(), trust_admin_groups_without_proof, ); session.update_dst_saved_peer_info_version(pi, from_peer_id); @@ -3789,7 +3796,7 @@ impl PublicIpv6SyncTrigger for OspfPublicIpv6SyncTrigger { pub struct PeerRoute { my_peer_id: PeerId, - global_ctx: ArcGlobalCtx, + context: ArcPeerContext, peer_rpc: Weak, service_impl: Arc, @@ -3812,13 +3819,14 @@ impl Debug for PeerRoute { impl PeerRoute { pub fn new( my_peer_id: PeerId, - global_ctx: ArcGlobalCtx, + context: ArcPeerContext, + public_ipv6_runtime: Arc, peer_rpc: Arc, ) -> Arc { - let service_impl = Arc::new(PeerRouteServiceImpl::new(my_peer_id, global_ctx.clone())); + let service_impl = Arc::new(PeerRouteServiceImpl::new(my_peer_id, context.clone())); let session_mgr = RouteSessionManager::new(service_impl.clone(), peer_rpc.clone()); let public_ipv6_service = Arc::new(PublicIpv6Service::new( - global_ctx.clone(), + public_ipv6_runtime, Arc::downgrade(&peer_rpc), Arc::new(OspfPublicIpv6RouteHandle { service_impl: Arc::downgrade(&service_impl), @@ -3831,7 +3839,7 @@ impl PeerRoute { Arc::new(PeerRoute { my_peer_id, - global_ctx, + context, peer_rpc: Arc::downgrade(&peer_rpc), service_impl, @@ -3844,7 +3852,7 @@ impl PeerRoute { async fn clear_expired_peer(service_impl: Arc) { loop { - tokio::time::sleep(Duration::from_secs(60)).await; + crate::foundation::time::sleep(Duration::from_secs(60)).await; service_impl.clear_expired_peer().await; // TODO: use debug log level for this. tracing::debug!(?service_impl, "clear_expired_peer"); @@ -3862,7 +3870,8 @@ impl PeerRoute { service_impl: Arc, session_mgr: RouteSessionManager, ) { - let mut global_event_receiver = service_impl.global_ctx.subscribe(); + let mut peer_event_receiver = service_impl.context.subscribe_peer_events(); + let mut runtime_change_receiver = service_impl.context.subscribe_runtime_changes(); service_impl.mark_interface_peers_dirty(); loop { if service_impl.update_my_infos().await { @@ -3878,17 +3887,39 @@ impl PeerRoute { } } + let peer_event = async { + match peer_event_receiver.as_mut() { + Some(receiver) => Some(receiver.recv().await), + None => std::future::pending().await, + } + }; + let runtime_change = async { + match runtime_change_receiver.as_mut() { + Some(receiver) => Some(receiver.changed().await), + None => std::future::pending().await, + } + }; + select! { - ev = global_event_receiver.recv() => { + Some(ev) = peer_event => { if let Ok(ev_ref) = &ev { - service_impl.handle_global_ctx_event(ev_ref); + service_impl.handle_peer_context_event(ev_ref); } else { service_impl.mark_interface_peers_dirty(); - global_event_receiver = global_event_receiver.resubscribe(); + peer_event_receiver = service_impl.context.subscribe_peer_events(); } - tracing::info!(?ev, "global event received in update_my_peer_info_routine"); + tracing::info!( + ?ev, + "peer context event received in update_my_peer_info_routine" + ); } - _ = tokio::time::sleep(Duration::from_secs(1)) => {} + Some(change) = runtime_change => { + if change.is_err() { + runtime_change_receiver = + service_impl.context.subscribe_runtime_changes(); + } + } + _ = crate::foundation::time::sleep(Duration::from_secs(1)) => {} } } } @@ -3904,11 +3935,11 @@ impl PeerRoute { peer_rpc.rpc_server().registry().register( OspfRouteRpcServer::new(self.session_mgr.clone()), - &self.global_ctx.get_network_name(), + &self.context.network_name(), ); peer_rpc.rpc_server().registry().register( PublicIpv6AddrRpcServer::new(self.public_ipv6_service.rpc_server()), - &self.global_ctx.get_network_name(), + &self.context.network_name(), ); self.tasks @@ -3942,30 +3973,61 @@ impl PeerRoute { .unwrap() .spawn(self.public_ipv6_service.clone().client_routine()); } -} - -impl Drop for PeerRoute { - fn drop(&mut self) { - tracing::debug!( - self.my_peer_id, - network = ?self.global_ctx.get_network_identity(), - service = ?self.service_impl, - "PeerRoute drop" - ); + fn unregister_rpc_services(&self) { let Some(peer_rpc) = self.peer_rpc.upgrade() else { return; }; peer_rpc.rpc_server().registry().unregister( OspfRouteRpcServer::new(self.session_mgr.clone()), - &self.global_ctx.get_network_name(), + &self.context.network_name(), ); peer_rpc.rpc_server().registry().unregister( PublicIpv6AddrRpcServer::new(self.public_ipv6_service.rpc_server()), - &self.global_ctx.get_network_name(), + &self.context.network_name(), ); } + + async fn stop(&self) { + self.service_impl.stopped.store(true, Ordering::Release); + self.unregister_rpc_services(); + + let mut tasks = { + let mut tasks = self.tasks.lock().unwrap(); + std::mem::replace(&mut *tasks, JoinSet::new()) + }; + tasks.abort_all(); + while tasks.join_next().await.is_some() {} + + let _session_operation = self.service_impl.session_operation.write().await; + let sessions = self + .service_impl + .sessions + .iter() + .map(|session| session.value().clone()) + .collect::>(); + self.service_impl.sessions.clear(); + self.service_impl.sessions.shrink_to_fit(); + for session in sessions { + session.task.stop().await; + } + + *self.service_impl.interface.lock().await = None; + } +} + +impl Drop for PeerRoute { + fn drop(&mut self) { + tracing::debug!( + self.my_peer_id, + network = ?self.context.network_identity(), + service = ?self.service_impl, + "PeerRoute drop" + ); + + self.unregister_rpc_services(); + } } #[async_trait::async_trait] @@ -3976,7 +4038,9 @@ impl Route for PeerRoute { Ok(1) } - async fn close(&self) {} + async fn close(&self) { + self.stop().await; + } async fn get_next_hop(&self, dst_peer_id: PeerId) -> Option { let route_table = &self.service_impl.route_table; @@ -4000,57 +4064,22 @@ impl Route for PeerRoute { .map(|x| x.next_hop_peer_id) } - async fn list_routes(&self) -> Vec { - let route_table = &self.service_impl.route_table; - let route_table_with_cost = &self.service_impl.route_table_with_cost; - let mut routes = Vec::new(); - for item in route_table.peer_infos.iter() { - if *item.key() == self.my_peer_id { - continue; - } - let Some(next_hop_peer) = route_table.get_next_hop(*item.key()) else { - continue; - }; - let next_hop_peer_latency_first = route_table_with_cost.get_next_hop(*item.key()); - let mut route: crate::proto::api::instance::Route = item.value().clone().into(); - route.next_hop_peer_id = next_hop_peer.next_hop_peer_id; - route.cost = next_hop_peer.path_len as i32; - route.path_latency = next_hop_peer.path_latency; - - route.next_hop_peer_id_latency_first = - next_hop_peer_latency_first.map(|x| x.next_hop_peer_id); - route.cost_latency_first = next_hop_peer_latency_first.map(|x| x.path_len as i32); - route.path_latency_latency_first = next_hop_peer_latency_first.map(|x| x.path_latency); - - route.feature_flag = item.feature_flag; - - routes.push(route); - } - routes + async fn list_routes(&self) -> Vec { + self.service_impl + .route_table + .list_routes(self.my_peer_id, &self.service_impl.route_table_with_cost) } async fn list_proxy_cidrs(&self) -> BTreeSet { - let my_peer_id = self.my_peer_id; self.service_impl .route_table - .cidr_peer_id_map - .load() - .iter() - .filter(|(_, pv)| pv.peer_id != my_peer_id) - .map(|(cidr, _)| *cidr) - .collect() + .list_proxy_cidrs_excluding(self.my_peer_id) } async fn list_proxy_cidrs_v6(&self) -> BTreeSet { - let my_peer_id = self.my_peer_id; self.service_impl .route_table - .cidr_v6_peer_id_map - .load() - .iter() - .filter(|(_, pv)| pv.peer_id != my_peer_id) - .map(|(cidr, _)| *cidr) - .collect() + .list_proxy_cidrs_v6_excluding(self.my_peer_id) } async fn list_public_ipv6_routes(&self) -> BTreeSet { @@ -4065,14 +4094,12 @@ impl Route for PeerRoute { self.public_ipv6_service.provider_peer_id_for_client() } - async fn get_local_public_ipv6_info( - &self, - ) -> crate::proto::api::instance::ListPublicIpv6InfoResponse { + async fn get_local_public_ipv6_info(&self) -> CoreListPublicIpv6InfoResponse { let Some((provider, leases)) = self.public_ipv6_service.local_provider_state() else { - return crate::proto::api::instance::ListPublicIpv6InfoResponse::default(); + return CoreListPublicIpv6InfoResponse::default(); }; - crate::proto::api::instance::ListPublicIpv6InfoResponse { + CoreListPublicIpv6InfoResponse { provider_prefix: Some( Ipv6Inet::new( provider.prefix.first_address(), @@ -4083,7 +4110,7 @@ impl Route for PeerRoute { ), provider_leases: leases .into_iter() - .map(|lease| crate::proto::api::instance::PublicIpv6LeaseInfo { + .map(|lease| CorePublicIpv6LeaseInfo { peer_id: lease.peer_id, inst_id: lease.inst_id.to_string(), leased_addr: Some(lease.addr.into()), @@ -4100,15 +4127,12 @@ impl Route for PeerRoute { async fn get_peer_id_by_ipv4(&self, ipv4_addr: &Ipv4Addr) -> Option { let route_table = &self.service_impl.route_table; - if let Some(p) = route_table.ipv4_peer_id_map.get(ipv4_addr) { - return Some(p.peer_id); + if let Some(peer_id) = route_table.get_peer_id_by_ipv4(ipv4_addr) { + return Some(peer_id); } // only get peer id for proxy when the dst ipv4 is not in same network with us - if self - .global_ctx - .is_ip_in_same_network(&std::net::IpAddr::V4(*ipv4_addr)) - { + if PeerContext::is_ip_in_same_network(self.context.as_ref(), &IpAddr::V4(*ipv4_addr)) { tracing::trace!(?ipv4_addr, "ipv4 addr is in same network with us"); return None; } @@ -4123,15 +4147,12 @@ impl Route for PeerRoute { async fn get_peer_id_by_ipv6(&self, ipv6_addr: &Ipv6Addr) -> Option { let route_table = &self.service_impl.route_table; - if let Some(p) = route_table.ipv6_peer_id_map.get(ipv6_addr) { - return Some(p.peer_id); + if let Some(peer_id) = route_table.get_peer_id_by_ipv6(ipv6_addr) { + return Some(peer_id); } // only get peer id for proxy when the dst ipv4 is not in same network with us - if self - .global_ctx - .is_ip_in_same_network(&std::net::IpAddr::V6(*ipv6_addr)) - { + if PeerContext::is_ip_in_same_network(self.context.as_ref(), &IpAddr::V6(*ipv6_addr)) { tracing::trace!(?ipv6_addr, "ipv6 addr is in same network with us"); return None; } @@ -4187,7 +4208,7 @@ impl Route for PeerRoute { async fn list_peers_own_foreign_network( &self, - network_identity: &NetworkIdentity, + network_identity: &CoreNetworkIdentity, ) -> Vec { self.service_impl .foreign_network_owner_map @@ -4208,11 +4229,7 @@ impl Route for PeerRoute { } async fn get_peer_info(&self, peer_id: PeerId) -> Option { - self.service_impl - .route_table - .peer_infos - .get(&peer_id) - .map(|x| x.clone()) + self.service_impl.route_table.get_peer_info(peer_id) } async fn get_peer_info_last_update_time(&self) -> Instant { @@ -4230,106 +4247,39 @@ impl Route for PeerRoute { } } -impl PeerPacketFilter for Arc {} +impl PeerPacketFilter for PeerRoute {} #[cfg(test)] mod tests { - use cidr::{Ipv4Cidr, Ipv4Inet, Ipv6Inet}; - use dashmap::DashMap; + use super::*; + use crate::packet::ZCPacket; + use crate::peers::peer_rpc::PeerRpcManagerTransport; + use crate::peers::route::{DefaultRouteCostCalculator, RouteInterface}; + use crate::peers::test_support::NoopPeerContext; use parking_lot::Mutex; - use prefix_trie::PrefixMap; - use prost::Message; - use prost_reflect::{DynamicMessage, ReflectMessage}; - use prost_wkt_types::Timestamp; - use std::net::IpAddr; - use std::{ - collections::{BTreeSet, HashMap}, - sync::{ - Arc, - atomic::{AtomicU32, Ordering}, - }, - time::{Duration, SystemTime}, - }; + use tokio::sync::Notify; - use super::{ - NextHopInfo, PeerRoute, REMOVE_DEAD_PEER_INFO_AFTER, RouteConnInfo, SyncRouteSession, - }; - use crate::proto::common::TimestampExt; - use crate::{ - common::{ - PeerId, - config::NetworkIdentity, - global_ctx::{ - GlobalCtxEvent, TrustedKeySource, - tests::{get_mock_global_ctx, get_mock_global_ctx_with_network}, - }, - }, - connector::udp_hole_punch::tests::replace_stun_info_collector, - peers::{ - create_packet_recv_chan, - peer_manager::{PeerManager, RouteAlgoType}, - peer_ospf_route::{FORCE_USE_CONN_LIST, PeerIdVersion, PeerRouteServiceImpl}, - route_trait::{NextHopPolicy, Route, RouteCostCalculatorInterface, RouteInterface}, - tests::{connect_peer_manager, create_mock_peer_manager, wait_route_appear}, - }, - proto::{ - acl::{Acl, AclV1, GroupIdentity, GroupInfo}, - common::{NatType, PeerFeatureFlag}, - peer_rpc::{ - ForeignNetworkRouteInfoEntry, ForeignNetworkRouteInfoKey, PeerGroupInfo, - PeerIdentityType, RoutePeerInfo, RoutePeerInfos, SyncRouteInfoRequest, - TrustedCredentialPubkey, TrustedCredentialPubkeyProof, - }, - }, - tunnel::common::tests::wait_for_condition, - }; - use base64::Engine as _; - use base64::prelude::BASE64_STANDARD; - - struct AuthOnlyInterface { - my_peer_id: PeerId, - identity_type: DashMap, - peer_public_key: DashMap>, - } - - #[async_trait::async_trait] - impl RouteInterface for AuthOnlyInterface { - async fn list_peers(&self) -> Vec { - Vec::new() - } - - fn my_peer_id(&self) -> PeerId { - self.my_peer_id - } - - async fn get_peer_public_key(&self, peer_id: PeerId) -> Option> { - self.peer_public_key - .get(&peer_id) - .map(|x| x.value().clone()) - } - - async fn get_peer_identity_type(&self, peer_id: PeerId) -> Option { - self.identity_type.get(&peer_id).map(|x| *x.value()) + impl PeerRouteServiceImpl { + pub(crate) async fn list_peers_from_interface>(&self) -> T { + self.interface_peer_snapshot() + .await + .peers + .iter() + .copied() + .collect() } } - struct TrackingInterface { - my_peer_id: PeerId, - closed_peers: Arc>>, - } - - #[async_trait::async_trait] - impl RouteInterface for TrackingInterface { - async fn list_peers(&self) -> Vec { - Vec::new() - } - - fn my_peer_id(&self) -> PeerId { - self.my_peer_id - } - - async fn close_peer(&self, peer_id: PeerId) { - self.closed_peers.lock().push(peer_id); + impl PeerRoute { + pub(crate) fn task_count(&self) -> usize { + let route_tasks = self.tasks.lock().unwrap().len(); + let session_tasks = self + .service_impl + .sessions + .iter() + .filter(|session| session.task.is_running()) + .count(); + route_tasks + session_tasks } } @@ -4341,6 +4291,74 @@ mod tests { get_peer_identity_type_calls: Arc, } + struct BlockingInterface { + entered: Arc, + release: Arc, + } + + #[async_trait::async_trait] + impl RouteInterface for BlockingInterface { + async fn list_peers(&self) -> Vec { + self.entered.notify_one(); + self.release.notified().await; + vec![2] + } + + async fn get_peer_identity_type(&self, _peer_id: PeerId) -> Option { + Some(PeerIdentityType::Admin) + } + + fn my_peer_id(&self) -> PeerId { + 1 + } + } + + struct TestPeerRpcTransport; + + #[async_trait::async_trait] + impl PeerRpcManagerTransport for TestPeerRpcTransport { + fn my_peer_id(&self) -> PeerId { + 1 + } + + async fn send(&self, _msg: ZCPacket, _dst_peer_id: PeerId) -> anyhow::Result<()> { + Ok(()) + } + + async fn recv(&self) -> anyhow::Result { + std::future::pending().await + } + } + + struct TestPublicIpv6Runtime; + + #[async_trait::async_trait] + impl PublicIpv6Runtime for TestPublicIpv6Runtime { + fn ipv6_public_addr_auto(&self) -> bool { + false + } + + fn ipv6_public_addr_provider(&self) -> bool { + false + } + + fn instance_id(&self) -> uuid::Uuid { + uuid::Uuid::nil() + } + + fn network_name(&self) -> String { + "default".to_owned() + } + + async fn collect_reserved_public_ipv6_addrs(&self, _prefix: Ipv6Cidr) -> HashSet { + HashSet::new() + } + + fn public_ipv6_lease_changed(&self, _old: Option, _new: Option) {} + + fn public_ipv6_routes_changed(&self, _added: Vec, _removed: Vec) {} + } + #[async_trait::async_trait] impl RouteInterface for CountingInterface { async fn list_peers(&self) -> Vec { @@ -4363,9 +4381,34 @@ mod tests { } } + fn test_service_impl(my_peer_id: PeerId) -> PeerRouteServiceImpl { + PeerRouteServiceImpl::new(my_peer_id, Arc::new(NoopPeerContext::default())) + } + + fn peer(peer_id: PeerId) -> OspfPeerInfo { + OspfPeerInfo { + peer_id, + info: RoutePeerInfo { + peer_id, + version: 1, + ..Default::default() + }, + } + } + + fn connected( + peer_id: PeerId, + connected_peers: impl IntoIterator, + ) -> OspfPeerConnInfo { + OspfPeerConnInfo { + peer_id, + connected_peers: connected_peers.into_iter().collect(), + } + } + #[tokio::test] async fn interface_peer_cache_refreshes_only_when_marked_dirty() { - let service_impl = PeerRouteServiceImpl::new(1, get_mock_global_ctx()); + let service_impl = test_service_impl(1); let peers = Arc::new(Mutex::new(vec![2, 3])); let peer_identity_types = Arc::new(Mutex::new(HashMap::new())); let list_peers_calls = Arc::new(AtomicU32::new(0)); @@ -4386,7 +4429,7 @@ mod tests { assert_eq!(list_peers_calls.load(Ordering::Relaxed), 1); *peers.lock() = vec![2, 4]; - service_impl.handle_global_ctx_event(&GlobalCtxEvent::PeerConnAdded(Default::default())); + service_impl.handle_peer_context_event(&PeerContextEvent::PeerConnAdded); let third: BTreeSet<_> = service_impl.list_peers_from_interface().await; assert_eq!(third, BTreeSet::from([2, 4])); @@ -4395,7 +4438,7 @@ mod tests { #[tokio::test] async fn update_my_conn_info_skips_interface_scan_when_topology_is_unchanged() { - let service_impl = PeerRouteServiceImpl::new(1, get_mock_global_ctx()); + let service_impl = test_service_impl(1); let peers = Arc::new(Mutex::new(vec![2, 3])); let peer_identity_types = Arc::new(Mutex::new(HashMap::new())); let list_peers_calls = Arc::new(AtomicU32::new(0)); @@ -4417,7 +4460,7 @@ mod tests { assert_eq!(get_peer_identity_type_calls.load(Ordering::Relaxed), 2); *peers.lock() = vec![2, 4]; - service_impl.handle_global_ctx_event(&GlobalCtxEvent::PeerConnRemoved(Default::default())); + service_impl.handle_peer_context_event(&PeerContextEvent::PeerConnRemoved); assert!(service_impl.update_my_conn_info().await); assert_eq!(list_peers_calls.load(Ordering::Relaxed), 2); @@ -4430,7 +4473,7 @@ mod tests { #[tokio::test] async fn get_peer_identity_type_reuses_snapshot_until_topology_changes() { - let service_impl = PeerRouteServiceImpl::new(1, get_mock_global_ctx()); + let service_impl = test_service_impl(1); let peers = Arc::new(Mutex::new(vec![2, 3])); let peer_identity_types = Arc::new(Mutex::new(HashMap::from([ (2, Some(PeerIdentityType::Credential)), @@ -4462,7 +4505,7 @@ mod tests { assert_eq!(get_peer_identity_type_calls.load(Ordering::Relaxed), 2); *peers.lock() = vec![2, 4]; - service_impl.handle_global_ctx_event(&GlobalCtxEvent::PeerConnRemoved(Default::default())); + service_impl.handle_peer_context_event(&PeerContextEvent::PeerConnRemoved); assert_eq!( service_impl.get_peer_identity_type_from_interface(4).await, @@ -4479,2645 +4522,95 @@ mod tests { assert_eq!(get_peer_identity_type_calls.load(Ordering::Relaxed), 4); } - async fn create_mock_route(peer_mgr: Arc) -> Arc { - let peer_route = PeerRoute::new( - peer_mgr.my_peer_id(), - peer_mgr.get_global_ctx(), - peer_mgr.get_peer_rpc_mgr(), - ); - peer_mgr.add_route(peer_route.clone()).await; - peer_route - } - - fn get_rpc_counter(route: &Arc, peer_id: PeerId) -> (u32, u32) { - let session = route.service_impl.get_session(peer_id).unwrap(); - ( - session.rpc_tx_count.load(Ordering::Relaxed), - session.rpc_rx_count.load(Ordering::Relaxed), - ) - } - - fn get_is_initiator(route: &Arc, peer_id: PeerId) -> (bool, bool) { - let session = route.service_impl.get_session(peer_id).unwrap(); - ( - session.we_are_initiator.load(Ordering::Relaxed), - session.dst_is_initiator.load(Ordering::Relaxed), - ) - } - - fn make_credential_route_peer_info( - peer_id: PeerId, - noise_static_pubkey: &[u8], - ) -> RoutePeerInfo { - let mut peer_info = RoutePeerInfo::new(); - peer_info.peer_id = peer_id; - peer_info.version = 1; - peer_info.noise_static_pubkey = noise_static_pubkey.to_vec(); - peer_info.feature_flag = Some(PeerFeatureFlag { - is_credential_peer: true, - ..Default::default() - }); - peer_info - } - - fn make_admin_route_peer_info( - peer_id: PeerId, - credential_key: &[u8], - network_secret: &str, - now: i64, - ) -> RoutePeerInfo { - let mut admin_info = RoutePeerInfo::new(); - admin_info.peer_id = peer_id; - admin_info.version = 1; - admin_info.feature_flag = Some(PeerFeatureFlag { - is_credential_peer: false, - ..Default::default() - }); - admin_info.trusted_credential_pubkeys = vec![TrustedCredentialPubkeyProof::new_signed( - TrustedCredentialPubkey { - pubkey: credential_key.to_vec(), - expiry_unix: now + 600, - reusable: Some(false), - ..Default::default() - }, - network_secret, - )]; - admin_info - } - - fn make_route_conn_info(connected_peers: I, last_update: SystemTime) -> RouteConnInfo - where - I: IntoIterator, - { - RouteConnInfo { - connected_peers: connected_peers.into_iter().collect(), - version: 1.into(), - last_update, - } - } - - async fn create_mock_pmgr() -> Arc { - let (s, _r) = create_packet_recv_chan(); - let peer_mgr = Arc::new(PeerManager::new( - RouteAlgoType::None, - get_mock_global_ctx(), - s, - )); - replace_stun_info_collector(peer_mgr.clone(), NatType::Unknown); - peer_mgr.run().await.unwrap(); - peer_mgr - } - - fn check_rpc_counter(route: &Arc, peer_id: PeerId, max_tx: u32, max_rx: u32) { - let (tx1, rx1) = get_rpc_counter(route, peer_id); - assert!(tx1 <= max_tx); - assert!(rx1 <= max_rx); - } - #[tokio::test] - async fn credential_flag_controls_role_classification() { - let service_impl = PeerRouteServiceImpl::new(1, get_mock_global_ctx()); - - let mut admin_info = RoutePeerInfo::new(); - admin_info.peer_id = 10; - admin_info.version = 1; - admin_info.feature_flag = Some(PeerFeatureFlag { - is_credential_peer: false, - ..Default::default() - }); - - let mut credential_info = RoutePeerInfo::new(); - credential_info.peer_id = 11; - credential_info.version = 1; - credential_info.feature_flag = Some(PeerFeatureFlag { - is_credential_peer: true, - ..Default::default() - }); - - { - let mut guard = service_impl.synced_route_info.peer_infos.write(); - guard.insert(admin_info.peer_id, admin_info.clone()); - guard.insert(credential_info.peer_id, credential_info.clone()); - } - - assert!(service_impl.synced_route_info.is_admin_peer(&admin_info)); - assert!( - !service_impl - .synced_route_info - .is_admin_peer(&credential_info) - ); - assert!( - service_impl - .synced_route_info - .is_credential_peer(credential_info.peer_id) - ); - assert!( - !service_impl - .synced_route_info - .is_credential_peer(admin_info.peer_id) - ); - } - - #[tokio::test] - async fn trusted_credentials_only_from_admin_publishers() { - let service_impl = PeerRouteServiceImpl::new(1, get_mock_global_ctx()); - let network_secret = "sec1"; - let now = std::time::SystemTime::now() - .duration_since(std::time::UNIX_EPOCH) - .unwrap() - .as_secs() as i64; - - let admin_key = vec![1; 32]; - let credential_key = vec![2; 32]; - - let mut admin_info = RoutePeerInfo::new(); - admin_info.peer_id = 20; - admin_info.version = 1; - admin_info.feature_flag = Some(PeerFeatureFlag { - is_credential_peer: false, - ..Default::default() - }); - admin_info.trusted_credential_pubkeys = vec![TrustedCredentialPubkeyProof::new_signed( - TrustedCredentialPubkey { - pubkey: admin_key.clone(), - expiry_unix: now + 600, - ..Default::default() - }, - network_secret, - )]; - - let mut credential_info = RoutePeerInfo::new(); - credential_info.peer_id = 21; - credential_info.version = 1; - credential_info.feature_flag = Some(PeerFeatureFlag { - is_credential_peer: true, - ..Default::default() - }); - credential_info.trusted_credential_pubkeys = - vec![TrustedCredentialPubkeyProof::new_signed( - TrustedCredentialPubkey { - pubkey: credential_key.clone(), - expiry_unix: now + 600, - ..Default::default() - }, - network_secret, - )]; - - { - let mut guard = service_impl.synced_route_info.peer_infos.write(); - guard.insert(admin_info.peer_id, admin_info); - guard.insert(credential_info.peer_id, credential_info); - } - - service_impl - .synced_route_info - .verify_and_update_credential_trusts(Some(network_secret)); - - assert!( - service_impl - .synced_route_info - .trusted_credential_pubkeys - .contains_key(&admin_key) - ); - assert!( - !service_impl - .synced_route_info - .trusted_credential_pubkeys - .contains_key(&credential_key) - ); - } - - #[tokio::test] - async fn credential_groups_merge_with_proof_groups_and_recompute_cleanly() { - let service_impl = PeerRouteServiceImpl::new(1, get_mock_global_ctx()); - let network_secret = "sec1"; - let now = std::time::SystemTime::now() - .duration_since(std::time::UNIX_EPOCH) - .unwrap() - .as_secs() as i64; - let credential_peer_id = 31; - let credential_pubkey = vec![7; 32]; - - let mut credential_info = RoutePeerInfo::new(); - credential_info.peer_id = credential_peer_id; - credential_info.version = 1; - credential_info.noise_static_pubkey = credential_pubkey.clone(); - credential_info.groups = vec![PeerGroupInfo::generate_with_proof( - "proof-group".to_string(), - "proof-secret".to_string(), - credential_peer_id, - )]; - - let mut admin_info = RoutePeerInfo::new(); - admin_info.peer_id = 32; - admin_info.version = 1; - admin_info.feature_flag = Some(PeerFeatureFlag { - is_credential_peer: false, - ..Default::default() - }); - admin_info.trusted_credential_pubkeys = vec![TrustedCredentialPubkeyProof::new_signed( - TrustedCredentialPubkey { - pubkey: credential_pubkey.clone(), - groups: vec!["cred-group".to_string()], - expiry_unix: now + 600, - ..Default::default() - }, - network_secret, - )]; - - { - let mut guard = service_impl.synced_route_info.peer_infos.write(); - guard.insert(admin_info.peer_id, admin_info.clone()); - guard.insert(credential_peer_id, credential_info.clone()); - } - - service_impl - .synced_route_info - .verify_and_update_group_trusts( - &[credential_info], - &[GroupIdentity { - group_name: "proof-group".to_string(), - group_secret: "proof-secret".to_string(), - }], - false, - ); - service_impl - .synced_route_info - .verify_and_update_credential_trusts(Some(network_secret)); - - let groups = service_impl.get_peer_groups(credential_peer_id); - assert!(groups.contains(&"proof-group".to_string())); - assert!(groups.contains(&"cred-group".to_string())); - - let guard = service_impl.synced_route_info.peer_infos.write(); - let admin_info = guard.get(&32).unwrap().clone(); - drop(guard); - - let mut updated_admin = admin_info; - updated_admin.trusted_credential_pubkeys = vec![TrustedCredentialPubkeyProof::new_signed( - TrustedCredentialPubkey { - pubkey: credential_pubkey.clone(), - groups: vec!["replacement-group".to_string()], - expiry_unix: now + 600, - ..Default::default() - }, - network_secret, - )]; - service_impl - .synced_route_info - .peer_infos - .write() - .insert(updated_admin.peer_id, updated_admin); - - service_impl - .synced_route_info - .verify_and_update_credential_trusts(Some(network_secret)); - - let groups = service_impl.get_peer_groups(credential_peer_id); - assert!(groups.contains(&"proof-group".to_string())); - assert!(groups.contains(&"replacement-group".to_string())); - assert!(!groups.contains(&"cred-group".to_string())); - } - - #[tokio::test] - async fn remove_peers_batches_cleanup_and_version_increment() { - let service_impl = PeerRouteServiceImpl::new(1, get_mock_global_ctx()); - let removed_peer_ids = [41, 42]; - let retained_peer_id = 43; - - { - let mut peer_infos = service_impl.synced_route_info.peer_infos.write(); - let mut conn_map = service_impl.synced_route_info.conn_map.write(); - for peer_id in removed_peer_ids { - let mut info = RoutePeerInfo::new(); - info.peer_id = peer_id; - info.version = 1; - peer_infos.insert(peer_id, info); - conn_map.insert(peer_id, RouteConnInfo::default()); - } - - let mut retained_info = RoutePeerInfo::new(); - retained_info.peer_id = retained_peer_id; - retained_info.version = 1; - peer_infos.insert(retained_peer_id, retained_info); - conn_map.insert(retained_peer_id, RouteConnInfo::default()); - } - - for peer_id in removed_peer_ids { - service_impl.synced_route_info.raw_peer_infos.insert( - peer_id, - DynamicMessage::new(RoutePeerInfo::default().descriptor()), - ); - service_impl.synced_route_info.group_trust_map.insert( - peer_id, - HashMap::from([("guest".to_string(), vec![1, 2, 3])]), - ); - service_impl - .synced_route_info - .group_trust_map_cache - .insert(peer_id, Arc::new(vec!["guest".to_string()])); - service_impl.synced_route_info.foreign_network.insert( - ForeignNetworkRouteInfoKey { - peer_id, - ..Default::default() - }, - ForeignNetworkRouteInfoEntry::default(), - ); - } - - service_impl.synced_route_info.foreign_network.insert( - ForeignNetworkRouteInfoKey { - peer_id: retained_peer_id, - ..Default::default() - }, - ForeignNetworkRouteInfoEntry::default(), - ); - - let initial_version = service_impl.synced_route_info.version.get(); - service_impl - .synced_route_info - .remove_peers(removed_peer_ids); - - assert_eq!( - service_impl.synced_route_info.version.get(), - initial_version + 1 - ); - for peer_id in removed_peer_ids { - assert!( - !service_impl - .synced_route_info - .peer_infos - .read() - .contains_key(&peer_id) - ); - assert!( - !service_impl - .synced_route_info - .conn_map - .read() - .contains_key(&peer_id) - ); - assert!( - !service_impl - .synced_route_info - .raw_peer_infos - .contains_key(&peer_id) - ); - assert!( - !service_impl - .synced_route_info - .group_trust_map - .contains_key(&peer_id) - ); - assert!( - !service_impl - .synced_route_info - .group_trust_map_cache - .contains_key(&peer_id) - ); - assert!( - !service_impl.synced_route_info.foreign_network.contains_key( - &ForeignNetworkRouteInfoKey { - peer_id, - ..Default::default() - } - ) - ); - } - - assert!( - service_impl - .synced_route_info - .peer_infos - .read() - .contains_key(&retained_peer_id) - ); - assert!(service_impl.synced_route_info.foreign_network.contains_key( - &ForeignNetworkRouteInfoKey { - peer_id: retained_peer_id, - ..Default::default() - } - )); - } - - #[tokio::test] - async fn verify_trusted_credential_hmac_with_raw_payload_bytes() { - let service_impl = PeerRouteServiceImpl::new(1, get_mock_global_ctx()); - let network_secret = "sec1"; - let now = std::time::SystemTime::now() - .duration_since(std::time::UNIX_EPOCH) - .unwrap() - .as_secs() as i64; - - let credential_key = vec![7; 32]; - - let mut admin_info = RoutePeerInfo::new(); - admin_info.peer_id = 30; - admin_info.version = 1; - - let credential = TrustedCredentialPubkey { - pubkey: credential_key.clone(), - expiry_unix: now + 600, - reusable: Some(true), - ..Default::default() - }; - let mut raw_credential_bytes = credential.encode_to_vec(); - prost::encoding::encode_key( - 9999, - prost::encoding::WireType::Varint, - &mut raw_credential_bytes, - ); - prost::encoding::encode_varint(42, &mut raw_credential_bytes); - - let (admin_info, raw_admin_info) = make_route_info_with_raw_trusted_credential_proof( - &admin_info, - &raw_credential_bytes, - &TrustedCredentialPubkeyProof::generate_credential_hmac_from_bytes( - &raw_credential_bytes, - network_secret, - ), - ); - assert_eq!(admin_info.trusted_credential_pubkeys.len(), 1); - assert!( - !admin_info.trusted_credential_pubkeys[0].verify_credential_hmac(network_secret), - "typed verification should fail after nested unknown fields are dropped" - ); - - let mut credential_info = RoutePeerInfo::new(); - credential_info.peer_id = 41; - credential_info.version = 1; - credential_info.noise_static_pubkey = credential_key.clone(); - credential_info.feature_flag = Some(PeerFeatureFlag { - is_credential_peer: true, - ..Default::default() - }); - - let mut raw_credential_info = DynamicMessage::new(RoutePeerInfo::default().descriptor()); - raw_credential_info - .transcode_from(&credential_info) - .unwrap(); - - { - let mut guard = service_impl.synced_route_info.peer_infos.write(); - guard.insert(admin_info.peer_id, admin_info); - guard.insert(credential_info.peer_id, credential_info); - } - service_impl - .synced_route_info - .raw_peer_infos - .insert(30, raw_admin_info); - service_impl - .synced_route_info - .raw_peer_infos - .insert(41, raw_credential_info); - - let (untrusted_peers, _) = service_impl - .synced_route_info - .verify_and_update_credential_trusts(Some(network_secret)); - assert!(untrusted_peers.is_empty()); - assert!( - service_impl - .synced_route_info - .trusted_credential_pubkeys - .contains_key(&credential_key) - ); - } - - #[tokio::test] - async fn non_reusable_credential_elects_lowest_peer_id() { - let service_impl = PeerRouteServiceImpl::new(1, get_mock_global_ctx()); - let network_secret = "sec1"; - let now = std::time::SystemTime::now() - .duration_since(std::time::UNIX_EPOCH) - .unwrap() - .as_secs() as i64; - - let credential_key = vec![7; 32]; - - let admin_info = make_admin_route_peer_info(30, &credential_key, network_secret, now); - - let mut original_peer = RoutePeerInfo::new(); - original_peer.peer_id = 41; - original_peer.version = 1; - original_peer.noise_static_pubkey = credential_key.clone(); - original_peer.feature_flag = Some(PeerFeatureFlag { - is_credential_peer: true, - ..Default::default() - }); - - { - let mut guard = service_impl.synced_route_info.peer_infos.write(); - guard.insert(admin_info.peer_id, admin_info.clone()); - guard.insert(original_peer.peer_id, original_peer); - } - - let (first_untrusted, _) = service_impl - .synced_route_info - .verify_and_update_credential_trusts(Some(network_secret)); - assert!(first_untrusted.is_empty()); - assert_eq!( - service_impl - .synced_route_info - .non_reusable_credential_owners - .get(&credential_key) - .map(|entry| *entry.value()), - Some(41) - ); - - let mut new_peer = RoutePeerInfo::new(); - new_peer.peer_id = 39; - new_peer.version = 1; - new_peer.noise_static_pubkey = credential_key.clone(); - new_peer.feature_flag = Some(PeerFeatureFlag { - is_credential_peer: true, - ..Default::default() - }); - service_impl - .synced_route_info - .peer_infos - .write() - .insert(new_peer.peer_id, new_peer); - service_impl - .synced_route_info - .non_reusable_credential_owners - .insert(credential_key.clone(), 41); - - let (second_untrusted, _) = service_impl - .synced_route_info - .verify_and_update_credential_trusts(Some(network_secret)); - assert!(second_untrusted.is_empty()); - assert!( - service_impl - .synced_route_info - .peer_infos - .read() - .contains_key(&41) - ); - assert!( - service_impl - .synced_route_info - .peer_infos - .read() - .contains_key(&39) - ); - assert_eq!( - service_impl - .synced_route_info - .non_reusable_credential_owners - .get(&credential_key) - .map(|entry| *entry.value()), - Some(39) - ); - assert!(service_impl.synced_route_info.is_route_suppressed(41)); - assert!(!service_impl.synced_route_info.is_route_suppressed(39)); - } - - #[tokio::test] - async fn non_reusable_credential_ignores_unreachable_stale_owner() { - let service_impl = PeerRouteServiceImpl::new(1, get_mock_global_ctx()); - let network_secret = "sec1"; - let now = std::time::SystemTime::now() - .duration_since(std::time::UNIX_EPOCH) - .unwrap() - .as_secs() as i64; - - let credential_key = vec![8; 32]; - let stale_peer_id = 41; - let replacement_peer_id = 39; - - let admin_info = make_admin_route_peer_info(30, &credential_key, network_secret, now); - - let mut stale_peer = RoutePeerInfo::new(); - stale_peer.peer_id = stale_peer_id; - stale_peer.version = 1; - stale_peer.noise_static_pubkey = credential_key.clone(); - stale_peer.feature_flag = Some(PeerFeatureFlag { - is_credential_peer: true, - ..Default::default() - }); - - let mut replacement_peer = RoutePeerInfo::new(); - replacement_peer.peer_id = replacement_peer_id; - replacement_peer.version = 1; - replacement_peer.noise_static_pubkey = credential_key.clone(); - replacement_peer.feature_flag = Some(PeerFeatureFlag { - is_credential_peer: true, - ..Default::default() - }); - - { - let mut guard = service_impl.synced_route_info.peer_infos.write(); - guard.insert(admin_info.peer_id, admin_info); - guard.insert(stale_peer.peer_id, stale_peer); - guard.insert(replacement_peer.peer_id, replacement_peer); - } - service_impl - .synced_route_info - .non_reusable_credential_owners - .insert(credential_key.clone(), stale_peer_id); - - service_impl.route_table.next_hop_map.insert( - replacement_peer_id, - NextHopInfo { - next_hop_peer_id: replacement_peer_id, - path_latency: 0, - path_len: 1, - version: 1, - }, - ); - service_impl.route_table.next_hop_map_version.set(1); - - let (untrusted_peers, _) = service_impl - .synced_route_info - .verify_and_update_credential_trusts_with_active_peers( - Some(network_secret), - |peer_id| service_impl.is_active_non_reusable_credential_peer(peer_id), - ); - assert!(untrusted_peers.is_empty()); - assert!( - service_impl - .synced_route_info - .peer_infos - .read() - .contains_key(&stale_peer_id) - ); - assert!( - service_impl - .synced_route_info - .peer_infos - .read() - .contains_key(&replacement_peer_id) - ); - assert_eq!( - service_impl - .synced_route_info - .non_reusable_credential_owners - .get(&credential_key) - .map(|entry| *entry.value()), - Some(replacement_peer_id) - ); - assert!( - !service_impl - .synced_route_info - .is_route_suppressed(stale_peer_id) - ); - assert!( - !service_impl - .synced_route_info - .is_route_suppressed(replacement_peer_id) - ); - } - - #[tokio::test] - async fn suppressed_non_reusable_credential_peer_stays_synced_and_can_be_reactivated() { - const NETWORK_SECRET: &str = "sec1"; - const SELF_PEER_ID: PeerId = 1; - const ADMIN_PEER_ID: PeerId = 30; - const FIRST_PEER_ID: PeerId = 39; - const SECOND_PEER_ID: PeerId = 41; - - let service_impl = PeerRouteServiceImpl::new( - SELF_PEER_ID, - get_mock_global_ctx_with_network(Some(NetworkIdentity::new( - "test-net".to_string(), - NETWORK_SECRET.to_string(), - ))), - ); - let now_unix = SystemTime::now() - .duration_since(std::time::UNIX_EPOCH) - .unwrap() - .as_secs() as i64; - let now = SystemTime::now(); - let credential_key = vec![10; 32]; - - let mut self_info = RoutePeerInfo::new(); - self_info.peer_id = SELF_PEER_ID; - self_info.version = 1; - - let admin_info = - make_admin_route_peer_info(ADMIN_PEER_ID, &credential_key, NETWORK_SECRET, now_unix); - let mut first_peer = make_credential_route_peer_info(FIRST_PEER_ID, &credential_key); - first_peer.ipv4_addr = Some(std::net::Ipv4Addr::new(10, 144, 0, 39).into()); - let mut second_peer = make_credential_route_peer_info(SECOND_PEER_ID, &credential_key); - second_peer.ipv4_addr = Some(std::net::Ipv4Addr::new(10, 144, 0, 41).into()); - second_peer.proxy_cidrs.push("10.244.41.0/24".into()); - - { - let mut peer_infos = service_impl.synced_route_info.peer_infos.write(); - peer_infos.insert(self_info.peer_id, self_info); - peer_infos.insert(admin_info.peer_id, admin_info); - peer_infos.insert(first_peer.peer_id, first_peer); - peer_infos.insert(second_peer.peer_id, second_peer); - } - { - let mut conn_map = service_impl.synced_route_info.conn_map.write(); - conn_map.insert(SELF_PEER_ID, make_route_conn_info([ADMIN_PEER_ID], now)); - conn_map.insert( - ADMIN_PEER_ID, - make_route_conn_info([SELF_PEER_ID, FIRST_PEER_ID, SECOND_PEER_ID], now), - ); - conn_map.insert(FIRST_PEER_ID, make_route_conn_info([ADMIN_PEER_ID], now)); - conn_map.insert(SECOND_PEER_ID, make_route_conn_info([ADMIN_PEER_ID], now)); - } - service_impl.synced_route_info.version.set(1); - - let first_untrusted = service_impl.refresh_credential_trusts_with_current_topology(); - assert!(first_untrusted.is_empty()); - assert_eq!( - service_impl - .synced_route_info - .non_reusable_credential_owners - .get(&credential_key) - .map(|entry| *entry.value()), - Some(FIRST_PEER_ID) - ); - assert!( - service_impl - .synced_route_info - .peer_infos - .read() - .contains_key(&SECOND_PEER_ID) - ); - assert!( - service_impl - .synced_route_info - .is_route_suppressed(SECOND_PEER_ID) - ); - assert!( - service_impl - .route_table - .topology_peer_reachable(SECOND_PEER_ID) - ); - assert!(service_impl.route_table.peer_reachable(FIRST_PEER_ID)); - assert!(!service_impl.route_table.peer_reachable(SECOND_PEER_ID)); - assert!( - service_impl - .route_table - .peer_infos - .contains_key(&FIRST_PEER_ID) - ); - assert!( - !service_impl - .route_table - .peer_infos - .contains_key(&SECOND_PEER_ID) - ); - assert_eq!( - service_impl - .route_table - .ipv4_peer_id_map - .get(&"10.144.0.41".parse().unwrap()) - .map(|entry| entry.peer_id), - None - ); - assert_eq!( - service_impl - .route_table - .get_peer_id_for_proxy(&"10.244.41.1".parse().unwrap()), - None - ); - let sync_session = SyncRouteSession::new(SELF_PEER_ID, ADMIN_PEER_ID); - let sync_peer_ids: BTreeSet<_> = service_impl - .build_route_info(&sync_session) - .unwrap() - .into_iter() - .map(|info| info.peer_id) - .collect(); - assert!(sync_peer_ids.contains(&SECOND_PEER_ID)); - - { - let mut conn_map = service_impl.synced_route_info.conn_map.write(); - conn_map.insert( - ADMIN_PEER_ID, - make_route_conn_info([SELF_PEER_ID, SECOND_PEER_ID], now), - ); - conn_map.insert(FIRST_PEER_ID, make_route_conn_info([], now)); - conn_map.insert(SECOND_PEER_ID, make_route_conn_info([ADMIN_PEER_ID], now)); - } - service_impl.synced_route_info.version.inc(); - - let second_untrusted = service_impl.refresh_credential_trusts_with_current_topology(); - assert!(second_untrusted.is_empty()); - assert_eq!( - service_impl - .synced_route_info - .non_reusable_credential_owners - .get(&credential_key) - .map(|entry| *entry.value()), - Some(SECOND_PEER_ID) - ); - assert!( - !service_impl - .synced_route_info - .is_route_suppressed(SECOND_PEER_ID) - ); - assert!( - service_impl - .route_table - .topology_peer_reachable(SECOND_PEER_ID) - ); - assert!(!service_impl.route_table.peer_reachable(FIRST_PEER_ID)); - assert!(service_impl.route_table.peer_reachable(SECOND_PEER_ID)); - assert!( - !service_impl - .route_table - .peer_infos - .contains_key(&FIRST_PEER_ID) - ); - assert!( - service_impl - .route_table - .peer_infos - .contains_key(&SECOND_PEER_ID) - ); - assert_eq!( - service_impl - .route_table - .ipv4_peer_id_map - .get(&"10.144.0.41".parse().unwrap()) - .map(|entry| entry.peer_id), - Some(SECOND_PEER_ID) - ); - assert_eq!( - service_impl - .route_table - .get_peer_id_for_proxy(&"10.244.41.1".parse().unwrap()), - Some(SECOND_PEER_ID) - ); - } - - #[tokio::test] - async fn suppressed_non_reusable_credential_peer_is_not_transit_next_hop() { - const NETWORK_SECRET: &str = "sec1"; - const SELF_PEER_ID: PeerId = 1; - const ADMIN_PEER_ID: PeerId = 30; - const FIRST_PEER_ID: PeerId = 39; - const SECOND_PEER_ID: PeerId = 41; - const DOWNSTREAM_PEER_ID: PeerId = 50; - - let service_impl = PeerRouteServiceImpl::new( - SELF_PEER_ID, - get_mock_global_ctx_with_network(Some(NetworkIdentity::new( - "test-net".to_string(), - NETWORK_SECRET.to_string(), - ))), - ); - let now_unix = SystemTime::now() - .duration_since(std::time::UNIX_EPOCH) - .unwrap() - .as_secs() as i64; - let now = SystemTime::now(); - let credential_key = vec![10; 32]; - - let mut self_info = RoutePeerInfo::new(); - self_info.peer_id = SELF_PEER_ID; - self_info.version = 1; - - let admin_info = - make_admin_route_peer_info(ADMIN_PEER_ID, &credential_key, NETWORK_SECRET, now_unix); - let first_peer = make_credential_route_peer_info(FIRST_PEER_ID, &credential_key); - let second_peer = make_credential_route_peer_info(SECOND_PEER_ID, &credential_key); - let mut downstream_peer = RoutePeerInfo::new(); - downstream_peer.peer_id = DOWNSTREAM_PEER_ID; - downstream_peer.version = 1; - - { - let mut peer_infos = service_impl.synced_route_info.peer_infos.write(); - peer_infos.insert(self_info.peer_id, self_info); - peer_infos.insert(admin_info.peer_id, admin_info); - peer_infos.insert(first_peer.peer_id, first_peer); - peer_infos.insert(second_peer.peer_id, second_peer); - peer_infos.insert(downstream_peer.peer_id, downstream_peer); - } - { - let mut conn_map = service_impl.synced_route_info.conn_map.write(); - conn_map.insert(SELF_PEER_ID, make_route_conn_info([ADMIN_PEER_ID], now)); - conn_map.insert( - ADMIN_PEER_ID, - make_route_conn_info([SELF_PEER_ID, FIRST_PEER_ID, SECOND_PEER_ID], now), - ); - conn_map.insert(FIRST_PEER_ID, make_route_conn_info([ADMIN_PEER_ID], now)); - conn_map.insert( - SECOND_PEER_ID, - make_route_conn_info([ADMIN_PEER_ID, DOWNSTREAM_PEER_ID], now), - ); - conn_map.insert( - DOWNSTREAM_PEER_ID, - make_route_conn_info([SECOND_PEER_ID], now), - ); - } - service_impl.synced_route_info.version.set(1); - - let untrusted = service_impl.refresh_credential_trusts_with_current_topology(); - assert!(untrusted.is_empty()); - assert_eq!( - service_impl - .synced_route_info - .non_reusable_credential_owners - .get(&credential_key) - .map(|entry| *entry.value()), - Some(FIRST_PEER_ID) - ); - assert!( - service_impl - .synced_route_info - .is_route_suppressed(SECOND_PEER_ID) - ); - assert!( - service_impl - .route_table - .topology_peer_reachable(SECOND_PEER_ID) - ); - assert!(!service_impl.route_table.peer_reachable(SECOND_PEER_ID)); - assert!(!service_impl.route_table.peer_reachable(DOWNSTREAM_PEER_ID)); - assert!( - service_impl - .route_table - .get_next_hop(DOWNSTREAM_PEER_ID) - .is_none() - ); - assert!( - !service_impl - .route_table - .peer_infos - .contains_key(&DOWNSTREAM_PEER_ID) - ); - } - - #[tokio::test] - async fn credential_trust_refresh_does_not_remove_self_peer() { - let my_peer_id = 11; - let remote_peer_id = 12; - let credential_key = vec![8; 32]; - let service_impl = PeerRouteServiceImpl::new(my_peer_id, get_mock_global_ctx()); - - let self_info = make_credential_route_peer_info(my_peer_id, &credential_key); - let remote_info = make_credential_route_peer_info(remote_peer_id, &credential_key); - - { - let mut guard = service_impl.synced_route_info.peer_infos.write(); - guard.insert(self_info.peer_id, self_info); - guard.insert(remote_info.peer_id, remote_info); - } - service_impl - .synced_route_info - .trusted_credential_pubkeys - .insert( - credential_key.clone(), - TrustedCredentialPubkey { - pubkey: credential_key, - expiry_unix: i64::MAX, - ..Default::default() - }, - ); - - let (untrusted_peers, _, _) = service_impl - .synced_route_info - .verify_and_update_credential_trusts_with_active_peers_protecting( - None, - |_| true, - Some(my_peer_id), - ); - - assert_eq!(untrusted_peers, vec![remote_peer_id]); - assert!( - service_impl - .synced_route_info - .peer_infos - .read() - .contains_key(&my_peer_id) - ); - assert!( - !service_impl - .synced_route_info - .peer_infos - .read() - .contains_key(&remote_peer_id) - ); - } - - #[tokio::test] - async fn credential_refresh_rebuilds_reachability_before_owner_election() { - const NETWORK_SECRET: &str = "sec1"; - const SELF_PEER_ID: PeerId = 1; - - let service_impl = PeerRouteServiceImpl::new( - SELF_PEER_ID, - get_mock_global_ctx_with_network(Some(NetworkIdentity::new( - "test-net".to_string(), - NETWORK_SECRET.to_string(), - ))), - ); - let now = std::time::SystemTime::now() - .duration_since(std::time::UNIX_EPOCH) - .unwrap() - .as_secs() as i64; - - let credential_key = vec![9; 32]; - let admin_peer_id = 30; - let stale_peer_id = 41; - let replacement_peer_id = 39; - - let mut self_info = RoutePeerInfo::new(); - self_info.peer_id = SELF_PEER_ID; - self_info.version = 1; - - let admin_info = - make_admin_route_peer_info(admin_peer_id, &credential_key, NETWORK_SECRET, now); - - let stale_peer = make_credential_route_peer_info(stale_peer_id, &credential_key); - let replacement_peer = - make_credential_route_peer_info(replacement_peer_id, &credential_key); - - { - let mut guard = service_impl.synced_route_info.peer_infos.write(); - guard.insert(self_info.peer_id, self_info); - guard.insert(admin_info.peer_id, admin_info); - guard.insert(stale_peer.peer_id, stale_peer); - guard.insert(replacement_peer.peer_id, replacement_peer); - } - - let now = std::time::SystemTime::now(); - { - let mut guard = service_impl.synced_route_info.conn_map.write(); - guard.insert(SELF_PEER_ID, make_route_conn_info([admin_peer_id], now)); - guard.insert( - admin_peer_id, - make_route_conn_info([SELF_PEER_ID, replacement_peer_id], now), - ); - guard.insert( - replacement_peer_id, - make_route_conn_info([admin_peer_id], now), - ); - guard.insert(stale_peer_id, make_route_conn_info([], now)); - } - service_impl.synced_route_info.version.set(2); - - service_impl.update_route_table_and_cached_local_conn_bitmap(); - assert!(!service_impl.is_active_non_reusable_credential_peer(stale_peer_id)); - assert!(service_impl.is_active_non_reusable_credential_peer(replacement_peer_id)); - - service_impl.route_table.next_hop_map.clear(); - service_impl.route_table.next_hop_map.insert( - stale_peer_id, - NextHopInfo { - next_hop_peer_id: stale_peer_id, - path_latency: 0, - path_len: 1, - version: 1, - }, - ); - service_impl.route_table.next_hop_map_version.set(1); - - let untrusted = service_impl.refresh_credential_trusts_with_current_topology(); - assert!(untrusted.is_empty()); - assert!(!service_impl.is_active_non_reusable_credential_peer(stale_peer_id)); - assert!(service_impl.is_active_non_reusable_credential_peer(replacement_peer_id)); - assert_eq!( - service_impl - .synced_route_info - .non_reusable_credential_owners - .get(&credential_key) - .map(|entry| *entry.value()), - Some(replacement_peer_id) - ); - } - - #[tokio::test] - async fn update_my_infos_refreshes_non_reusable_owner_on_conn_change() { - const NETWORK_SECRET: &str = "sec1"; - const ADMIN_PEER_ID: PeerId = 30; - const FIRST_PEER_ID: PeerId = 39; - const SECOND_PEER_ID: PeerId = 41; - - let global_ctx = get_mock_global_ctx_with_network(Some(NetworkIdentity::new( - "test-net".to_string(), - NETWORK_SECRET.to_string(), - ))); - let (_credential_id, credential_secret) = global_ctx - .get_credential_manager() - .generate_credential_with_options( - vec![], - false, - vec![], - Duration::from_secs(3600), - None, - false, - ); - let credential_secret_bytes: [u8; 32] = BASE64_STANDARD - .decode(&credential_secret) - .unwrap() - .try_into() - .unwrap(); - let credential_secret = x25519_dalek::StaticSecret::from(credential_secret_bytes); - let credential_key = x25519_dalek::PublicKey::from(&credential_secret) - .as_bytes() - .to_vec(); - - let service_impl = PeerRouteServiceImpl::new(ADMIN_PEER_ID, global_ctx); - let peers = Arc::new(Mutex::new(vec![FIRST_PEER_ID, SECOND_PEER_ID])); - let peer_identity_types = Arc::new(Mutex::new(HashMap::from([ - (FIRST_PEER_ID, Some(PeerIdentityType::Credential)), - (SECOND_PEER_ID, Some(PeerIdentityType::Credential)), - ]))); - *service_impl.interface.lock().await = Some(Box::new(CountingInterface { - my_peer_id: ADMIN_PEER_ID, - peers: peers.clone(), - peer_identity_types, - list_peers_calls: Arc::new(AtomicU32::new(0)), - get_peer_identity_type_calls: Arc::new(AtomicU32::new(0)), - })); - - { - let mut peer_infos = service_impl.synced_route_info.peer_infos.write(); - peer_infos.insert( - FIRST_PEER_ID, - make_credential_route_peer_info(FIRST_PEER_ID, &credential_key), - ); - peer_infos.insert( - SECOND_PEER_ID, - make_credential_route_peer_info(SECOND_PEER_ID, &credential_key), - ); - } - let now = SystemTime::now(); - { - let mut conn_map = service_impl.synced_route_info.conn_map.write(); - conn_map.insert(FIRST_PEER_ID, make_route_conn_info([ADMIN_PEER_ID], now)); - conn_map.insert(SECOND_PEER_ID, make_route_conn_info([ADMIN_PEER_ID], now)); - } - - assert!(service_impl.update_my_infos().await); - assert_eq!( - service_impl - .synced_route_info - .non_reusable_credential_owners - .get(&credential_key) - .map(|entry| *entry.value()), - Some(FIRST_PEER_ID) - ); - assert!(service_impl.route_table.peer_reachable(FIRST_PEER_ID)); - assert!(!service_impl.route_table.peer_reachable(SECOND_PEER_ID)); - - *peers.lock() = vec![SECOND_PEER_ID]; - service_impl.handle_global_ctx_event(&GlobalCtxEvent::PeerConnRemoved(Default::default())); - - assert!(service_impl.update_my_infos().await); - assert_eq!( - service_impl - .synced_route_info - .non_reusable_credential_owners - .get(&credential_key) - .map(|entry| *entry.value()), - Some(SECOND_PEER_ID) - ); - assert!(!service_impl.route_table.peer_reachable(FIRST_PEER_ID)); - assert!(service_impl.route_table.peer_reachable(SECOND_PEER_ID)); - } - - #[tokio::test] - async fn sync_route_info_marks_credential_sender_and_filters_entries() { - let peer_mgr = create_mock_pmgr().await; - let route = create_mock_route(peer_mgr.clone()).await; - let from_peer_id: PeerId = 10001; - let forwarded_peer_id: PeerId = 10002; - let credential_pubkey = vec![3u8; 32]; - - let identity_type = DashMap::new(); - identity_type.insert(from_peer_id, PeerIdentityType::Credential); - let peer_public_key = DashMap::new(); - peer_public_key.insert(from_peer_id, credential_pubkey.clone()); - *route.service_impl.interface.lock().await = Some(Box::new(AuthOnlyInterface { - my_peer_id: peer_mgr.my_peer_id(), - identity_type, - peer_public_key, - })); - route - .service_impl - .synced_route_info - .trusted_credential_pubkeys - .insert( - credential_pubkey.clone(), - TrustedCredentialPubkey { - pubkey: credential_pubkey, - expiry_unix: i64::MAX, - ..Default::default() - }, - ); - - let mut sender_info = RoutePeerInfo::new(); - sender_info.peer_id = from_peer_id; - sender_info.version = 1; - sender_info.proxy_cidrs = vec!["10.10.0.0/24".to_string()]; - - let mut forwarded_info = RoutePeerInfo::new(); - forwarded_info.peer_id = forwarded_peer_id; - forwarded_info.version = 1; - - let make_raw = |info: &RoutePeerInfo| { - let mut raw = DynamicMessage::new(RoutePeerInfo::default().descriptor()); - raw.transcode_from(info).unwrap(); - raw - }; - let raw_infos = vec![make_raw(&sender_info), make_raw(&forwarded_info)]; - - route - .session_mgr - .do_sync_route_info( - from_peer_id, - 1, - true, - Some(vec![sender_info, forwarded_info]), - Some(raw_infos), - None, - None, - ) - .await - .unwrap(); - - let guard = route.service_impl.synced_route_info.peer_infos.read(); - let stored = guard.get(&from_peer_id).unwrap(); - assert!( - stored - .feature_flag - .as_ref() - .map(|x| x.is_credential_peer) - .unwrap_or(false) - ); - assert!(stored.proxy_cidrs.is_empty()); - assert!(guard.get(&forwarded_peer_id).is_none()); - } - - // shared node doesn't have hmac. - #[tokio::test] - async fn sync_route_info_shared_sender_cannot_publish_trusted_credentials() { - let peer_mgr = create_mock_pmgr().await; - let route = create_mock_route(peer_mgr.clone()).await; - let from_peer_id: PeerId = 10021; - let forwarded_peer_id: PeerId = 10022; - let credential_key = vec![9u8; 32]; - - let identity_type = DashMap::new(); - identity_type.insert(from_peer_id, PeerIdentityType::SharedNode); - *route.service_impl.interface.lock().await = Some(Box::new(AuthOnlyInterface { - my_peer_id: peer_mgr.my_peer_id(), - identity_type, - peer_public_key: DashMap::new(), - })); - - let mut sender_info = RoutePeerInfo::new(); - sender_info.peer_id = from_peer_id; - sender_info.version = 1; - - let mut forwarded_info = RoutePeerInfo::new(); - forwarded_info.peer_id = forwarded_peer_id; - forwarded_info.version = 1; - forwarded_info.trusted_credential_pubkeys = vec![TrustedCredentialPubkeyProof { - credential: Some(TrustedCredentialPubkey { - pubkey: credential_key.clone(), - expiry_unix: i64::MAX, - ..Default::default() - }), - credential_hmac: vec![1; 32], - }]; - - let make_raw = |info: &RoutePeerInfo| { - let mut raw = DynamicMessage::new(RoutePeerInfo::default().descriptor()); - raw.transcode_from(info).unwrap(); - raw - }; - let raw_infos = vec![make_raw(&sender_info), make_raw(&forwarded_info)]; - - route - .session_mgr - .do_sync_route_info( - from_peer_id, - 1, - true, - Some(vec![sender_info, forwarded_info]), - Some(raw_infos), - None, - None, - ) - .await - .unwrap(); - - assert!( - !route - .service_impl - .synced_route_info - .trusted_credential_pubkeys - .contains_key(&credential_key) - ); - } - - #[tokio::test] - async fn clear_expired_peer_recomputes_trust_after_last_admin_disappears() { - let service_impl = PeerRouteServiceImpl::new(1, get_mock_global_ctx()); - let admin_peer_id: PeerId = 10051; - let credential_peer_id: PeerId = 10052; - let admin_pubkey = vec![5u8; 32]; - let credential_pubkey = vec![6u8; 32]; - let network_name = service_impl - .global_ctx - .get_network_identity() - .network_name - .clone(); - let now = SystemTime::now(); - let closed_peers = Arc::new(Mutex::new(Vec::new())); - - *service_impl.interface.lock().await = Some(Box::new(TrackingInterface { - my_peer_id: service_impl.my_peer_id, - closed_peers: closed_peers.clone(), - })); - - { - let mut guard = service_impl.synced_route_info.peer_infos.write(); - - let mut admin_info = RoutePeerInfo::new(); - admin_info.peer_id = admin_peer_id; - admin_info.version = 1; - admin_info.last_update = - Some((now - REMOVE_DEAD_PEER_INFO_AFTER - Duration::from_secs(1)).into()); - admin_info.noise_static_pubkey = admin_pubkey; - admin_info.trusted_credential_pubkeys = vec![TrustedCredentialPubkeyProof { - credential: Some(TrustedCredentialPubkey { - pubkey: credential_pubkey.clone(), - groups: vec!["guest".to_string()], - expiry_unix: i64::MAX, - ..Default::default() - }), - credential_hmac: vec![1; 32], - }]; - - let mut credential_info = RoutePeerInfo::new(); - credential_info.peer_id = credential_peer_id; - credential_info.version = 1; - credential_info.last_update = Some(now.into()); - credential_info.noise_static_pubkey = credential_pubkey.clone(); - - guard.insert(admin_peer_id, admin_info); - guard.insert(credential_peer_id, credential_info); - } - - let (_, global_trusted_keys) = service_impl - .synced_route_info - .verify_and_update_credential_trusts(None); - service_impl - .global_ctx - .update_trusted_keys(global_trusted_keys, &network_name); - - assert!( - service_impl - .synced_route_info - .trusted_credential_pubkeys - .contains_key(&credential_pubkey) - ); - assert!( - service_impl - .get_peer_groups(credential_peer_id) - .contains(&"guest".to_string()) - ); - - service_impl.clear_expired_peer().await; - - assert!(!service_impl.global_ctx.is_pubkey_trusted_with_source( - &credential_pubkey, - &network_name, - TrustedKeySource::OspfCredential, - )); - assert!(closed_peers.lock().contains(&credential_peer_id)); - assert!( - !service_impl - .synced_route_info - .peer_infos - .read() - .contains_key(&admin_peer_id) - ); - assert!( - !service_impl - .synced_route_info - .peer_infos - .read() - .contains_key(&credential_peer_id) - ); - assert!( - !service_impl - .synced_route_info - .group_trust_map_cache - .contains_key(&credential_peer_id) - ); - } - - #[tokio::test] - async fn refresh_acl_groups_returns_true_when_untrusted_peers_are_disconnected() { - let service_impl = PeerRouteServiceImpl::new(1, get_mock_global_ctx()); - let credential_peer_id: PeerId = 10061; - let credential_pubkey = vec![8u8; 32]; - let closed_peers = Arc::new(Mutex::new(Vec::new())); - - *service_impl.interface.lock().await = Some(Box::new(TrackingInterface { - my_peer_id: service_impl.my_peer_id, - closed_peers: closed_peers.clone(), - })); - - let mut credential_info = RoutePeerInfo::new(); - credential_info.peer_id = credential_peer_id; - credential_info.version = 1; - credential_info.noise_static_pubkey = credential_pubkey.clone(); - credential_info.feature_flag = Some(PeerFeatureFlag { - is_credential_peer: true, - ..Default::default() - }); - let self_info = RoutePeerInfo::new_updated_self( - service_impl.my_peer_id, - service_impl.my_peer_route_id, - &service_impl.global_ctx, - None, - ); - let mut self_info = self_info; - self_info.version = 1; - self_info.last_update = Some(Timestamp::now()); - { - let mut guard = service_impl.synced_route_info.peer_infos.write(); - guard.insert(service_impl.my_peer_id, self_info); - guard.insert(credential_peer_id, credential_info); - } - service_impl - .synced_route_info - .trusted_credential_pubkeys - .insert( - credential_pubkey.clone(), - TrustedCredentialPubkey { - pubkey: credential_pubkey.clone(), - expiry_unix: i64::MAX, - ..Default::default() - }, - ); - - assert!(service_impl.refresh_acl_groups().await); - assert!(closed_peers.lock().contains(&credential_peer_id)); - assert!( - !service_impl - .synced_route_info - .peer_infos - .read() - .contains_key(&credential_peer_id) - ); - assert!( - !service_impl - .synced_route_info - .trusted_credential_pubkeys - .contains_key(&credential_pubkey) - ); - } - - #[tokio::test] - async fn refresh_acl_groups_updates_local_membership_immediately() { - let peer_mgr = create_mock_pmgr().await; - let route = create_mock_route(peer_mgr.clone()).await; - let my_peer_id = peer_mgr.my_peer_id(); - - assert!(route.service_impl.get_peer_groups(my_peer_id).is_empty()); - - peer_mgr.get_global_ctx().config.set_acl(Some(Acl { - acl_v1: Some(AclV1 { - group: Some(GroupInfo { - declares: vec![GroupIdentity { - group_name: "admin".to_string(), - group_secret: "admin-secret".to_string(), - }], - members: vec!["admin".to_string()], - }), - ..Default::default() - }), - })); - - route.refresh_acl_groups().await; - - let groups = route.service_impl.get_peer_groups(my_peer_id); - assert!(groups.contains(&"admin".to_string())); - assert_eq!(groups.len(), 1); - } - - #[tokio::test] - async fn refresh_acl_groups_revalidates_cached_remote_groups() { - let peer_mgr = create_mock_pmgr().await; - let route = create_mock_route(peer_mgr.clone()).await; - let remote_peer_id = 200; - let remote_group = PeerGroupInfo::generate_with_proof( - "ops".to_string(), - "secret-v1".to_string(), - remote_peer_id, - ); - - peer_mgr.get_global_ctx().config.set_acl(Some(Acl { - acl_v1: Some(AclV1 { - group: Some(GroupInfo { - declares: vec![GroupIdentity { - group_name: "ops".to_string(), - group_secret: "secret-v1".to_string(), - }], - members: vec![], - }), - ..Default::default() - }), - })); - - let mut remote_info = RoutePeerInfo::new(); - remote_info.peer_id = remote_peer_id; - remote_info.version = 1; - remote_info.groups = vec![remote_group]; - route - .service_impl - .synced_route_info - .peer_infos - .write() - .insert(remote_peer_id, remote_info.clone()); - route - .service_impl - .synced_route_info - .verify_and_update_group_trusts( - &[remote_info], - &peer_mgr.get_global_ctx().get_acl_group_declarations(), - false, - ); - - assert!( - route - .service_impl - .get_peer_groups(remote_peer_id) - .contains(&"ops".to_string()) - ); - - peer_mgr.get_global_ctx().config.set_acl(Some(Acl { - acl_v1: Some(AclV1 { - group: Some(GroupInfo { - declares: vec![GroupIdentity { - group_name: "ops".to_string(), - group_secret: "secret-v2".to_string(), - }], - members: vec![], - }), - ..Default::default() - }), - })); - - route.refresh_acl_groups().await; - - assert!( - route - .service_impl - .get_peer_groups(remote_peer_id) - .is_empty() - ); - } - - #[tokio::test] - async fn credential_verifier_trusts_admin_self_groups_from_multiple_admins() { - let service_impl = PeerRouteServiceImpl::new( + async fn stop_waits_for_in_flight_route_sync_before_draining_sessions() { + let peer_rpc = Arc::new(PeerRpcManager::new(TestPeerRpcTransport)); + let route = PeerRoute::new( 1, - get_mock_global_ctx_with_network(Some( - crate::common::config::NetworkIdentity::new_credential("net1".to_string()), - )), + Arc::new(NoopPeerContext::default()), + Arc::new(TestPublicIpv6Runtime), + peer_rpc, ); + let entered = Arc::new(Notify::new()); + let release = Arc::new(Notify::new()); + *route.service_impl.interface.lock().await = Some(Box::new(BlockingInterface { + entered: entered.clone(), + release: release.clone(), + })); - let mut admin_a = RoutePeerInfo::new(); - admin_a.peer_id = 501; - admin_a.version = 1; - admin_a.groups = vec![ - PeerGroupInfo { - group_name: "ops".to_string(), - group_proof: vec![1; 32], - }, - PeerGroupInfo { - group_name: "core-admin".to_string(), - group_proof: vec![2; 32], - }, - ]; + let sync_task = tokio::spawn({ + let session_mgr = route.session_mgr.clone(); + async move { + session_mgr + .do_sync_route_info(2, 1, false, None, None, None, None) + .await + } + }); + crate::foundation::time::timeout(Duration::from_secs(1), entered.notified()) + .await + .expect("route sync did not enter the interface call"); - let mut admin_b = RoutePeerInfo::new(); - admin_b.peer_id = 502; - admin_b.version = 1; - admin_b.groups = vec![PeerGroupInfo { - group_name: "audit".to_string(), - group_proof: vec![3; 32], - }]; + let stop_task = tokio::spawn({ + let route = route.clone(); + async move { route.stop().await } + }); + crate::foundation::time::timeout(Duration::from_secs(1), async { + while !route.service_impl.stopped.load(Ordering::Acquire) { + tokio::task::yield_now().await; + } + }) + .await + .expect("route did not enter the stopped state"); - service_impl - .synced_route_info - .verify_and_update_group_trusts(&[admin_a.clone(), admin_b.clone()], &[], true); + assert!(!stop_task.is_finished()); + assert!(matches!( + route.session_mgr.get_or_start_session(3), + Err(Error::Stopped) + )); - let admin_a_groups = service_impl.get_peer_groups(admin_a.peer_id); - assert!(admin_a_groups.contains(&"ops".to_string())); - assert!(admin_a_groups.contains(&"core-admin".to_string())); + release.notify_one(); + sync_task + .await + .expect("route sync task panicked") + .expect("in-flight route sync failed"); + crate::foundation::time::timeout(Duration::from_secs(1), stop_task) + .await + .expect("route stop did not finish") + .expect("route stop task panicked"); - let admin_b_groups = service_impl.get_peer_groups(admin_b.peer_id); - assert!(admin_b_groups.contains(&"audit".to_string())); + assert!(route.service_impl.sessions.is_empty()); + assert_eq!(route.task_count(), 0); } - #[tokio::test] - async fn credential_verifier_still_checks_credential_self_declared_groups() { - let service_impl = PeerRouteServiceImpl::new( + #[test] + fn builds_next_hop_and_proxy_lookup_from_snapshot() { + let mut remote_proxy_peer = peer(3); + remote_proxy_peer + .info + .proxy_cidrs + .push("10.10.0.0/16".into()); + + let snapshot = OspfRouteSnapshot { + peer_infos: vec![peer(1), peer(2), remote_proxy_peer], + conn_map: vec![connected(1, [2]), connected(2, [1, 3]), connected(3, [2])], + suppressed_peer_ids: BTreeSet::new(), + version: 1, + }; + + let table = OspfRouteTable::new(); + table.build_from_snapshot( 1, - get_mock_global_ctx_with_network(Some( - crate::common::config::NetworkIdentity::new_credential("net1".to_string()), - )), + &snapshot, + NextHopPolicy::LeastHop, + &DefaultRouteCostCalculator, ); - let now = std::time::SystemTime::now() - .duration_since(std::time::UNIX_EPOCH) - .unwrap() - .as_secs() as i64; - let credential_peer_id = 601; - let credential_pubkey = vec![9; 32]; - - let mut admin_info = RoutePeerInfo::new(); - admin_info.peer_id = 600; - admin_info.version = 1; - admin_info.trusted_credential_pubkeys = vec![TrustedCredentialPubkeyProof { - credential: Some(TrustedCredentialPubkey { - pubkey: credential_pubkey.clone(), - groups: vec!["cred-acl".to_string()], - expiry_unix: now + 600, - ..Default::default() - }), - credential_hmac: vec![7; 32], - }]; - - let mut credential_info = RoutePeerInfo::new(); - credential_info.peer_id = credential_peer_id; - credential_info.version = 1; - credential_info.noise_static_pubkey = credential_pubkey.clone(); - credential_info.feature_flag = Some(PeerFeatureFlag { - is_credential_peer: true, - ..Default::default() - }); - credential_info.groups = vec![ - PeerGroupInfo::generate_with_proof( - "proof-group".to_string(), - "proof-secret".to_string(), - credential_peer_id, - ), - PeerGroupInfo::generate_with_proof( - "invalid-group".to_string(), - "wrong-secret".to_string(), - credential_peer_id, - ), - ]; - - { - let mut guard = service_impl.synced_route_info.peer_infos.write(); - guard.insert(admin_info.peer_id, admin_info.clone()); - guard.insert(credential_info.peer_id, credential_info.clone()); - } - - service_impl - .synced_route_info - .verify_and_update_group_trusts( - &[admin_info, credential_info], - &[ - GroupIdentity { - group_name: "proof-group".to_string(), - group_secret: "proof-secret".to_string(), - }, - GroupIdentity { - group_name: "invalid-group".to_string(), - group_secret: "actual-secret".to_string(), - }, - ], - true, - ); - service_impl - .synced_route_info - .verify_and_update_credential_trusts(None); - - let groups = service_impl.get_peer_groups(credential_peer_id); - assert!(groups.contains(&"proof-group".to_string())); - assert!(groups.contains(&"cred-acl".to_string())); - assert!(!groups.contains(&"invalid-group".to_string())); - } - - #[rstest::rstest] - #[tokio::test] - async fn ospf_route_2node(#[values(true, false)] enable_conn_list_sync: bool) { - FORCE_USE_CONN_LIST.store(enable_conn_list_sync, Ordering::Relaxed); - - let p_a = create_mock_pmgr().await; - let p_b = create_mock_pmgr().await; - connect_peer_manager(p_a.clone(), p_b.clone()).await; - - let r_a = create_mock_route(p_a.clone()).await; - let r_b = create_mock_route(p_b.clone()).await; - - for r in [r_a.clone(), r_b.clone()].iter() { - wait_for_condition( - || async { - println!("route: {:?}", r.list_routes().await); - r.list_routes().await.len() == 1 - }, - Duration::from_secs(5), - ) - .await; - } - - tokio::time::sleep(Duration::from_secs(3)).await; + let next_hop = table.get_next_hop(3).unwrap(); + assert_eq!(next_hop.next_hop_peer_id, 2); + assert_eq!(next_hop.path_len, 2); assert_eq!( - 2, - r_a.service_impl.synced_route_info.peer_infos.read().len() - ); - assert_eq!( - 2, - r_b.service_impl.synced_route_info.peer_infos.read().len() - ); - - for s in r_a.service_impl.sessions.iter() { - assert!(s.value().task.is_running()); - } - - assert_eq!( - r_a.service_impl - .synced_route_info - .peer_infos - .read() - .get(&p_a.my_peer_id()) - .unwrap() - .version, - r_a.service_impl - .get_session(p_b.my_peer_id()) - .unwrap() - .dst_saved_peer_info_versions - .get(&p_a.my_peer_id()) - .unwrap() - .value() - .get() - ); - - assert_eq!((1, 1), get_rpc_counter(&r_a, p_b.my_peer_id())); - assert_eq!((1, 1), get_rpc_counter(&r_b, p_a.my_peer_id())); - - let i_a = get_is_initiator(&r_a, p_b.my_peer_id()); - let i_b = get_is_initiator(&r_b, p_a.my_peer_id()); - assert_eq!(i_a.0, i_b.1); - assert_eq!(i_b.0, i_a.1); - - println!("after drop p_b, r_b"); - - drop(r_b); - drop(p_b); - - wait_for_condition( - || async { r_a.list_routes().await.is_empty() }, - Duration::from_secs(5), - ) - .await; - - wait_for_condition( - || async { r_a.service_impl.sessions.is_empty() }, - Duration::from_secs(5), - ) - .await; - } - - #[rstest::rstest] - #[tokio::test] - async fn ospf_route_multi_node(#[values(true, false)] enable_conn_list_sync: bool) { - FORCE_USE_CONN_LIST.store(enable_conn_list_sync, Ordering::Relaxed); - - let p_a = create_mock_pmgr().await; - let p_b = create_mock_pmgr().await; - let p_c = create_mock_pmgr().await; - connect_peer_manager(p_a.clone(), p_b.clone()).await; - connect_peer_manager(p_c.clone(), p_b.clone()).await; - - let r_a = create_mock_route(p_a.clone()).await; - let r_b = create_mock_route(p_b.clone()).await; - let r_c = create_mock_route(p_c.clone()).await; - - for r in [r_a.clone(), r_b.clone(), r_c.clone()].iter() { - wait_for_condition( - || async { r.service_impl.synced_route_info.peer_infos.read().len() == 3 }, - Duration::from_secs(5), - ) - .await; - } - - connect_peer_manager(p_a.clone(), p_c.clone()).await; - // for full-connected 3 nodes, the sessions between them may be a cycle or a line - wait_for_condition( - || async { - let mut lens = vec![ - r_a.service_impl.sessions.len(), - r_b.service_impl.sessions.len(), - r_c.service_impl.sessions.len(), - ]; - lens.sort(); - - lens == vec![1, 1, 2] || lens == vec![2, 2, 2] - }, - Duration::from_secs(3), - ) - .await; - - let p_d = create_mock_pmgr().await; - let r_d = create_mock_route(p_d.clone()).await; - connect_peer_manager(p_d.clone(), p_a.clone()).await; - connect_peer_manager(p_d.clone(), p_b.clone()).await; - connect_peer_manager(p_d.clone(), p_c.clone()).await; - - // find the smallest peer_id, which should be a center node - let mut all_route = [r_a.clone(), r_b.clone(), r_c.clone(), r_d.clone()]; - all_route.sort_by_key(|r| r.my_peer_id); - let mut all_peer_mgr = [p_a.clone(), p_b.clone(), p_c.clone(), p_d.clone()]; - all_peer_mgr.sort_by_key(|p| p.my_peer_id()); - - wait_for_condition( - || async { all_route[0].service_impl.sessions.len() == 3 }, - Duration::from_secs(3), - ) - .await; - - for r in all_route.iter() { - println!("session: {}", r.session_mgr.dump_sessions().unwrap()); - } - - let p_e = create_mock_pmgr().await; - let r_e = create_mock_route(p_e.clone()).await; - let last_p = all_peer_mgr.last().unwrap(); - connect_peer_manager(p_e.clone(), last_p.clone()).await; - - wait_for_condition( - || async { r_e.session_mgr.list_session_peers().len() == 1 }, - Duration::from_secs(3), - ) - .await; - - for s in r_e.service_impl.sessions.iter() { - assert!(s.value().task.is_running()); - } - - tokio::time::sleep(Duration::from_secs(2)).await; - - check_rpc_counter(&r_e, last_p.my_peer_id(), 2, 2); - - for r in all_route.iter() { - if r.my_peer_id != last_p.my_peer_id() { - wait_for_condition( - || async { - r.get_next_hop(p_e.my_peer_id()).await == Some(last_p.my_peer_id()) - }, - Duration::from_secs(3), - ) - .await; - } else { - wait_for_condition( - || async { r.get_next_hop(p_e.my_peer_id()).await == Some(p_e.my_peer_id()) }, - Duration::from_secs(3), - ) - .await; - } - } - } - - async fn check_route_sanity(p: &Arc, routable_peers: Vec>) { - let synced_info = &p.service_impl.synced_route_info; - for routable_peer in routable_peers.iter() { - // check conn map - let conns = { - let guard = synced_info.conn_map.read(); - guard.get(&routable_peer.my_peer_id()).cloned().unwrap() - }; - - assert_eq!( - conns.connected_peers, - routable_peer - .get_peer_map() - .list_peers() - .into_iter() - .collect::>() - ); - - // check peer infos - let peer_info = synced_info - .peer_infos - .read() - .get(&routable_peer.my_peer_id()) - .cloned() - .unwrap(); - assert_eq!(peer_info.peer_id, routable_peer.my_peer_id()); - } - } - - async fn print_routes(peers: Vec>) { - for p in peers.iter() { - println!("p:{:?}, route: {:#?}", p.my_peer_id, p.list_routes().await); - } - } - - #[rstest::rstest] - #[tokio::test] - async fn ospf_route_3node_disconnect(#[values(true, false)] enable_conn_list_sync: bool) { - FORCE_USE_CONN_LIST.store(enable_conn_list_sync, Ordering::Relaxed); - let p_a = create_mock_pmgr().await; - let p_b = create_mock_pmgr().await; - let p_c = create_mock_pmgr().await; - connect_peer_manager(p_a.clone(), p_b.clone()).await; - connect_peer_manager(p_c.clone(), p_b.clone()).await; - - let mgrs = vec![p_a.clone(), p_b.clone(), p_c.clone()]; - - let r_a = create_mock_route(p_a.clone()).await; - let r_b = create_mock_route(p_b.clone()).await; - let r_c = create_mock_route(p_c.clone()).await; - - for r in [r_a.clone(), r_b.clone(), r_c.clone()].iter() { - wait_for_condition( - || async { r.service_impl.synced_route_info.peer_infos.read().len() == 3 }, - Duration::from_secs(5), - ) - .await; - } - - tokio::time::sleep(tokio::time::Duration::from_secs(1)).await; - print_routes(vec![r_a.clone(), r_b.clone(), r_c.clone()]).await; - check_route_sanity(&r_a, mgrs.clone()).await; - check_route_sanity(&r_b, mgrs.clone()).await; - check_route_sanity(&r_c, mgrs.clone()).await; - - assert_eq!(2, r_a.list_routes().await.len()); - - drop(mgrs); - drop(r_c); - drop(p_c); - - for r in [r_a.clone(), r_b.clone()].iter() { - wait_for_condition( - || async { r.list_routes().await.len() == 1 }, - Duration::from_secs(5), - ) - .await; - } - } - - #[rstest::rstest] - #[tokio::test] - async fn peer_reconnect(#[values(true, false)] enable_conn_list_sync: bool) { - FORCE_USE_CONN_LIST.store(enable_conn_list_sync, Ordering::Relaxed); - let p_a = create_mock_pmgr().await; - let p_b = create_mock_pmgr().await; - let r_a = create_mock_route(p_a.clone()).await; - let r_b = create_mock_route(p_b.clone()).await; - - connect_peer_manager(p_a.clone(), p_b.clone()).await; - - wait_for_condition( - || async { r_a.list_routes().await.len() == 1 }, - Duration::from_secs(5), - ) - .await; - - assert_eq!(1, r_b.list_routes().await.len()); - - check_rpc_counter(&r_a, p_b.my_peer_id(), 2, 2); - - p_a.get_peer_map() - .close_peer(p_b.my_peer_id()) - .await - .unwrap(); - wait_for_condition( - || async { r_a.list_routes().await.is_empty() }, - Duration::from_secs(5), - ) - .await; - - // reconnect - connect_peer_manager(p_a.clone(), p_b.clone()).await; - wait_for_condition( - || async { r_a.list_routes().await.len() == 1 }, - Duration::from_secs(5), - ) - .await; - - // wait session init - tokio::time::sleep(Duration::from_secs(1)).await; - - println!("session: {:?}", r_a.session_mgr.dump_sessions()); - check_rpc_counter(&r_a, p_b.my_peer_id(), 2, 2); - } - - #[rstest::rstest] - #[tokio::test] - async fn test_cost_calculator(#[values(true, false)] enable_conn_list_sync: bool) { - FORCE_USE_CONN_LIST.store(enable_conn_list_sync, Ordering::Relaxed); - let p_a = create_mock_pmgr().await; - let p_b = create_mock_pmgr().await; - let p_c = create_mock_pmgr().await; - let p_d = create_mock_pmgr().await; - connect_peer_manager(p_a.clone(), p_b.clone()).await; - connect_peer_manager(p_a.clone(), p_c.clone()).await; - connect_peer_manager(p_d.clone(), p_b.clone()).await; - connect_peer_manager(p_d.clone(), p_c.clone()).await; - connect_peer_manager(p_b.clone(), p_c.clone()).await; - - let _r_a = create_mock_route(p_a.clone()).await; - let _r_b = create_mock_route(p_b.clone()).await; - let _r_c = create_mock_route(p_c.clone()).await; - let r_d = create_mock_route(p_d.clone()).await; - - // in normal mode, packet from p_c should directly forward to p_a - wait_for_condition( - || async { (r_d.get_next_hop(p_a.my_peer_id()).await).is_some() }, - Duration::from_secs(5), - ) - .await; - - struct TestCostCalculator { - p_a_peer_id: PeerId, - p_b_peer_id: PeerId, - p_c_peer_id: PeerId, - p_d_peer_id: PeerId, - } - - impl RouteCostCalculatorInterface for TestCostCalculator { - fn calculate_cost(&self, src: PeerId, dst: PeerId) -> i32 { - if src == self.p_d_peer_id && dst == self.p_b_peer_id { - return 100; - } - - if src == self.p_d_peer_id && dst == self.p_c_peer_id { - return 1; - } - - if src == self.p_c_peer_id && dst == self.p_a_peer_id { - return 101; - } - - if src == self.p_b_peer_id && dst == self.p_a_peer_id { - return 1; - } - - if src == self.p_c_peer_id && dst == self.p_b_peer_id { - return 2; - } - - 1 - } - } - - r_d.set_route_cost_fn(Box::new(TestCostCalculator { - p_a_peer_id: p_a.my_peer_id(), - p_b_peer_id: p_b.my_peer_id(), - p_c_peer_id: p_c.my_peer_id(), - p_d_peer_id: p_d.my_peer_id(), - })) - .await; - - // after set cost, packet from p_c should forward to p_b first - wait_for_condition( - || async { - r_d.get_next_hop_with_policy(p_a.my_peer_id(), NextHopPolicy::LeastCost) - .await - == Some(p_c.my_peer_id()) - }, - Duration::from_secs(5), - ) - .await; - - wait_for_condition( - || async { - r_d.get_next_hop_with_policy(p_a.my_peer_id(), NextHopPolicy::LeastHop) - .await - == Some(p_b.my_peer_id()) - }, - Duration::from_secs(5), - ) - .await; - } - - #[rstest::rstest] - #[tokio::test] - async fn test_raw_peer_info(#[values(true, false)] enable_conn_list_sync: bool) { - FORCE_USE_CONN_LIST.store(enable_conn_list_sync, Ordering::Relaxed); - let mut req = SyncRouteInfoRequest::default(); - let raw_info_map: DashMap = DashMap::new(); - - req.peer_infos = Some(RoutePeerInfos { - items: vec![RoutePeerInfo { - peer_id: 1, - ..Default::default() - }], - }); - - let mut raw_req = DynamicMessage::new(RoutePeerInfo::default().descriptor()); - raw_req - .transcode_from(&req.peer_infos.as_ref().unwrap().items[0]) - .unwrap(); - raw_info_map.insert(1, raw_req); - - let out = PeerRouteServiceImpl::build_sync_route_raw_req(&req, &raw_info_map); - - let out_bytes = out.encode_to_vec(); - - let req2 = SyncRouteInfoRequest::decode(out_bytes.as_slice()).unwrap(); - - assert_eq!(req, req2); - } - - #[rstest::rstest] - #[tokio::test] - async fn test_peer_id_map_override(#[values(true, false)] enable_conn_list_sync: bool) { - FORCE_USE_CONN_LIST.store(enable_conn_list_sync, Ordering::Relaxed); - let p_a = create_mock_peer_manager().await; - let p_b = create_mock_peer_manager().await; - let p_c = create_mock_peer_manager().await; - - connect_peer_manager(p_a.clone(), p_b.clone()).await; - connect_peer_manager(p_b.clone(), p_c.clone()).await; - - let ip: Ipv4Inet = "10.0.0.1/24".parse().unwrap(); - let ipv6: Ipv6Inet = "2001:db8::1/64".parse().unwrap(); - let proxy: Ipv4Cidr = "10.3.0.0/24".parse().unwrap(); - let check_route_peer_id = async |p: Arc| { - let p = p.clone(); - wait_for_condition( - || async { - p_a.get_route().get_peer_id_by_ipv4(&ip.address()).await == Some(p.my_peer_id()) - && p_a.get_route().get_peer_id_by_ipv6(&ipv6.address()).await - == Some(p.my_peer_id()) - && p_a - .get_route() - .get_peer_id_by_ipv4(&proxy.first_address()) - .await - == Some(p.my_peer_id()) - }, - Duration::from_secs(5), - ) - .await; - }; - - p_c.get_global_ctx().set_ipv4(Some(ip)); - p_c.get_global_ctx().set_ipv6(Some(ipv6)); - p_c.get_global_ctx() - .config - .add_proxy_cidr(proxy, None) - .unwrap(); - check_route_peer_id(p_c.clone()).await; - - p_b.get_global_ctx().set_ipv4(Some(ip)); - p_b.get_global_ctx().set_ipv6(Some(ipv6)); - p_b.get_global_ctx() - .config - .add_proxy_cidr(proxy, None) - .unwrap(); - check_route_peer_id(p_b.clone()).await; - - p_b.get_global_ctx() - .set_ipv4(Some("10.0.0.2/24".parse().unwrap())); - p_b.get_global_ctx() - .set_ipv6(Some("2001:db8::2/64".parse().unwrap())); - p_b.get_global_ctx().config.remove_proxy_cidr(proxy); - check_route_peer_id(p_c.clone()).await; - } - #[rstest::rstest] - #[tokio::test] - async fn test_subnet_proxy_conflict(#[values(true, false)] enable_conn_list_sync: bool) { - FORCE_USE_CONN_LIST.store(enable_conn_list_sync, Ordering::Relaxed); - // Create three peer managers: A, B, C - let p_a = create_mock_peer_manager().await; - let p_b = create_mock_peer_manager().await; - let p_c = create_mock_peer_manager().await; - - // Connect A-B-C in a line topology - connect_peer_manager(p_a.clone(), p_b.clone()).await; - connect_peer_manager(p_b.clone(), p_c.clone()).await; - - // Create routes for testing - let route_a = p_a.get_route(); - let route_b = p_b.get_route(); - - // Define the proxy CIDR that will be used by both A and B - let proxy_cidr: Ipv4Cidr = "192.168.100.0/24".parse().unwrap(); - let test_ip = proxy_cidr.first_address(); - - let mut cidr_peer_id_map: PrefixMap = PrefixMap::new(); - cidr_peer_id_map.insert( - proxy_cidr, - PeerIdVersion { - peer_id: p_c.my_peer_id(), - version: 0, - }, - ); - assert_eq!( - cidr_peer_id_map - .get_lpm(&Ipv4Cidr::new(test_ip, 32).unwrap()) - .map(|v| v.1.peer_id) - .unwrap_or(0), - p_c.my_peer_id(), - ); - - // First, add proxy CIDR to node C to establish a baseline route - p_c.get_global_ctx() - .config - .add_proxy_cidr(proxy_cidr, None) - .unwrap(); - - // Wait for route convergence - A should route to C for the proxy CIDR - wait_for_condition( - || async { - let peer_id_for_proxy = route_a.get_peer_id_by_ipv4(&test_ip).await; - peer_id_for_proxy == Some(p_c.my_peer_id()) - }, - Duration::from_secs(10), - ) - .await; - - // Now add the same proxy CIDR to node A (creating a conflict) - p_a.get_global_ctx() - .config - .add_proxy_cidr(proxy_cidr, None) - .unwrap(); - - // Wait for route convergence - A should now route to itself for the proxy CIDR - wait_for_condition( - || async { route_a.get_peer_id_by_ipv4(&test_ip).await == Some(p_a.my_peer_id()) }, - Duration::from_secs(10), - ) - .await; - - // Also add the same proxy CIDR to node B (creating another conflict) - p_b.get_global_ctx() - .config - .add_proxy_cidr(proxy_cidr, None) - .unwrap(); - - // Wait for route convergence - B should route to itself for the proxy CIDR - wait_for_condition( - || async { route_b.get_peer_id_by_ipv4(&test_ip).await == Some(p_b.my_peer_id()) }, - Duration::from_secs(5), - ) - .await; - - // Final verification: A should still route to itself even with multiple conflicts - assert_eq!( - route_a.get_peer_id_by_ipv4(&test_ip).await, - Some(p_a.my_peer_id()) - ); - - // remove proxy on A, a should route to B - p_a.get_global_ctx().config.remove_proxy_cidr(proxy_cidr); - wait_for_condition( - || async { - let peer_id_for_proxy = route_a.get_peer_id_by_ipv4(&test_ip).await; - peer_id_for_proxy == Some(p_b.my_peer_id()) - }, - Duration::from_secs(10), - ) - .await; - } - #[rstest::rstest] - #[tokio::test] - async fn test_connect_at_different_time(#[values(true, false)] enable_conn_list_sync: bool) { - FORCE_USE_CONN_LIST.store(enable_conn_list_sync, Ordering::Relaxed); - // Create three peer managers: A, B, C - let p_a = create_mock_peer_manager().await; - let p_b = create_mock_peer_manager().await; - let p_c = create_mock_peer_manager().await; - - // Connect A-B-C in a line topology - connect_peer_manager(p_a.clone(), p_b.clone()).await; - - wait_route_appear(p_a.clone(), p_b.clone()).await.unwrap(); - - connect_peer_manager(p_b.clone(), p_c.clone()).await; - wait_route_appear(p_a.clone(), p_c.clone()).await.unwrap(); - } - - /// Helper: create a raw DynamicMessage from a RoutePeerInfo with an extra - /// unknown field appended (field number 9999, varint value 42). - /// Returns the raw DynamicMessage and the encoded unknown field bytes. - fn make_raw_with_unknown_field(info: &RoutePeerInfo) -> (DynamicMessage, Vec) { - // Encode the info to bytes - let mut bytes = info.encode_to_vec(); - // Append an unknown field: field 9999, wire type 0 (varint), value 42 - // Tag = (9999 << 3) | 0 = 79992, encoded as varint - prost::encoding::encode_key(9999, prost::encoding::WireType::Varint, &mut bytes); - prost::encoding::encode_varint(42, &mut bytes); - let unknown_field_bytes = bytes[info.encoded_len()..].to_vec(); - // Decode as DynamicMessage — unknown fields are preserved - let raw = DynamicMessage::decode(RoutePeerInfo::default().descriptor(), bytes.as_slice()) - .unwrap(); - (raw, unknown_field_bytes) - } - - /// Check that a raw DynamicMessage still contains the unknown field bytes - /// by re-encoding and checking the suffix. - fn raw_has_unknown_bytes(raw: &DynamicMessage, unknown_bytes: &[u8]) -> bool { - let encoded = raw.encode_to_vec(); - // The unknown field bytes should appear somewhere in the encoded output - encoded - .windows(unknown_bytes.len()) - .any(|w| w == unknown_bytes) - } - - fn encode_length_delimited_field(field_number: u32, payload: &[u8], dst: &mut Vec) { - prost::encoding::encode_key( - field_number, - prost::encoding::WireType::LengthDelimited, - dst, - ); - prost::encoding::encode_varint(payload.len() as u64, dst); - dst.extend_from_slice(payload); - } - - fn make_route_info_with_raw_trusted_credential_proof( - info: &RoutePeerInfo, - raw_credential_bytes: &[u8], - credential_hmac: &[u8], - ) -> (RoutePeerInfo, DynamicMessage) { - let mut proof_bytes = Vec::new(); - encode_length_delimited_field(1, raw_credential_bytes, &mut proof_bytes); - encode_length_delimited_field(2, credential_hmac, &mut proof_bytes); - - let mut route_info_bytes = info.encode_to_vec(); - encode_length_delimited_field(19, &proof_bytes, &mut route_info_bytes); - - let typed_info = RoutePeerInfo::decode(route_info_bytes.as_slice()).unwrap(); - let raw_info = DynamicMessage::decode( - RoutePeerInfo::default().descriptor(), - route_info_bytes.as_slice(), - ) - .unwrap(); - - (typed_info, raw_info) - } - - #[tokio::test] - async fn sync_route_preserves_unknown_fields_for_credential_sender() { - let peer_mgr = create_mock_pmgr().await; - let route = create_mock_route(peer_mgr.clone()).await; - let from_peer_id: PeerId = 20001; - let credential_pubkey = vec![4u8; 32]; - - let identity_type = DashMap::new(); - identity_type.insert(from_peer_id, PeerIdentityType::Credential); - let peer_public_key = DashMap::new(); - peer_public_key.insert(from_peer_id, credential_pubkey.clone()); - *route.service_impl.interface.lock().await = Some(Box::new(AuthOnlyInterface { - my_peer_id: peer_mgr.my_peer_id(), - identity_type, - peer_public_key, - })); - route - .service_impl - .synced_route_info - .trusted_credential_pubkeys - .insert( - credential_pubkey.clone(), - TrustedCredentialPubkey { - pubkey: credential_pubkey, - expiry_unix: i64::MAX, - ..Default::default() - }, - ); - - let mut sender_info = RoutePeerInfo::new(); - sender_info.peer_id = from_peer_id; - sender_info.version = 1; - - let (raw, unknown_bytes) = make_raw_with_unknown_field(&sender_info); - - route - .session_mgr - .do_sync_route_info( - from_peer_id, - 1, - true, - Some(vec![sender_info]), - Some(vec![raw]), - None, - None, - ) - .await - .unwrap(); - - let stored_raw = route - .service_impl - .synced_route_info - .raw_peer_infos - .get(&from_peer_id) - .expect("raw peer info should be stored"); - assert!( - raw_has_unknown_bytes(stored_raw.value(), &unknown_bytes), - "unknown fields should be preserved for credential sender" - ); - } - - #[tokio::test] - async fn sync_route_preserves_unknown_fields_for_shared_sender() { - let peer_mgr = create_mock_pmgr().await; - let route = create_mock_route(peer_mgr.clone()).await; - let from_peer_id: PeerId = 20011; - let forwarded_peer_id: PeerId = 20012; - - let identity_type = DashMap::new(); - identity_type.insert(from_peer_id, PeerIdentityType::SharedNode); - *route.service_impl.interface.lock().await = Some(Box::new(AuthOnlyInterface { - my_peer_id: peer_mgr.my_peer_id(), - identity_type, - peer_public_key: DashMap::new(), - })); - - let mut sender_info = RoutePeerInfo::new(); - sender_info.peer_id = from_peer_id; - sender_info.version = 1; - - let mut forwarded_info = RoutePeerInfo::new(); - forwarded_info.peer_id = forwarded_peer_id; - forwarded_info.version = 1; - forwarded_info.trusted_credential_pubkeys = vec![TrustedCredentialPubkeyProof { - credential: Some(TrustedCredentialPubkey { - pubkey: vec![9u8; 32], - expiry_unix: i64::MAX, - ..Default::default() - }), - credential_hmac: vec![1; 32], - }]; - - let (raw_sender, unknown_sender) = make_raw_with_unknown_field(&sender_info); - let (raw_forwarded, unknown_forwarded) = make_raw_with_unknown_field(&forwarded_info); - - route - .session_mgr - .do_sync_route_info( - from_peer_id, - 1, - true, - Some(vec![sender_info, forwarded_info]), - Some(vec![raw_sender, raw_forwarded]), - None, - None, - ) - .await - .unwrap(); - - // Shared node: trusted_credential_pubkeys cleared but unknown fields preserved - let stored_sender = route - .service_impl - .synced_route_info - .raw_peer_infos - .get(&from_peer_id) - .expect("sender raw should be stored"); - assert!( - raw_has_unknown_bytes(stored_sender.value(), &unknown_sender), - "unknown fields should be preserved for shared sender's own info" - ); - - let stored_forwarded = route - .service_impl - .synced_route_info - .raw_peer_infos - .get(&forwarded_peer_id) - .expect("forwarded raw should be stored"); - assert!( - raw_has_unknown_bytes(stored_forwarded.value(), &unknown_forwarded), - "unknown fields should be preserved for shared sender's forwarded info" - ); - } - - #[tokio::test] - async fn sync_route_preserves_unknown_fields_for_admin_sender() { - let peer_mgr = create_mock_pmgr().await; - let route = create_mock_route(peer_mgr.clone()).await; - let from_peer_id: PeerId = 20021; - - let identity_type = DashMap::new(); - identity_type.insert(from_peer_id, PeerIdentityType::Admin); - *route.service_impl.interface.lock().await = Some(Box::new(AuthOnlyInterface { - my_peer_id: peer_mgr.my_peer_id(), - identity_type, - peer_public_key: DashMap::new(), - })); - - let mut sender_info = RoutePeerInfo::new(); - sender_info.peer_id = from_peer_id; - sender_info.version = 1; - // Set is_credential_peer=true so the mark_credential_peer(false) path triggers - sender_info.feature_flag = Some(PeerFeatureFlag { - is_credential_peer: true, - ..Default::default() - }); - - let (raw, unknown_bytes) = make_raw_with_unknown_field(&sender_info); - - route - .session_mgr - .do_sync_route_info( - from_peer_id, - 1, - true, - Some(vec![sender_info]), - Some(vec![raw]), - None, - None, - ) - .await - .unwrap(); - - let stored_raw = route - .service_impl - .synced_route_info - .raw_peer_infos - .get(&from_peer_id) - .expect("raw peer info should be stored"); - assert!( - raw_has_unknown_bytes(stored_raw.value(), &unknown_bytes), - "unknown fields should be preserved for admin sender (mark non-credential path)" - ); - } - - #[tokio::test] - async fn sync_route_info_prioritizes_local_over_remote_for_overlapped_proxy_cidrs() { - let peer_mgr = create_mock_pmgr().await; - let route = create_mock_route(peer_mgr.clone()).await; - let from_peer_id: PeerId = 11001; - - let peers = Arc::new(Mutex::new(vec![from_peer_id])); - let peer_identity_types = Arc::new(Mutex::new(HashMap::from([( - from_peer_id, - Some(PeerIdentityType::Admin), - )]))); - *route.service_impl.interface.lock().await = Some(Box::new(CountingInterface { - my_peer_id: peer_mgr.my_peer_id(), - peers, - peer_identity_types, - list_peers_calls: Arc::new(AtomicU32::new(0)), - get_peer_identity_type_calls: Arc::new(AtomicU32::new(0)), - })); - route.service_impl.mark_interface_peers_dirty(); - assert!(route.service_impl.update_my_conn_info().await); - - route - .service_impl - .global_ctx - .config - .add_proxy_cidr("10.10.0.0/16".parse().unwrap(), None) - .unwrap(); - assert!(route.service_impl.update_my_peer_info()); - - let mut sender_info = RoutePeerInfo::new(); - sender_info.peer_id = from_peer_id; - sender_info.version = 1; - sender_info.proxy_cidrs = vec![ - "10.10.0.0/16".to_string(), - "10.10.1.0/24".to_string(), - "10.11.0.0/16".to_string(), - ]; - - let make_raw = |info: &RoutePeerInfo| { - let mut raw = DynamicMessage::new(RoutePeerInfo::default().descriptor()); - raw.transcode_from(info).unwrap(); - raw - }; - - route - .session_mgr - .do_sync_route_info( - from_peer_id, - 1, - true, - Some(vec![sender_info.clone()]), - Some(vec![make_raw(&sender_info)]), - None, - None, - ) - .await - .unwrap(); - - // Keep route table in sync with interface-derived adjacency during assertion window. - route - .service_impl - .update_route_table_and_cached_local_conn_bitmap(); - - // Control plane: keep what remote announced. - let guard = route.service_impl.synced_route_info.peer_infos.read(); - let stored = guard.get(&from_peer_id).unwrap(); - assert_eq!(stored.proxy_cidrs, sender_info.proxy_cidrs); - drop(guard); - - // Route-table filtering: local announced /16 should dominate remote equal/subset. - assert_eq!( - route - .service_impl - .route_table - .get_peer_id_for_proxy(&"10.10.1.1".parse::().unwrap()), - Some(peer_mgr.my_peer_id()) - ); - // Non-overlapped remote prefix should still route to remote. - assert_eq!( - route - .service_impl - .route_table - .get_peer_id_for_proxy(&"10.11.0.1".parse::().unwrap()), - Some(from_peer_id) + table.get_peer_id_for_proxy(&"10.10.1.1".parse::().unwrap()), + Some(3) ); } } diff --git a/easytier-core/src/peers/route/route_peer_wire.rs b/easytier-core/src/peers/route/route_peer_wire.rs new file mode 100644 index 00000000..5a5cddd9 --- /dev/null +++ b/easytier-core/src/peers/route/route_peer_wire.rs @@ -0,0 +1,399 @@ +//! Minimal protobuf wire editing used by OSPF route reflection. +//! +//! Route calculation uses generated prost types. This module only keeps the +//! original `RoutePeerInfo` bytes and replaces the two fields that credential +//! filtering is allowed to change, leaving every other field byte-for-byte +//! intact. + +use bytes::Bytes; +use prost::{ + Message, + encoding::{ + DecodeContext, WireType, decode_key, decode_varint, encode_key, encode_varint, skip_field, + }, +}; +use thiserror::Error; + +use crate::proto::peer_rpc::{RoutePeerInfo, SyncRouteInfoRequest}; + +const SYNC_ROUTE_PEER_INFOS_TAG: u32 = 4; +const ROUTE_PEER_INFOS_ITEM_TAG: u32 = 1; +const ROUTE_PEER_INFO_PROXY_CIDRS_TAG: u32 = 5; +const ROUTE_PEER_INFO_FEATURE_FLAG_TAG: u32 = 11; +const ROUTE_PEER_INFO_CREDENTIAL_PROOF_TAG: u32 = 19; +const FEATURE_FLAG_IS_CREDENTIAL_PEER_TAG: u32 = 8; +const CREDENTIAL_PROOF_CREDENTIAL_TAG: u32 = 1; + +pub(crate) type RawRoutePeerInfo = Bytes; + +#[derive(Debug, Error)] +pub(crate) enum WireError { + #[error(transparent)] + Decode(#[from] prost::DecodeError), + #[error("protobuf field {tag} has wire type {actual:?}, expected {expected:?}")] + WrongWireType { + tag: u32, + actual: WireType, + expected: WireType, + }, + #[error("protobuf length-delimited field is larger than the remaining input")] + TruncatedLengthDelimited, + #[error("raw RoutePeerInfo count does not match the decoded request")] + PeerInfoCountMismatch, +} + +type Result = std::result::Result; + +#[derive(Clone, Copy)] +struct WireField<'a> { + tag: u32, + wire_type: WireType, + encoded: &'a [u8], + length_delimited: Option<&'a [u8]>, +} + +fn parse_fields(message: &[u8]) -> Result>> { + let mut input = message; + let mut fields = Vec::new(); + + while !input.is_empty() { + let start = message.len() - input.len(); + let (tag, wire_type) = decode_key(&mut input)?; + let length_delimited = if wire_type == WireType::LengthDelimited { + let len = usize::try_from(decode_varint(&mut input)?) + .map_err(|_| WireError::TruncatedLengthDelimited)?; + if len > input.len() { + return Err(WireError::TruncatedLengthDelimited); + } + let (payload, remaining) = input.split_at(len); + input = remaining; + Some(payload) + } else { + skip_field(wire_type, tag, &mut input, DecodeContext::default())?; + None + }; + let end = message.len() - input.len(); + fields.push(WireField { + tag, + wire_type, + encoded: &message[start..end], + length_delimited, + }); + } + + Ok(fields) +} + +fn length_delimited_fields(message: &[u8], tag: u32) -> Result> { + parse_fields(message)? + .into_iter() + .filter(|field| field.tag == tag) + .map(|field| { + field.length_delimited.ok_or(WireError::WrongWireType { + tag, + actual: field.wire_type, + expected: WireType::LengthDelimited, + }) + }) + .collect() +} + +fn replace_fields( + message: &[u8], + tag: u32, + expected_wire_type: WireType, + replacement: &[u8], +) -> Result> { + let mut output = Vec::with_capacity(message.len() + replacement.len()); + for field in parse_fields(message)? { + if field.tag != tag { + output.extend_from_slice(field.encoded); + continue; + } + if field.wire_type != expected_wire_type { + return Err(WireError::WrongWireType { + tag, + actual: field.wire_type, + expected: expected_wire_type, + }); + } + } + output.extend_from_slice(replacement); + Ok(output) +} + +fn append_length_delimited_field(output: &mut Vec, tag: u32, payload: &[u8]) { + encode_key(tag, WireType::LengthDelimited, output); + encode_varint(payload.len() as u64, output); + output.extend_from_slice(payload); +} + +fn append_varint_field(output: &mut Vec, tag: u32, value: u64) { + encode_key(tag, WireType::Varint, output); + encode_varint(value, output); +} + +fn merged_length_delimited_field(message: &[u8], tag: u32) -> Result> { + let values = length_delimited_fields(message, tag)?; + let total_len = values.iter().map(|value| value.len()).sum(); + let mut merged = Vec::with_capacity(total_len); + for value in values { + merged.extend_from_slice(value); + } + Ok(merged) +} + +pub(crate) fn raw_route_peer_info(info: &RoutePeerInfo) -> RawRoutePeerInfo { + Bytes::from(info.encode_to_vec()) +} + +pub(crate) fn extract_route_peer_infos(request: &[u8]) -> Result> { + let mut result = Vec::new(); + for peer_infos in length_delimited_fields(request, SYNC_ROUTE_PEER_INFOS_TAG)? { + result.extend( + length_delimited_fields(peer_infos, ROUTE_PEER_INFOS_ITEM_TAG)? + .into_iter() + .map(Bytes::copy_from_slice), + ); + } + Ok(result) +} + +pub(crate) fn encode_sync_route_request( + request: &SyncRouteInfoRequest, + raw_peer_infos: &[RawRoutePeerInfo], +) -> Result> { + let decoded_count = request + .peer_infos + .as_ref() + .map(|peer_infos| peer_infos.items.len()) + .unwrap_or_default(); + if decoded_count != raw_peer_infos.len() { + return Err(WireError::PeerInfoCountMismatch); + } + + let mut request_without_peer_infos = request.clone(); + request_without_peer_infos.peer_infos = None; + let mut output = request_without_peer_infos.encode_to_vec(); + if request.peer_infos.is_some() { + let mut peer_infos = Vec::new(); + for info in raw_peer_infos { + append_length_delimited_field(&mut peer_infos, ROUTE_PEER_INFOS_ITEM_TAG, info); + } + append_length_delimited_field(&mut output, SYNC_ROUTE_PEER_INFOS_TAG, &peer_infos); + } + Ok(output) +} + +pub(crate) fn patch_credential_route_peer_info( + raw: &RawRoutePeerInfo, + proxy_cidrs: &[String], +) -> Result { + let mut proxy_cidr_fields = Vec::new(); + for cidr in proxy_cidrs { + append_length_delimited_field( + &mut proxy_cidr_fields, + ROUTE_PEER_INFO_PROXY_CIDRS_TAG, + cidr.as_bytes(), + ); + } + let route_info = replace_fields( + raw, + ROUTE_PEER_INFO_PROXY_CIDRS_TAG, + WireType::LengthDelimited, + &proxy_cidr_fields, + )?; + + // Singular message fields merge when they occur more than once. Concatenating + // their payloads preserves that protobuf behavior before changing tag 8. + let feature_flag = + merged_length_delimited_field(&route_info, ROUTE_PEER_INFO_FEATURE_FLAG_TAG)?; + let mut credential_flag = Vec::new(); + append_varint_field(&mut credential_flag, FEATURE_FLAG_IS_CREDENTIAL_PEER_TAG, 1); + let feature_flag = replace_fields( + &feature_flag, + FEATURE_FLAG_IS_CREDENTIAL_PEER_TAG, + WireType::Varint, + &credential_flag, + )?; + let mut feature_flag_field = Vec::new(); + append_length_delimited_field( + &mut feature_flag_field, + ROUTE_PEER_INFO_FEATURE_FLAG_TAG, + &feature_flag, + ); + + Ok(Bytes::from(replace_fields( + &route_info, + ROUTE_PEER_INFO_FEATURE_FLAG_TAG, + WireType::LengthDelimited, + &feature_flag_field, + )?)) +} + +pub(crate) fn raw_credential_bytes( + raw_route_info: &RawRoutePeerInfo, + proof_idx: usize, +) -> Result> { + let proofs = length_delimited_fields(raw_route_info, ROUTE_PEER_INFO_CREDENTIAL_PROOF_TAG)?; + let Some(proof) = proofs.get(proof_idx) else { + return Ok(None); + }; + let credentials = length_delimited_fields(proof, CREDENTIAL_PROOF_CREDENTIAL_TAG)?; + if credentials.is_empty() { + return Ok(None); + } + + // Multiple occurrences of a singular message field merge. Concatenation is + // the wire-equivalent merged message and keeps nested unknown fields intact. + let total_len = credentials.iter().map(|value| value.len()).sum(); + let mut merged = Vec::with_capacity(total_len); + for credential in credentials { + merged.extend_from_slice(credential); + } + Ok(Some(Bytes::from(merged))) +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::proto::{ + common::PeerFeatureFlag, + peer_rpc::{RoutePeerInfos, TrustedCredentialPubkey, TrustedCredentialPubkeyProof}, + }; + + fn encoded_varint_field(tag: u32, value: u64) -> Vec { + let mut output = Vec::new(); + append_varint_field(&mut output, tag, value); + output + } + + fn encoded_length_delimited_field(tag: u32, payload: &[u8]) -> Vec { + let mut output = Vec::new(); + append_length_delimited_field(&mut output, tag, payload); + output + } + + fn encoded_fields(message: &[u8], tag: u32) -> Vec> { + parse_fields(message) + .unwrap() + .into_iter() + .filter(|field| field.tag == tag) + .map(|field| field.encoded.to_vec()) + .collect() + } + + #[test] + fn patch_preserves_top_level_and_nested_unknown_fields() { + const TOP_LEVEL_UNKNOWN_TAG: u32 = 100; + const FEATURE_UNKNOWN_TAG: u32 = 101; + + let mut feature_flag = PeerFeatureFlag { + avoid_relay_data: true, + ..Default::default() + } + .encode_to_vec(); + feature_flag.extend(encoded_varint_field(FEATURE_UNKNOWN_TAG, 42)); + + let info = RoutePeerInfo { + peer_id: 7, + proxy_cidrs: vec!["10.0.0.0/8".to_owned()], + ..Default::default() + }; + let mut raw = info.encode_to_vec(); + raw.extend(encoded_length_delimited_field( + ROUTE_PEER_INFO_FEATURE_FLAG_TAG, + &feature_flag, + )); + raw.extend(encoded_varint_field(TOP_LEVEL_UNKNOWN_TAG, 99)); + + let top_level_unknown = encoded_fields(&raw, TOP_LEVEL_UNKNOWN_TAG); + let nested_unknown = encoded_fields(&feature_flag, FEATURE_UNKNOWN_TAG); + let patched = + patch_credential_route_peer_info(&Bytes::from(raw), &["10.1.0.0/16".to_owned()]) + .unwrap(); + + assert_eq!( + encoded_fields(&patched, TOP_LEVEL_UNKNOWN_TAG), + top_level_unknown + ); + let patched_feature = + merged_length_delimited_field(&patched, ROUTE_PEER_INFO_FEATURE_FLAG_TAG).unwrap(); + assert_eq!( + encoded_fields(&patched_feature, FEATURE_UNKNOWN_TAG), + nested_unknown + ); + + let decoded = RoutePeerInfo::decode(patched).unwrap(); + assert_eq!(decoded.proxy_cidrs, ["10.1.0.0/16"]); + assert!(decoded.feature_flag.unwrap().is_credential_peer); + } + + #[test] + fn multi_hop_sync_keeps_raw_peer_info_exactly() { + let info = RoutePeerInfo { + peer_id: 9, + ..Default::default() + }; + let mut raw = info.encode_to_vec(); + raw.extend(encoded_length_delimited_field(120, b"future")); + let raw = Bytes::from(raw); + let request = SyncRouteInfoRequest { + my_peer_id: 1, + peer_infos: Some(RoutePeerInfos { items: vec![info] }), + ..Default::default() + }; + + let first_hop = encode_sync_route_request(&request, std::slice::from_ref(&raw)).unwrap(); + let first_hop_raw = extract_route_peer_infos(&first_hop).unwrap(); + assert_eq!(first_hop_raw.as_slice(), std::slice::from_ref(&raw)); + assert_eq!( + SyncRouteInfoRequest::decode(first_hop.as_slice()) + .unwrap() + .peer_infos, + request.peer_infos + ); + + let second_hop = encode_sync_route_request(&request, &first_hop_raw).unwrap(); + assert_eq!(extract_route_peer_infos(&second_hop).unwrap(), [raw]); + } + + #[test] + fn credential_hmac_uses_exact_nested_message_bytes() { + let secret = "wire-test-secret"; + let credential = TrustedCredentialPubkey { + pubkey: vec![3; 32], + ..Default::default() + }; + let mut raw_credential = credential.encode_to_vec(); + raw_credential.extend(encoded_varint_field(100, 1234)); + let hmac = TrustedCredentialPubkeyProof::generate_credential_hmac_from_bytes( + &raw_credential, + secret, + ); + + let mut raw_proof = + encoded_length_delimited_field(CREDENTIAL_PROOF_CREDENTIAL_TAG, &raw_credential); + raw_proof.extend(encoded_length_delimited_field(2, &hmac)); + let mut raw_route_info = RoutePeerInfo { + peer_id: 11, + ..Default::default() + } + .encode_to_vec(); + raw_route_info.extend(encoded_length_delimited_field( + ROUTE_PEER_INFO_CREDENTIAL_PROOF_TAG, + &raw_proof, + )); + + let extracted = raw_credential_bytes(&Bytes::from(raw_route_info), 0) + .unwrap() + .unwrap(); + assert_eq!(extracted, raw_credential); + + let proof = TrustedCredentialPubkeyProof { + credential: Some(credential), + credential_hmac: hmac, + }; + assert!(proof.verify_credential_hmac_with_bytes(&extracted, secret)); + } +} diff --git a/easytier-core/src/peers/test_support.rs b/easytier-core/src/peers/test_support.rs new file mode 100644 index 00000000..7da13bb1 --- /dev/null +++ b/easytier-core/src/peers/test_support.rs @@ -0,0 +1,100 @@ +//! Test-only peer-context fakes shared by peer-domain unit tests +//! (`peers::tests`, `peers::route::peer_ospf_route::tests`, and +//! `context::tests`). Kept out of `context.rs` so the context unit tests and +//! their consumers share one definition. + +use std::net::IpAddr; + +use cidr::{Ipv4Inet, Ipv6Inet}; +use easytier_proto::common::{FlagsInConfig, SecureModeConfig}; +use hmac::Hmac; +use sha2::Sha256; + +use crate::{ + config::peers::PeerRuntimeConfig, + config::{CoreConfig, IpPrefix, NodeConfig, PeerPolicyConfig, RouteConfig, TrafficConfig}, + peers::context::{NetworkIdentity, PeerContext, secret_proof_from_secret}, +}; + +pub(crate) trait PeerContextTestExt: PeerContext { + fn runtime_config(&self) -> PeerRuntimeConfig { + let network_identity = self.network_identity(); + let hostname = self.hostname(); + PeerRuntimeConfig { + core: CoreConfig { + node: NodeConfig { + peer_id: None, + instance_id: Some(*self.instance_id().as_bytes()), + hostname: (!hostname.is_empty()).then_some(hostname), + network_name: network_identity.network_name.clone(), + }, + routes: RouteConfig { + ipv4: self.ipv4().map(ipv4_inet_to_config), + ipv6: self.ipv6().map(ipv6_inet_to_config), + ..Default::default() + }, + peer_policy: PeerPolicyConfig::default(), + traffic: TrafficConfig::default(), + }, + network_identity, + stun_info: self.stun_info(), + feature_flags: self.feature_flags(), + secure_mode: self.secure_mode(), + host_routing: self.host_routing_policy(), + } + } +} + +fn ipv4_inet_to_config(value: Ipv4Inet) -> IpPrefix { + IpPrefix::new(IpAddr::V4(value.address()), value.network_length()) + .expect("Ipv4Inet should always have a valid IPv4 prefix length") +} + +fn ipv6_inet_to_config(value: Ipv6Inet) -> IpPrefix { + IpPrefix::new(IpAddr::V6(value.address()), value.network_length()) + .expect("Ipv6Inet should always have a valid IPv6 prefix length") +} + +#[derive(Debug, Clone)] +pub(crate) struct NoopPeerContext { + network_identity: NetworkIdentity, + flags: FlagsInConfig, + secure_mode: Option, +} + +impl NoopPeerContext { + pub(crate) fn new(network_identity: NetworkIdentity) -> Self { + Self { + network_identity, + flags: FlagsInConfig::default(), + secure_mode: None, + } + } +} + +impl Default for NoopPeerContext { + fn default() -> Self { + Self::new(NetworkIdentity::default()) + } +} + +impl PeerContext for NoopPeerContext { + fn network_identity(&self) -> NetworkIdentity { + self.network_identity.clone() + } + + fn flags(&self) -> FlagsInConfig { + self.flags.clone() + } + + fn secure_mode(&self) -> Option { + self.secure_mode.clone() + } + + fn secret_proof(&self, challenge: &[u8]) -> Option> { + let secret = self.network_identity.network_secret.as_ref()?; + secret_proof_from_secret(secret, challenge) + } +} + +impl PeerContextTestExt for NoopPeerContext {} diff --git a/easytier-core/src/peers/tests.rs b/easytier-core/src/peers/tests.rs new file mode 100644 index 00000000..9a398eb3 --- /dev/null +++ b/easytier-core/src/peers/tests.rs @@ -0,0 +1,113 @@ +use std::sync::Arc; + +use crate::foundation::time::{Duration, timeout}; + +use crate::{ + packet::{PacketType, ZCPacket}, + peers::{ + conn::{peer_conn::PeerConn, peer_map::PeerMap, peer_session::PeerSessionStore}, + context::NetworkIdentity, + create_packet_recv_chan, + error::Error, + test_support::NoopPeerContext, + }, + tunnel::ring::create_ring_tunnel_pair, +}; + +impl PeerConn { + #[tracing::instrument] + async fn do_handshake_as_server(&mut self) -> Result<(), Error> { + self.do_handshake_as_server_ext(|_, _| Ok(())).await + } +} + +#[tokio::test] +async fn peer_conn_handshake_over_memory_tunnel() { + let peer_session_store = Arc::new(PeerSessionStore::new()); + let (client_tunnel, server_tunnel) = create_ring_tunnel_pair(); + let client_ctx = Arc::new(NoopPeerContext::default()); + let server_ctx = Arc::new(NoopPeerContext::default()); + + let mut client = PeerConn::new(1, client_ctx, client_tunnel, peer_session_store.clone()); + let mut server = PeerConn::new(2, server_ctx, server_tunnel, peer_session_store); + + let (client_ret, server_ret) = tokio::join!( + client.do_handshake_as_client(), + server.do_handshake_as_server() + ); + + client_ret.unwrap(); + server_ret.unwrap(); + assert_eq!(client.get_peer_id(), 2); + assert_eq!(server.get_peer_id(), 1); +} + +#[tokio::test] +async fn peer_conn_handshake_matches_plaintext_secret_identity() { + let peer_session_store = Arc::new(PeerSessionStore::new()); + let (client_tunnel, server_tunnel) = create_ring_tunnel_pair(); + let client_ctx = Arc::new(NoopPeerContext::new(NetworkIdentity { + network_name: "net".to_string(), + network_secret: Some("secret".to_string()), + network_secret_digest: None, + })); + let server_ctx = Arc::new(NoopPeerContext::new(NetworkIdentity { + network_name: "net".to_string(), + network_secret: Some("secret".to_string()), + network_secret_digest: None, + })); + + let mut client = PeerConn::new(1, client_ctx, client_tunnel, peer_session_store.clone()); + let mut server = PeerConn::new(2, server_ctx, server_tunnel, peer_session_store); + + let (client_ret, server_ret) = tokio::join!( + client.do_handshake_as_client(), + server.do_handshake_as_server() + ); + + client_ret.unwrap(); + server_ret.unwrap(); + assert!(client.matches_local_network_secret()); + assert!(server.matches_local_network_secret()); +} + +#[tokio::test] +async fn peer_map_forwards_packet_over_memory_tunnel() { + let peer_session_store = Arc::new(PeerSessionStore::new()); + let (client_tunnel, server_tunnel) = create_ring_tunnel_pair(); + let client_ctx = Arc::new(NoopPeerContext::default()); + let server_ctx = Arc::new(NoopPeerContext::default()); + + let mut client_conn = PeerConn::new( + 1, + client_ctx.clone(), + client_tunnel, + peer_session_store.clone(), + ); + let mut server_conn = PeerConn::new(2, server_ctx.clone(), server_tunnel, peer_session_store); + + let (client_ret, server_ret) = tokio::join!( + client_conn.do_handshake_as_client(), + server_conn.do_handshake_as_server() + ); + client_ret.unwrap(); + server_ret.unwrap(); + + let (client_tx, _client_rx) = create_packet_recv_chan(); + let (server_tx, mut server_rx) = create_packet_recv_chan(); + let client_map = PeerMap::new(client_tx, client_ctx, 1); + let server_map = PeerMap::new(server_tx, server_ctx, 2); + + client_map.add_new_peer_conn(client_conn).await.unwrap(); + server_map.add_new_peer_conn(server_conn).await.unwrap(); + + let mut packet = ZCPacket::new_with_payload(b"hello"); + packet.fill_peer_manager_hdr(1, 2, PacketType::Data as u8); + client_map.send_msg_directly(packet, 2).await.unwrap(); + + let received = timeout(Duration::from_secs(1), server_rx.recv()) + .await + .unwrap() + .unwrap(); + assert_eq!(received.payload(), b"hello"); +} diff --git a/easytier/src/peers/traffic_metrics.rs b/easytier-core/src/peers/traffic_metrics.rs similarity index 83% rename from easytier/src/peers/traffic_metrics.rs rename to easytier-core/src/peers/traffic_metrics.rs index 398aeaa5..810ab797 100644 --- a/easytier/src/peers/traffic_metrics.rs +++ b/easytier-core/src/peers/traffic_metrics.rs @@ -3,17 +3,16 @@ use std::{future::Future, sync::Arc}; use dashmap::DashMap; use futures::future::BoxFuture; -use crate::common::{ - PeerId, shrink_dashmap, - stats_manager::{CounterHandle, LabelSet, LabelType, MetricName, StatsManager}, -}; +use crate::config::PeerId; +use crate::foundation::stats::{CounterHandle, LabelSet, LabelType, MetricName, StatsManager}; +use crate::packet::PacketType; +use crate::peers::util::shrink_dashmap; use crate::proto::peer_rpc::RoutePeerInfo; -use crate::tunnel::packet_def::PacketType; -pub(crate) const UNKNOWN_INSTANCE_ID: &str = "unknown"; +pub const UNKNOWN_INSTANCE_ID: &str = "unknown"; #[derive(Clone, Copy)] -pub(crate) enum InstanceLabelKind { +pub enum InstanceLabelKind { To, From, } @@ -31,55 +30,6 @@ impl TrafficCounters { } } -#[derive(Clone)] -pub(crate) struct AggregateTrafficMetrics { - tx: TrafficCounters, - rx: TrafficCounters, -} - -impl AggregateTrafficMetrics { - pub(crate) fn control(stats_mgr: Arc, network_name: String) -> Self { - Self::new( - stats_mgr, - network_name, - MetricName::TrafficControlBytesTx, - MetricName::TrafficControlPacketsTx, - MetricName::TrafficControlBytesRx, - MetricName::TrafficControlPacketsRx, - ) - } - - fn new( - stats_mgr: Arc, - network_name: String, - tx_bytes_metric: MetricName, - tx_packets_metric: MetricName, - rx_bytes_metric: MetricName, - rx_packets_metric: MetricName, - ) -> Self { - let label_set = - LabelSet::new().with_label_type(LabelType::NetworkName(network_name.clone())); - Self { - tx: TrafficCounters { - bytes: stats_mgr.get_counter(tx_bytes_metric, label_set.clone()), - packets: stats_mgr.get_counter(tx_packets_metric, label_set.clone()), - }, - rx: TrafficCounters { - bytes: stats_mgr.get_counter(rx_bytes_metric, label_set.clone()), - packets: stats_mgr.get_counter(rx_packets_metric, label_set), - }, - } - } - - pub(crate) fn record_tx(&self, bytes: u64) { - self.tx.add_sample(bytes); - } - - pub(crate) fn record_rx(&self, bytes: u64) { - self.rx.add_sample(bytes); - } -} - #[derive(Clone)] enum CachedPeerTrafficCounters { Unknown(TrafficCounters), @@ -99,7 +49,7 @@ impl CachedPeerTrafficCounters { } } -pub(crate) struct LogicalTrafficMetrics { +pub struct LogicalTrafficMetrics { stats_mgr: Arc, network_name: String, instance_bytes_metric: MetricName, @@ -110,7 +60,7 @@ pub(crate) struct LogicalTrafficMetrics { } impl LogicalTrafficMetrics { - pub(crate) fn new( + pub fn new( stats_mgr: Arc, network_name: String, total_bytes_metric: MetricName, @@ -135,12 +85,8 @@ impl LogicalTrafficMetrics { } } - pub(crate) async fn record_with_resolver( - &self, - peer_id: PeerId, - bytes: u64, - resolver: F, - ) where + pub async fn record_with_resolver(&self, peer_id: PeerId, bytes: u64, resolver: F) + where F: FnOnce() -> Fut, Fut: Future>, { @@ -186,22 +132,16 @@ impl LogicalTrafficMetrics { } } - pub(crate) fn remove_peer(&self, peer_id: PeerId) { + pub fn remove_peer(&self, peer_id: PeerId) { self.per_peer.remove(&peer_id); shrink_dashmap(&self.per_peer, None); } - pub(crate) fn clear_peer_cache(&self) { + pub fn clear_peer_cache(&self) { self.per_peer.clear(); shrink_dashmap(&self.per_peer, None); } - #[cfg(test)] - fn peer_cache_size(&self) -> usize { - self.per_peer.len() - } - - #[cfg(test)] fn contains_peer_cache(&self, peer_id: PeerId) -> bool { self.per_peer.contains_key(&peer_id) } @@ -226,12 +166,12 @@ impl LogicalTrafficMetrics { } #[derive(Clone, Copy, Debug, PartialEq, Eq)] -pub(crate) enum TrafficKind { +pub enum TrafficKind { Data, Control, } -pub(crate) fn traffic_kind(packet_type: u8) -> TrafficKind { +pub fn traffic_kind(packet_type: u8) -> TrafficKind { if packet_type == PacketType::Data as u8 || packet_type == PacketType::KcpSrc as u8 || packet_type == PacketType::KcpDst as u8 @@ -246,7 +186,7 @@ pub(crate) fn traffic_kind(packet_type: u8) -> TrafficKind { } } -pub(crate) fn is_relay_data_packet_type(packet_type: u8) -> bool { +pub fn is_relay_data_packet_type(packet_type: u8) -> bool { // Relay handshakes are control-plane setup; payload data is blocked by its // original packet type after the session exists. traffic_kind(packet_type) == TrafficKind::Data @@ -270,7 +210,7 @@ impl TrafficMetricGroup { type InstanceIdResolver = dyn Fn(PeerId) -> BoxFuture<'static, Option> + Send + Sync; -pub(crate) struct TrafficMetricRecorder { +pub struct TrafficMetricRecorder { my_peer_id: PeerId, tx_metrics: TrafficMetricGroup, rx_metrics: TrafficMetricGroup, @@ -278,7 +218,7 @@ pub(crate) struct TrafficMetricRecorder { } impl TrafficMetricRecorder { - pub(crate) fn new( + pub fn new( my_peer_id: PeerId, tx_data: Arc, tx_control: Arc, @@ -304,7 +244,7 @@ impl TrafficMetricRecorder { } } - pub(crate) async fn record_tx(&self, peer_id: PeerId, packet_type: u8, bytes: u64) { + pub async fn record_tx(&self, peer_id: PeerId, packet_type: u8, bytes: u64) { if peer_id == self.my_peer_id { return; } @@ -314,7 +254,7 @@ impl TrafficMetricRecorder { .await; } - pub(crate) async fn record_rx(&self, peer_id: PeerId, packet_type: u8, bytes: u64) { + pub async fn record_rx(&self, peer_id: PeerId, packet_type: u8, bytes: u64) { if peer_id == self.my_peer_id { return; } @@ -324,22 +264,21 @@ impl TrafficMetricRecorder { .await; } - pub(crate) fn remove_peer(&self, peer_id: PeerId) { + pub fn remove_peer(&self, peer_id: PeerId) { self.tx_metrics.data.remove_peer(peer_id); self.tx_metrics.control.remove_peer(peer_id); self.rx_metrics.data.remove_peer(peer_id); self.rx_metrics.control.remove_peer(peer_id); } - pub(crate) fn clear_peer_cache(&self) { + pub fn clear_peer_cache(&self) { self.tx_metrics.data.clear_peer_cache(); self.tx_metrics.control.clear_peer_cache(); self.rx_metrics.data.clear_peer_cache(); self.rx_metrics.control.clear_peer_cache(); } - #[cfg(test)] - pub(crate) fn contains_peer_cache(&self, peer_id: PeerId) -> bool { + pub fn contains_peer_cache(&self, peer_id: PeerId) -> bool { self.tx_metrics.data.contains_peer_cache(peer_id) || self.tx_metrics.control.contains_peer_cache(peer_id) || self.rx_metrics.data.contains_peer_cache(peer_id) @@ -351,7 +290,7 @@ impl TrafficMetricRecorder { } } -pub(crate) fn route_peer_info_instance_id(route_peer_info: &RoutePeerInfo) -> Option { +pub fn route_peer_info_instance_id(route_peer_info: &RoutePeerInfo) -> Option { let instance_id = route_peer_info.inst_id.as_ref()?; let instance_id: uuid::Uuid = (*instance_id).into(); if instance_id.is_nil() { @@ -364,7 +303,12 @@ pub(crate) fn route_peer_info_instance_id(route_peer_info: &RoutePeerInfo) -> Op #[cfg(test)] mod tests { use super::*; - use crate::common::stats_manager::LabelSet; + + impl LogicalTrafficMetrics { + fn peer_cache_size(&self) -> usize { + self.per_peer.len() + } + } fn network_labels(network_name: &str) -> LabelSet { LabelSet::new().with_label_type(LabelType::NetworkName(network_name.to_string())) diff --git a/easytier-core/src/peers/util.rs b/easytier-core/src/peers/util.rs new file mode 100644 index 00000000..48718bfe --- /dev/null +++ b/easytier-core/src/peers/util.rs @@ -0,0 +1,10 @@ +use std::hash::Hash; + +use dashmap::DashMap; + +pub(crate) fn shrink_dashmap(map: &DashMap, threshold: Option) { + let threshold = threshold.unwrap_or(16); + if map.capacity() - map.len() > threshold { + map.shrink_to_fit(); + } +} diff --git a/easytier-core/src/peers/whitelist.rs b/easytier-core/src/peers/whitelist.rs new file mode 100644 index 00000000..a7cb7e6a --- /dev/null +++ b/easytier-core/src/peers/whitelist.rs @@ -0,0 +1,18 @@ +//! Relay-network whitelist matching shared by the peer context and the +//! foreign-network manager. Kept in the peers kernel so both can depend on it +//! without depending on each other. + +pub(crate) fn check_network_in_relay_whitelist( + relay_network_whitelist: &str, + network_name: &str, +) -> Result<(), anyhow::Error> { + if relay_network_whitelist + .split(' ') + .map(wildmatch::WildMatch::new) + .any(|whitelist| whitelist.matches(network_name)) + { + Ok(()) + } else { + Err(anyhow::anyhow!("network {} not in whitelist", network_name)) + } +} diff --git a/easytier-core/src/process_runtime.rs b/easytier-core/src/process_runtime.rs new file mode 100644 index 00000000..1b84f24c --- /dev/null +++ b/easytier-core/src/process_runtime.rs @@ -0,0 +1,415 @@ +#![cfg_attr( + not(any(feature = "management", test, target_os = "wasi")), + allow(dead_code) +)] + +//! Host-domain portable resources shared by core instances. + +use std::{ + collections::HashMap, + sync::{Arc, Mutex}, +}; + +use crate::{ + connectivity::{ + manual::{ + ManualConnectorHost, ManualConnectorOptions, ManualTunnelConnector, + discovery::{CoreManualEndpointResolver, ManualEndpointDiscoveryConfig}, + }, + protocol::ClientProtocolUpgrader, + }, + host::dns::{DnsRecordResolver, DnsResolver}, + socket::SocketListener, + socket::{ring::RingSocketId, tcp::VirtualTcpSocketFactory}, + tunnel::{Tunnel, ring::RingTunnelRegistry}, +}; + +/// Owns portable resources whose identity is shared across core instances in +/// one native process or one instantiated WASI module. +/// +/// Native composition roots pass this handle around. The WASI lifecycle keeps +/// it module-local and never exposes it through the Go ABI. Neither host can +/// receive the internal managers it owns. +#[derive(Default)] +pub struct CoreProcessRuntime { + ring_registry: Arc, + protected_tcp_ports: Arc, +} + +#[derive(Default)] +pub(crate) struct ProtectedTcpPortRegistry { + ports: Mutex>, +} + +impl ProtectedTcpPortRegistry { + fn protect(self: &Arc, port: u16) -> ProtectedTcpPortLease { + let mut ports = self.ports.lock().unwrap(); + *ports.entry(port).or_default() += 1; + ProtectedTcpPortLease { + registry: self.clone(), + port, + } + } + + pub(crate) fn contains(&self, port: u16) -> bool { + self.ports.lock().unwrap().contains_key(&port) + } +} + +pub(crate) struct ProtectedTcpPortLease { + registry: Arc, + port: u16, +} + +impl Drop for ProtectedTcpPortLease { + fn drop(&mut self) { + let mut ports = self.registry.ports.lock().unwrap(); + let ref_count = ports + .get_mut(&self.port) + .expect("protected TCP port lease must have a registry entry"); + *ref_count -= 1; + if *ref_count == 0 { + ports.remove(&self.port); + } + } +} + +impl CoreProcessRuntime { + pub fn new() -> Arc { + Arc::new(Self::default()) + } + + pub(crate) fn ring_registry(&self) -> Arc { + self.ring_registry.clone() + } + + pub(crate) fn protected_tcp_ports(&self) -> Arc { + self.protected_tcp_ports.clone() + } + + pub(crate) fn protect_tcp_port(&self, port: u16) -> ProtectedTcpPortLease { + self.protected_tcp_ports.protect(port) + } + + /// Binds an application-level Ring listener without exposing the registry + /// that owns its process namespace. + pub fn bind_ring_tunnel( + &self, + local_id: RingSocketId, + ) -> anyhow::Result>>> { + Ok(Box::new(self.ring_registry.bind(local_id)?)) + } + + /// Connects an application-level Ring tunnel in this runtime's namespace. + pub fn connect_ring_tunnel(&self, remote_id: RingSocketId) -> anyhow::Result> { + Ok(self.ring_registry.connect(remote_id)?.into_tunnel()) + } + + pub fn manual_connector( + &self, + host: Arc, + dns: Arc, + dns_records: Arc, + protocol: Arc::Socket>>, + endpoint_discovery: ManualEndpointDiscoveryConfig, + options: ManualConnectorOptions, + ) -> ManualTunnelConnector + where + H: ManualConnectorHost, + { + let endpoint_resolver = Arc::new(CoreManualEndpointResolver::new( + host.clone(), + dns.clone(), + dns_records, + endpoint_discovery, + )); + ManualTunnelConnector::new(host, dns, endpoint_resolver, protocol, options) + .with_ring_registry(self.ring_registry()) + } +} + +#[cfg(test)] +mod tests { + use std::{ + io, + net::SocketAddr, + pin::Pin, + sync::Mutex, + task::{Context, Poll}, + }; + + use async_trait::async_trait; + use futures::{SinkExt, StreamExt}; + use tokio::io::{AsyncRead, AsyncWrite, ReadBuf}; + + use super::*; + use crate::{ + connectivity::{ + manual::ManualInterfaceAddrs, + protocol::{CoreClientProtocolConfig, CoreClientProtocolUpgrader}, + }, + host::dns::{DnsQuery, DnsRecordResolver, DnsSrvRecord}, + packet::ZCPacket, + socket::{ + IpVersion, NetNamespace, SocketContext, + tcp::{TcpConnectOptions, VirtualTcpSocket}, + udp::{UdpBindOptions, VirtualUdpSocket, VirtualUdpSocketFactory}, + }, + }; + + struct TestTcpSocket; + + #[test] + fn protected_tcp_port_leases_are_ref_counted_and_runtime_scoped() { + let runtime = CoreProcessRuntime::new(); + let other_runtime = CoreProcessRuntime::new(); + + let first = runtime.protect_tcp_port(15888); + let second = runtime.protect_tcp_port(15888); + assert!(runtime.protected_tcp_ports().contains(15888)); + assert!(!other_runtime.protected_tcp_ports().contains(15888)); + + drop(first); + assert!(runtime.protected_tcp_ports().contains(15888)); + + drop(second); + assert!(!runtime.protected_tcp_ports().contains(15888)); + } + + impl AsyncRead for TestTcpSocket { + fn poll_read( + self: Pin<&mut Self>, + _cx: &mut Context<'_>, + _buf: &mut ReadBuf<'_>, + ) -> Poll> { + Poll::Pending + } + } + + impl AsyncWrite for TestTcpSocket { + fn poll_write( + self: Pin<&mut Self>, + _cx: &mut Context<'_>, + _buf: &[u8], + ) -> Poll> { + Poll::Pending + } + + fn poll_flush(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll> { + Poll::Ready(Ok(())) + } + + fn poll_shutdown(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll> { + Poll::Ready(Ok(())) + } + } + + impl VirtualTcpSocket for TestTcpSocket { + fn local_addr(&self) -> io::Result { + Ok("127.0.0.1:1".parse().unwrap()) + } + + fn peer_addr(&self) -> io::Result { + Ok("127.0.0.1:2".parse().unwrap()) + } + } + + struct TestUdpSocket; + + #[async_trait] + impl VirtualUdpSocket for TestUdpSocket { + fn local_addr(&self) -> io::Result { + Ok("127.0.0.1:1".parse().unwrap()) + } + + async fn send_to(&self, _data: &[u8], _addr: SocketAddr) -> io::Result { + unreachable!("Ring connector must not use UDP") + } + + async fn recv_from(&self, _buf: &mut [u8]) -> io::Result<(usize, SocketAddr)> { + unreachable!("Ring connector must not use UDP") + } + } + + struct TestHost; + + #[async_trait] + impl VirtualTcpSocketFactory for TestHost { + type Socket = TestTcpSocket; + + async fn connect_tcp(&self, _options: TcpConnectOptions) -> anyhow::Result { + anyhow::bail!("Ring connector must not use TCP") + } + } + + #[async_trait] + impl VirtualUdpSocketFactory for TestHost { + type Socket = TestUdpSocket; + + async fn bind_udp(&self, _options: UdpBindOptions) -> anyhow::Result> { + anyhow::bail!("Ring connector must not use UDP") + } + } + + #[async_trait] + impl ManualConnectorHost for TestHost { + async fn local_addr_for_remote( + &self, + _remote_addr: SocketAddr, + _context: SocketContext, + ) -> anyhow::Result { + anyhow::bail!("Ring connector must not probe routes") + } + + async fn interface_addrs(&self) -> anyhow::Result { + anyhow::bail!("Ring connector must not collect interfaces") + } + } + + struct TestDns; + + #[async_trait] + impl DnsResolver for TestDns { + async fn resolve(&self, _query: DnsQuery) -> anyhow::Result> { + anyhow::bail!("Ring connector must not resolve DNS") + } + } + + #[async_trait] + impl DnsRecordResolver for TestDns { + async fn resolve_txt(&self, _query: DnsQuery) -> anyhow::Result { + anyhow::bail!("Ring connector must not resolve TXT records") + } + + async fn resolve_srv(&self, _query: DnsQuery) -> anyhow::Result> { + anyhow::bail!("Ring connector must not resolve SRV records") + } + } + + struct RecordingDnsRecords { + endpoint: String, + queries: Arc>>, + } + + #[async_trait] + impl DnsRecordResolver for RecordingDnsRecords { + async fn resolve_txt(&self, query: DnsQuery) -> anyhow::Result { + self.queries.lock().unwrap().push(query); + Ok(self.endpoint.clone()) + } + + async fn resolve_srv(&self, _query: DnsQuery) -> anyhow::Result> { + anyhow::bail!("TXT discovery must not resolve SRV records") + } + } + + fn manual_connector(runtime: &Arc) -> ManualTunnelConnector { + runtime.manual_connector( + Arc::new(TestHost), + Arc::new(TestDns), + Arc::new(TestDns), + Arc::new(CoreClientProtocolUpgrader::::new( + CoreClientProtocolConfig::default(), + )), + ManualEndpointDiscoveryConfig::default(), + ManualConnectorOptions::default(), + ) + } + + #[test] + fn instances_share_ring_state_only_through_the_process_runtime() { + let runtime = CoreProcessRuntime::new(); + + assert!(Arc::ptr_eq( + &runtime.ring_registry(), + &runtime.ring_registry() + )); + assert!(!Arc::ptr_eq( + &runtime.ring_registry(), + &CoreProcessRuntime::new().ring_registry() + )); + } + + #[tokio::test] + async fn one_shot_ring_connector_uses_its_process_runtime_namespace() { + let runtime = CoreProcessRuntime::new(); + let isolated = CoreProcessRuntime::new(); + let listener_id = uuid::Uuid::new_v4(); + let mut listener = runtime.bind_ring_tunnel(listener_id).unwrap(); + let url: url::Url = format!("ring://{listener_id}").parse().unwrap(); + + assert!( + manual_connector(&isolated) + .connect(url.clone(), IpVersion::Both) + .await + .is_err() + ); + + let client = manual_connector(&runtime) + .connect(url, IpVersion::Both) + .await + .unwrap(); + let server = listener.accept().await.unwrap(); + let (_client_stream, mut client_sink) = client.split(); + let (mut server_stream, _server_sink) = server.split(); + client_sink + .send(ZCPacket::new_with_payload(b"process-runtime")) + .await + .unwrap(); + + assert_eq!( + server_stream.next().await.unwrap().unwrap().payload(), + b"process-runtime" + ); + } + + #[cfg(feature = "endpoint-discovery")] + #[tokio::test] + async fn one_shot_connector_owns_endpoint_discovery_wiring() { + let runtime = CoreProcessRuntime::new(); + let listener_id = uuid::Uuid::new_v4(); + let mut listener = runtime.bind_ring_tunnel(listener_id).unwrap(); + let queries = Arc::new(Mutex::new(Vec::new())); + let dns_context = SocketContext::new() + .with_ip_version(IpVersion::V6) + .with_socket_mark(Some(42)) + .with_netns(Some(NetNamespace::new("discovery-netns"))); + let connector = runtime.manual_connector( + Arc::new(TestHost), + Arc::new(TestDns), + Arc::new(RecordingDnsRecords { + endpoint: format!("ring://{listener_id}"), + queries: queries.clone(), + }), + Arc::new(CoreClientProtocolUpgrader::::new( + CoreClientProtocolConfig::default(), + )), + ManualEndpointDiscoveryConfig { + dns_record_context: dns_context.clone(), + ..Default::default() + }, + ManualConnectorOptions::default(), + ); + + let client = connector + .connect("txt://bootstrap.example".parse().unwrap(), IpVersion::Both) + .await + .unwrap(); + let server = listener.accept().await.unwrap(); + + assert_eq!( + *queries.lock().unwrap(), + vec![DnsQuery::new("bootstrap.example", dns_context)] + ); + let (_client_stream, mut client_sink) = client.split(); + let (mut server_stream, _server_sink) = server.split(); + client_sink + .send(ZCPacket::new_with_payload(b"endpoint-discovery")) + .await + .unwrap(); + assert_eq!( + server_stream.next().await.unwrap().unwrap().payload(), + b"endpoint-discovery" + ); + } +} diff --git a/easytier/src/proto/rpc_impl/bidirect.rs b/easytier-core/src/rpc/bidirect.rs similarity index 67% rename from easytier/src/proto/rpc_impl/bidirect.rs rename to easytier-core/src/rpc/bidirect.rs index e93babcb..2d61b6c9 100644 --- a/easytier/src/proto/rpc_impl/bidirect.rs +++ b/easytier-core/src/rpc/bidirect.rs @@ -1,16 +1,22 @@ -use std::sync::{Arc, Mutex, atomic::AtomicBool}; +use std::sync::{ + Arc, Mutex, + atomic::{AtomicBool, Ordering}, +}; use futures::{SinkExt as _, StreamExt}; -use guarden::defer; -use tokio::{task::JoinSet, time::timeout}; +use tokio::task::JoinSet; use crate::{ + foundation::{ + stats::{ArcRpcMetrics, RpcMetricsProvider}, + time::timeout, + }, + packet::PacketType, proto::rpc_types::error::Error, - tunnel::{Tunnel, packet_def::PacketType, ring::create_ring_tunnel_pair}, + tunnel::{Tunnel, ring::create_ring_tunnel_pair}, }; use super::{client::Client, server::Server, service_registry::ServiceRegistry}; -use crate::common::stats_manager::StatsManager; pub struct BidirectRpcManager { rpc_client: Client, @@ -45,7 +51,10 @@ impl BidirectRpcManager { } } - pub fn new_with_stats_manager(stats_manager: Arc) -> Self { + pub fn new_with_stats_manager(stats_manager: T) -> Self + where + T: Clone + RpcMetricsProvider, + { Self { rpc_client: Client::new_with_stats_manager(stats_manager.clone()), rpc_server: Server::new_with_registry_and_stats_manager( @@ -62,6 +71,23 @@ impl BidirectRpcManager { } } + pub fn new_with_metrics(metrics: ArcRpcMetrics) -> Self { + Self { + rpc_client: Client::new_with_metrics(metrics.clone()), + rpc_server: Server::new_with_registry_and_metrics( + Arc::new(ServiceRegistry::new()), + metrics, + ), + + rx_timeout: None, + error: Arc::new(Mutex::new(None)), + tunnel: Mutex::new(None), + running: Arc::new(AtomicBool::new(false)), + + tasks: Mutex::new(None), + } + } + pub fn set_rx_timeout(mut self, timeout: Option) -> Self { self.rx_timeout = timeout; self @@ -74,11 +100,19 @@ impl BidirectRpcManager { } pub fn run_with_tunnel(&self, inner: Box) { + let tunnel_info = inner.info(); + self.run_with_tunnel_info(inner, tunnel_info); + } + + pub(crate) fn run_with_tunnel_info( + &self, + inner: Box, + tunnel_info: Option, + ) { let mut tasks = JoinSet::new(); self.rpc_client.run(); - self.rpc_server.run(); - self.running - .store(true, std::sync::atomic::Ordering::Relaxed); + self.rpc_server.run_with_tunnel_info(tunnel_info); + self.running.store(true, Ordering::Relaxed); let (server_tx, mut server_rx) = ( self.rpc_server.get_transport_sink(), @@ -95,9 +129,6 @@ impl BidirectRpcManager { let e_clone = self.error.clone(); let r_clone = self.running.clone(); tasks.spawn(async move { - defer! { - r_clone.store(false, std::sync::atomic::Ordering::Relaxed); - } loop { let packet = tokio::select! { Some(Ok(packet)) = server_rx.next() => { @@ -110,6 +141,7 @@ impl BidirectRpcManager { } else => { tracing::warn!("rpc transport read aborted, exiting"); + r_clone.store(false, Ordering::Relaxed); break; } }; @@ -117,6 +149,8 @@ impl BidirectRpcManager { if let Err(e) = inner_tx.send(packet).await { tracing::error!(error = ?e, "send to peer failed"); e_clone.lock().unwrap().replace(Error::from(e)); + r_clone.store(false, Ordering::Relaxed); + break; } } }); @@ -125,15 +159,13 @@ impl BidirectRpcManager { let e_clone = self.error.clone(); let r_clone = self.running.clone(); tasks.spawn(async move { - defer! { - r_clone.store(false, std::sync::atomic::Ordering::Relaxed); - } loop { let ret = if let Some(recv_timeout) = recv_timeout { match timeout(recv_timeout, inner_rx.next()).await { Ok(ret) => ret, Err(e) => { e_clone.lock().unwrap().replace(e.into()); + r_clone.store(false, Ordering::Relaxed); break; } } @@ -146,11 +178,13 @@ impl BidirectRpcManager { Some(Err(e)) => { tracing::error!(error = ?e, "recv from peer failed"); e_clone.lock().unwrap().replace(Error::from(e)); + r_clone.store(false, Ordering::Relaxed); break; } None => { tracing::warn!("peer rpc transport read aborted, exiting"); e_clone.lock().unwrap().replace(Error::Shutdown); + r_clone.store(false, Ordering::Relaxed); break; } }; @@ -160,10 +194,20 @@ impl BidirectRpcManager { continue; }; if peer_manager_header.packet_type == PacketType::RpcReq as u8 { - server_tx.send(o).await.unwrap(); + if let Err(e) = server_tx.send(o).await { + tracing::error!(error = ?e, "send rpc request to server failed"); + e_clone.lock().unwrap().replace(Error::from(e)); + r_clone.store(false, Ordering::Relaxed); + break; + } continue; } else if peer_manager_header.packet_type == PacketType::RpcResp as u8 { - client_tx.send(o).await.unwrap(); + if let Err(e) = client_tx.send(o).await { + tracing::error!(error = ?e, "send rpc response to client failed"); + e_clone.lock().unwrap().replace(Error::from(e)); + r_clone.store(false, Ordering::Relaxed); + break; + } continue; } } @@ -181,6 +225,10 @@ impl BidirectRpcManager { } pub async fn stop(&self) { + self.running.store(false, Ordering::Relaxed); + self.tunnel.lock().unwrap().take(); + self.rpc_client.stop().await; + self.rpc_server.stop().await; let Some(mut tasks) = self.tasks.lock().unwrap().take() else { return; }; @@ -197,12 +245,15 @@ impl BidirectRpcManager { return; }; while tasks.join_next().await.is_some() { - // when any task is done, abort all tasks tasks.abort_all(); } + self.running.store(false, Ordering::Relaxed); + self.tunnel.lock().unwrap().take(); + self.rpc_client.stop().await; + self.rpc_server.stop().await; } pub fn is_running(&self) -> bool { - self.running.load(std::sync::atomic::Ordering::Relaxed) + self.running.load(Ordering::Relaxed) } } diff --git a/easytier/src/proto/rpc_impl/client.rs b/easytier-core/src/rpc/client.rs similarity index 59% rename from easytier/src/proto/rpc_impl/client.rs rename to easytier-core/src/rpc/client.rs index 13302557..52d03abf 100644 --- a/easytier/src/proto/rpc_impl/client.rs +++ b/easytier-core/src/rpc/client.rs @@ -1,44 +1,45 @@ -use std::marker::PhantomData; -use std::pin::Pin; -use std::sync::{Arc, Mutex}; +use std::{ + marker::PhantomData, + pin::Pin, + sync::{Arc, LazyLock, Mutex, atomic::Ordering}, +}; +use atomic_shim::AtomicI64; use bytes::Bytes; use dashmap::DashMap; -use guarden::defer; +use futures::StreamExt; use prost::Message; -use quanta::Instant; -use tokio::sync::mpsc; -use tokio::task::JoinSet; -use tokio::time::timeout; -use tokio_stream::StreamExt; +use tokio::{sync::mpsc, task::JoinSet}; -use crate::common::{ - PeerId, shrink_dashmap, - stats_manager::{LabelSet, LabelType, MetricName, StatsManager}, -}; -use crate::proto::common::{ - CompressionAlgoPb, RpcCompressionInfo, RpcDescriptor, RpcPacket, RpcRequest, RpcResponse, -}; -use crate::proto::rpc_impl::packet::{ - BuildRpcPacketArgs, build_rpc_packet, compress_packet, decompress_packet, -}; -use crate::proto::rpc_types::controller::Controller; -use crate::proto::rpc_types::descriptor::MethodDescriptor; -use crate::proto::rpc_types::{ - __rt::RpcClientFactory, descriptor::ServiceDescriptor, handler::Handler, +use crate::{ + config::PeerId, + foundation::{ + stats::{ArcRpcMetrics, RpcMetricLabels, RpcMetricsProvider}, + time::timeout, + }, + proto::{ + common::{ + CompressionAlgoPb, RpcCompressionInfo, RpcDescriptor, RpcPacket, RpcRequest, + RpcResponse, + }, + rpc_types::controller::Controller, + rpc_types::descriptor::MethodDescriptor, + rpc_types::error::{Error, Result}, + rpc_types::{__rt::RpcClientFactory, descriptor::ServiceDescriptor, handler::Handler}, + }, + rpc::packet::{ + BuildRpcPacketArgs, PacketMerger, build_rpc_packet, compress_packet, decompress_packet, + }, + tunnel::{ + Tunnel, TunnelError, ZCPacketStream, + mpsc::{MpscTunnel, MpscTunnelSender}, + ring::create_ring_tunnel_pair, + }, }; -use crate::proto::rpc_types::error::Result; -use crate::tunnel::mpsc::{MpscTunnel, MpscTunnelSender}; -use crate::tunnel::packet_def::ZCPacket; -use crate::tunnel::ring::create_ring_tunnel_pair; -use crate::tunnel::{Tunnel, TunnelError, ZCPacketStream}; - -use super::packet::PacketMerger; use super::{RpcTransactId, Transport}; -static CUR_TID: once_cell::sync::Lazy = - once_cell::sync::Lazy::new(|| atomic_shim::AtomicI64::new(rand::random())); +static CUR_TID: LazyLock = LazyLock::new(|| AtomicI64::new(rand::random())); type RpcPacketSender = mpsc::UnboundedSender; type RpcPacketReceiver = mpsc::UnboundedReceiver; @@ -53,7 +54,7 @@ struct InflightRequestKey { struct InflightRequest { sender: RpcPacketSender, merger: PacketMerger, - start_time: Instant, + start_time: std::time::Instant, } impl std::fmt::Debug for InflightRequest { @@ -65,15 +66,29 @@ impl std::fmt::Debug for InflightRequest { } } +struct InflightCleanup { + table: InflightRequestTable, + key: InflightRequestKey, +} + +impl Drop for InflightCleanup { + fn drop(&mut self) { + self.table.remove(&self.key); + if self.table.capacity() - self.table.len() > 4 { + self.table.shrink_to_fit(); + } + } +} + #[derive(Debug, Clone, Default)] -pub(crate) struct PeerInfo { +pub struct PeerInfo { pub peer_id: PeerId, pub compression_info: RpcCompressionInfo, - pub last_active: Option, + pub last_active: Option, } type InflightRequestTable = Arc>; -pub(crate) type PeerInfoTable = Arc>; +pub type PeerInfoTable = Arc>; pub struct Client { mpsc: Mutex>>, @@ -81,7 +96,7 @@ pub struct Client { inflight_requests: InflightRequestTable, peer_info: PeerInfoTable, tasks: Mutex>, - stats_manager: Option>, + metrics: Option, } impl Default for Client { @@ -99,14 +114,23 @@ impl Client { inflight_requests: Arc::new(DashMap::new()), peer_info: Arc::new(DashMap::new()), tasks: Mutex::new(JoinSet::new()), - stats_manager: None, + metrics: None, } } - pub fn new_with_stats_manager(stats_manager: Arc) -> Self { - let mut ret = Self::new(); - ret.stats_manager = Some(stats_manager); - ret + pub fn new_with_stats_manager(stats_manager: T) -> Self + where + T: RpcMetricsProvider, + { + let mut client = Self::new(); + client.metrics = stats_manager.into_rpc_metrics(); + client + } + + pub fn new_with_metrics(metrics: ArcRpcMetrics) -> Self { + let mut client = Self::new(); + client.metrics = Some(metrics); + client } pub fn get_transport_sink(&self) -> MpscTunnelSender { @@ -123,8 +147,8 @@ impl Client { let peer_infos = self.peer_info.clone(); tasks.spawn(async move { loop { - tokio::time::sleep(std::time::Duration::from_secs(30)).await; - let now = Instant::now(); + crate::foundation::time::sleep(std::time::Duration::from_secs(30)).await; + let now = std::time::Instant::now(); peer_infos.retain(|_, v| { if let Some(last_active) = v.last_active { return now.duration_since(last_active) @@ -140,11 +164,14 @@ impl Client { let inflight_requests = self.inflight_requests.clone(); tasks.spawn(async move { while let Some(packet) = rx.next().await { - if let Err(err) = packet { - tracing::error!(?err, "Failed to receive packet"); - continue; - } - let packet = match RpcPacket::decode(packet.unwrap().payload()) { + let packet = match packet { + Err(err) => { + tracing::error!(?err, "Failed to receive packet"); + continue; + } + Ok(packet) => packet, + }; + let packet = match RpcPacket::decode(packet.payload()) { Err(err) => { tracing::error!(?err, "Failed to decode packet"); continue; @@ -177,7 +204,15 @@ impl Client { let ret = inflight_request.merger.feed(packet); match ret { Ok(Some(rpc_packet)) => { - inflight_request.sender.send(rpc_packet).unwrap(); + if let Err(err) = inflight_request.sender.send(rpc_packet) { + tracing::warn!( + ?err, + ?key, + "RPC response receiver is gone, removing inflight request" + ); + drop(inflight_request); + inflight_requests.remove(&key); + } } Ok(None) => {} Err(err) => { @@ -202,14 +237,14 @@ impl Client { zc_packet_sender: MpscTunnelSender, inflight_requests: InflightRequestTable, peer_info: PeerInfoTable, - stats_manager: Option>, + metrics: Option, _phan: PhantomData, } impl HandlerImpl { async fn do_rpc( &self, - packets: Vec, + packets: Vec, rx: &mut RpcPacketReceiver, ) -> Result { for packet in packets { @@ -229,10 +264,10 @@ impl Client { &self, mut ctrl: Self::Controller, method: ::Method, - input: bytes::Bytes, - ) -> Result { - let start_time = Instant::now(); - let transaction_id = CUR_TID.fetch_add(1, std::sync::atomic::Ordering::Relaxed); + input: Bytes, + ) -> Result { + let start_time = std::time::Instant::now(); + let transaction_id = CUR_TID.fetch_add(1, Ordering::Relaxed); let (tx, mut rx) = mpsc::unbounded_channel(); let key = InflightRequestKey { from_peer_id: self.from_peer_id, @@ -240,14 +275,14 @@ impl Client { transaction_id, }; let desc = self.service_descriptor(); - let labels = LabelSet::new() - .with_label_type(LabelType::NetworkName(self.domain_name.to_string())) - .with_label_type(LabelType::SrcPeerId(self.from_peer_id)) - .with_label_type(LabelType::DstPeerId(self.to_peer_id)) - .with_label_type(LabelType::ServiceName(desc.name().to_string())) - .with_label_type(LabelType::MethodName(method.name().to_string())); + let labels = RpcMetricLabels { + network_name: self.domain_name.clone(), + src_peer_id: self.from_peer_id, + dst_peer_id: self.to_peer_id, + service_name: desc.name().to_string(), + method_name: method.name().to_string(), + }; - defer!(self.inflight_requests.remove(&key); shrink_dashmap(&self.inflight_requests, Some(4));); self.inflight_requests.insert( key.clone(), InflightRequest { @@ -256,12 +291,13 @@ impl Client { start_time, }, ); + let _cleanup = InflightCleanup { + table: self.inflight_requests.clone(), + key: key.clone(), + }; - // Record RPC client TX stats - if let Some(ref stats_manager) = self.stats_manager { - stats_manager - .get_counter(MetricName::PeerRpcClientTx, labels.clone()) - .inc(); + if let Some(metrics) = &self.metrics { + metrics.client_tx(&labels); } let rpc_desc = RpcDescriptor { @@ -307,7 +343,31 @@ impl Client { }, }); let timeout_dur = std::time::Duration::from_millis(ctrl.timeout_ms() as u64); - let mut rpc_packet = timeout(timeout_dur, self.do_rpc(packets, &mut rx)).await??; + let rpc_ret = timeout(timeout_dur, self.do_rpc(packets, &mut rx)).await; + let mut rpc_packet = match rpc_ret { + Ok(Ok(packet)) => packet, + Ok(Err(err)) => { + if let Some(metrics) = &self.metrics { + metrics.client_error( + &labels, + Some(format!("{:?}", err)), + start_time.elapsed().as_millis() as u64, + ); + } + return Err(err); + } + Err(err) => { + let err = Error::from(err); + if let Some(metrics) = &self.metrics { + metrics.client_error( + &labels, + Some(format!("{:?}", err)), + start_time.elapsed().as_millis() as u64, + ); + } + return Err(err); + } + }; if let Some(compression_info) = rpc_packet.compression_info { self.peer_info.insert( @@ -315,7 +375,7 @@ impl Client { PeerInfo { peer_id: self.to_peer_id, compression_info, - last_active: Some(Instant::now()), + last_active: Some(std::time::Instant::now()), }, ); @@ -328,21 +388,12 @@ impl Client { let rpc_resp = RpcResponse::decode(Bytes::from(rpc_packet.body))?; if let Some(err) = &rpc_resp.error { - // Record RPC error stats - if let Some(ref stats_manager) = self.stats_manager { - let labels = labels - .clone() - .with_label_type(LabelType::ErrorType(format!("{:?}", err.error_kind))) - .with_label_type(LabelType::Status("error".to_string())); - - stats_manager - .get_counter(MetricName::PeerRpcErrors, labels.clone()) - .inc(); - - let duration_ms = start_time.elapsed().as_millis() as u64; - stats_manager - .get_counter(MetricName::PeerRpcDuration, labels) - .add(duration_ms); + if let Some(metrics) = &self.metrics { + metrics.client_error( + &labels, + Some(format!("{:?}", err.error_kind)), + start_time.elapsed().as_millis() as u64, + ); } return Err(err.into()); } @@ -350,20 +401,8 @@ impl Client { let raw_output = Bytes::from(rpc_resp.response); ctrl.set_raw_output(raw_output.clone()); - // Record RPC client RX and duration stats - if let Some(ref stats_manager) = self.stats_manager { - let labels = labels - .clone() - .with_label_type(LabelType::Status("success".to_string())); - - stats_manager - .get_counter(MetricName::PeerRpcClientRx, labels.clone()) - .inc(); - - let duration_ms = start_time.elapsed().as_millis() as u64; - stats_manager - .get_counter(MetricName::PeerRpcDuration, labels) - .add(duration_ms); + if let Some(metrics) = &self.metrics { + metrics.client_rx(&labels, start_time.elapsed().as_millis() as u64); } Ok(raw_output) @@ -377,16 +416,35 @@ impl Client { zc_packet_sender: self.mpsc.lock().unwrap().get_sink(), inflight_requests: self.inflight_requests.clone(), peer_info: self.peer_info.clone(), - stats_manager: self.stats_manager.clone(), + metrics: self.metrics.clone(), _phan: PhantomData, }) } - pub fn inflight_count(&self) -> usize { - self.inflight_requests.len() - } - - pub(crate) fn peer_info_table(&self) -> PeerInfoTable { - self.peer_info.clone() + pub async fn stop(&self) { + self.transport.lock().unwrap().close(); + let mut tasks = { + let mut task_slot = self.tasks.lock().unwrap(); + std::mem::replace(&mut *task_slot, JoinSet::new()) + }; + tasks.abort_all(); + while tasks.join_next().await.is_some() {} + } +} + +#[cfg(any(test, feature = "test-utils"))] +mod test_utils { + use super::{Client, PeerInfoTable}; + + impl Client { + #[doc(hidden)] + pub fn inflight_count(&self) -> usize { + self.inflight_requests.len() + } + + #[doc(hidden)] + pub fn peer_info_table(&self) -> PeerInfoTable { + self.peer_info.clone() + } } } diff --git a/easytier/src/proto/rpc_impl/mod.rs b/easytier-core/src/rpc/mod.rs similarity index 75% rename from easytier/src/proto/rpc_impl/mod.rs rename to easytier-core/src/rpc/mod.rs index 09366a36..77baf687 100644 --- a/easytier/src/proto/rpc_impl/mod.rs +++ b/easytier-core/src/rpc/mod.rs @@ -1,6 +1,6 @@ use crate::tunnel::{Tunnel, mpsc::MpscTunnel}; -pub type RpcController = super::rpc_types::controller::BaseController; +pub type RpcController = crate::proto::rpc_types::controller::BaseController; pub mod bidirect; pub mod client; diff --git a/easytier/src/proto/rpc_impl/packet.rs b/easytier-core/src/rpc/packet.rs similarity index 51% rename from easytier/src/proto/rpc_impl/packet.rs rename to easytier-core/src/rpc/packet.rs index 9f095c89..6e06019c 100644 --- a/easytier/src/proto/rpc_impl/packet.rs +++ b/easytier-core/src/rpc/packet.rs @@ -1,36 +1,37 @@ use prost::{Message as _, length_delimiter_len}; -use quanta::Instant; - use crate::{ - common::{PeerId, compressor::DefaultCompressor}, + config::PeerId, + packet::{CompressorAlgo, PacketType, TAIL_RESERVED_SIZE, ZCPacket, ZCPacketType}, proto::{ common::{CompressionAlgoPb, RpcCompressionInfo, RpcDescriptor, RpcPacket}, rpc_types::error::Error, }, - tunnel::packet_def::{CompressorAlgo, PacketType, TAIL_RESERVED_SIZE, ZCPacket, ZCPacketType}, }; use super::RpcTransactId; -// Budget the final UDP payload size on the wire for peer RPC over `udp://`. -// This includes EasyTier's UDP tunnel header, peer header, and reserved tail -// space for encryption/compression metadata, but excludes the outer IP header. const RPC_PACKET_UDP_PAYLOAD_BUDGET: usize = 1300; pub async fn compress_packet( accepted_compression_algo: CompressionAlgoPb, content: &[u8], ) -> Result<(Vec, CompressionAlgoPb), Error> { - let compressor = DefaultCompressor::new(); - let algo = accepted_compression_algo - .try_into() + let algo = CompressorAlgo::try_from(accepted_compression_algo) + .ok() + .filter(|algo| algo.is_available()) .unwrap_or(CompressorAlgo::None); - let compressed = compressor.compress_raw(content, algo).await?; + let compressed = crate::packet::compressor::DefaultCompressor::new() + .compress_raw(content, algo) + .await + .map_err(Error::from)?; if compressed.len() >= content.len() { Ok((content.to_vec(), CompressionAlgoPb::None)) } else { - Ok((compressed, algo.try_into().unwrap())) + Ok(( + compressed, + CompressionAlgoPb::try_from(algo).expect("CompressorAlgo should map to protobuf"), + )) } } @@ -38,16 +39,17 @@ pub async fn decompress_packet( compression_algo: CompressionAlgoPb, content: &[u8], ) -> Result, Error> { - let compressor = DefaultCompressor::new(); - let algo = compression_algo.try_into()?; - let decompressed = compressor.decompress_raw(content, algo).await?; - Ok(decompressed) + let algo = CompressorAlgo::try_from(compression_algo).map_err(anyhow::Error::from)?; + crate::packet::compressor::DefaultCompressor::new() + .decompress_raw(content, algo) + .await + .map_err(Error::from) } pub(crate) struct PacketMerger { first_piece: Option, pieces: Vec, - last_updated: Instant, + last_updated: std::time::Instant, } impl Default for PacketMerger { @@ -61,7 +63,7 @@ impl PacketMerger { Self { first_piece: None, pieces: Vec::new(), - last_updated: Instant::now(), + last_updated: std::time::Instant::now(), } } @@ -71,19 +73,16 @@ impl PacketMerger { } for p in &self.pieces { - // some piece is missing if p.total_pieces == 0 { return None; } } - // all pieces are received let mut body = Vec::new(); for p in &self.pieces { body.extend_from_slice(&p.body); } - // only the first packet contains the complete info let mut tmpl_packet = self.pieces[0].clone(); tmpl_packet.total_pieces = 1; tmpl_packet.piece_idx = 0; @@ -96,7 +95,6 @@ impl PacketMerger { let total_pieces = rpc_packet.total_pieces; let piece_idx = rpc_packet.piece_idx; - // for compatibility with old version if total_pieces == 0 && piece_idx == 0 { return Ok(Some(rpc_packet)); } @@ -107,7 +105,6 @@ impl PacketMerger { )); } - // about 32MB max size if total_pieces > 32 * 1024 || total_pieces == 0 { return Err(Error::MalformatRpcPacket(format!( "total_pieces is invalid: {}", @@ -134,12 +131,12 @@ impl PacketMerger { .resize(total_pieces as usize, Default::default()); self.pieces[piece_idx as usize] = rpc_packet; - self.last_updated = Instant::now(); + self.last_updated = std::time::Instant::now(); Ok(self.try_merge_pieces()) } - pub(crate) fn last_updated(&self) -> Instant { + pub(crate) fn last_updated(&self) -> std::time::Instant { self.last_updated } } @@ -155,28 +152,14 @@ pub struct BuildRpcPacketArgs<'a> { pub compression_info: RpcCompressionInfo, } -// Fixed transport overhead for peer RPC carried by EasyTier's UDP tunnel: -// -// UDP payload budget -// +-------------------------------------------------------------------------+ -// | EasyTier UDP tunnel hdr | PeerManager hdr | RpcPacket bytes | tail room | -// +-------------------------------------------------------------------------+ -// |<------ ZCPacketType::UDP payload_offset ------>|<-- TAIL_RESERVED_SIZE -->| -// -// `udp_rpc_tunnel_overhead()` is everything except `RpcPacket bytes`. fn udp_rpc_tunnel_overhead() -> usize { ZCPacketType::UDP.get_packet_offsets().payload_offset + TAIL_RESERVED_SIZE } -// Maximum encoded RpcPacket size we can admit before adding it to a UDP tunnel. -// This budget excludes the outer UDP/IP headers because the caller only controls -// the EasyTier payload carried inside the UDP datagram. fn max_rpc_packet_encoded_len_for_udp() -> usize { RPC_PACKET_UDP_PAYLOAD_BUDGET.saturating_sub(udp_rpc_tunnel_overhead()) } -// Build one logical RpcPacket piece. This is reused both for the actual output -// packets and for sizing templates that estimate worst-case protobuf overhead. fn build_rpc_piece( args: &BuildRpcPacketArgs<'_>, total_pieces: u32, @@ -189,7 +172,6 @@ fn build_rpc_piece( descriptor: if piece_idx == 0 || args.compression_info.algo == CompressionAlgoPb::None as i32 { - // old version must have descriptor on every piece Some(args.rpc_desc.clone()) } else { None @@ -217,8 +199,6 @@ fn pick_piece_len_for_budget( return 0; } - // Minimum non-empty body field encoding cost: - // body tag (1 byte) + body length (1 byte) + body data (1 byte) if base_encoded_len_without_body + 3 > max_encoded_len { tracing::warn!( base_encoded_len_without_body, @@ -228,65 +208,25 @@ fn pick_piece_len_for_budget( return 1; } - // `budget` is what remains for the protobuf `body` field after all fixed - // RpcPacket metadata has been accounted for. let budget = max_encoded_len - base_encoded_len_without_body; - // Reserve the bytes field wrapper conservatively, then use the rest for - // the body itself. - // - // Encoded RpcPacket layout relevant to `body`: - // - // +------------------------------- max_encoded_len -------------------------------+ - // | fixed RpcPacket fields | body tag (1B) | body len varint (worst-case) | body | - // +--------------------------------------------------------------------------- --+ - // ^ ^ - // | `- reserve by using the varint width of `budget` - // `- base_encoded_len_without_body - // - // This is intentionally conservative. A few bytes may be left unused, but - // every piece stays within the UDP payload budget without iterative sizing. let reserved_for_body_header = 1 + length_delimiter_len(budget); remaining .min(budget.saturating_sub(reserved_for_body_header)) .max(1) } -// Pre-split the raw RPC content using conservative worst-case protobuf sizing. -// We compute separate base sizes for the first piece and later pieces because -// only the first piece carries `compression_info`, and old compatibility rules -// may also force `descriptor` to appear on every piece. -// -// Split flow: -// -// raw RPC content -// +--------------------------------------------------------------+ -// | args.content | -// +--------------------------------------------------------------+ -// | first piece uses first_piece_base_len -// | later pieces use other_piece_base_len -// v -// +-----------+-----------+-----------+----- ... -// | offset,len| offset,len| offset,len| -// +-----------+-----------+-----------+----- ... -// -// The result is only a slicing plan. Actual RpcPacket objects are built later -// with the real `total_pieces`. fn split_rpc_content_for_udp_budget(args: &BuildRpcPacketArgs<'_>) -> Vec<(usize, usize)> { if args.content.is_empty() { return vec![(0, 0)]; } let max_encoded_len = max_rpc_packet_encoded_len_for_udp().max(1); - // Use the worst-case varint width for piece counters so the budget remains - // valid without iterating on `total_pieces`/`piece_idx`. let first_piece_base_len = build_rpc_piece(args, u32::MAX, 0, &[]).encoded_len(); let other_piece_base_len = build_rpc_piece(args, u32::MAX, u32::MAX, &[]).encoded_len(); let mut pieces = Vec::new(); let mut offset = 0usize; while offset < args.content.len() { - // First and subsequent pieces have different metadata shapes, so they - // use different fixed-size templates. let base_len = if pieces.is_empty() { first_piece_base_len } else { @@ -301,9 +241,6 @@ fn split_rpc_content_for_udp_budget(args: &BuildRpcPacketArgs<'_>) -> Vec<(usize pieces } -// Build the final transport packets after the payload has been split. We do the -// actual `total_pieces` assignment only here so the wire packet stays accurate, -// while the earlier sizing step remains simple and conservatively safe. pub fn build_rpc_packet(args: BuildRpcPacketArgs<'_>) -> Vec { let mut ret = Vec::new(); let pieces = split_rpc_content_for_udp_budget(&args); @@ -332,61 +269,19 @@ pub fn build_rpc_packet(args: BuildRpcPacketArgs<'_>) -> Vec { ret } -#[cfg(test)] +#[cfg(all(test, not(feature = "zstd")))] mod tests { - use super::*; + use crate::proto::common::CompressionAlgoPb; - fn build_test_args<'a>( - content: &'a [u8], - compression_algo: CompressionAlgoPb, - ) -> BuildRpcPacketArgs<'a> { - BuildRpcPacketArgs { - from_peer: 11, - to_peer: 22, - rpc_desc: RpcDescriptor { - domain_name: "very-long-domain-name-for-rpc-packet-budget-check".repeat(2), - proto_name: "extremely.verbose.proto.name.for.rpc.packet.tests".repeat(2), - service_name: "LargeMetadataServiceForRpcPacketBudget".repeat(2), - method_index: 7, - }, - transaction_id: 33, - is_req: true, - content, - trace_id: 44, - compression_info: RpcCompressionInfo { - algo: compression_algo.into(), - accepted_algo: CompressionAlgoPb::Zstd.into(), - }, - } - } + use super::compress_packet; - fn udp_packet_size_after_tail(packet: &ZCPacket) -> usize { - ZCPacketType::UDP.get_packet_offsets().payload_offset - + packet.payload_len() - + TAIL_RESERVED_SIZE - } + #[tokio::test] + async fn compression_negotiation_falls_back_when_zstd_is_unavailable() { + let (content, algorithm) = compress_packet(CompressionAlgoPb::Zstd, b"rpc body") + .await + .unwrap(); - #[test] - fn build_rpc_packet_respects_udp_budget_with_large_metadata() { - let content = vec![0x5a; 4096]; - let packets = build_rpc_packet(build_test_args(&content, CompressionAlgoPb::None)); - - assert!(packets.len() > 1); - for packet in packets { - assert!( - udp_packet_size_after_tail(&packet) <= RPC_PACKET_UDP_PAYLOAD_BUDGET, - "packet size {} exceeded budget {}", - udp_packet_size_after_tail(&packet), - RPC_PACKET_UDP_PAYLOAD_BUDGET - ); - } - } - - #[test] - fn build_rpc_packet_respects_udp_budget_for_empty_payload() { - let packets = build_rpc_packet(build_test_args(&[], CompressionAlgoPb::Zstd)); - - assert_eq!(1, packets.len()); - assert!(udp_packet_size_after_tail(&packets[0]) <= RPC_PACKET_UDP_PAYLOAD_BUDGET); + assert_eq!(content, b"rpc body"); + assert_eq!(algorithm, CompressionAlgoPb::None); } } diff --git a/easytier/src/proto/rpc_impl/server.rs b/easytier-core/src/rpc/server.rs similarity index 61% rename from easytier/src/proto/rpc_impl/server.rs rename to easytier-core/src/rpc/server.rs index fa43c24e..dc365952 100644 --- a/easytier/src/proto/rpc_impl/server.rs +++ b/easytier-core/src/rpc/server.rs @@ -1,28 +1,31 @@ use std::{ pin::Pin, - sync::{Arc, Mutex}, + sync::{ + Arc, Mutex, + atomic::{AtomicBool, Ordering}, + }, }; use bytes::Bytes; use dashmap::DashMap; +use futures::StreamExt; use prost::Message; -use quanta::Instant; -use tokio::{task::JoinSet, time::timeout}; -use tokio_stream::StreamExt; +use tokio::task::JoinSet; use crate::{ - common::{ - PeerId, join_joinset_background, - stats_manager::{LabelSet, LabelType, MetricName, StatsManager}, + foundation::{ + stats::{ArcRpcMetrics, RpcMetricLabels, RpcMetricsProvider}, + task::reap_joinset_background, + time::timeout, }, proto::{ common::{ self, CompressionAlgoPb, RpcCompressionInfo, RpcPacket, RpcRequest, RpcResponse, TunnelInfo, }, - rpc_impl::packet::BuildRpcPacketArgs, rpc_types::{controller::Controller, error::Result}, }, + rpc::packet::BuildRpcPacketArgs, tunnel::{ Tunnel, ZCPacketStream, mpsc::{MpscTunnel, MpscTunnelSender}, @@ -38,20 +41,19 @@ use super::{ #[derive(Debug, Clone, PartialEq, Eq, Hash)] struct PacketMergerKey { - from_peer_id: PeerId, + from_peer_id: crate::config::PeerId, transaction_id: i64, } pub struct Server { registry: Arc, - mpsc: Mutex>>>, - transport: Mutex, - tasks: Arc>>, + handler_tasks: Arc>>, + stopped: Arc, packet_mergers: Arc>, - stats_manager: Option>, + metrics: Option, } impl Default for Server { @@ -73,18 +75,32 @@ impl Server { mpsc: Mutex::new(Some(MpscTunnel::new(ring_a, None))), transport: Mutex::new(MpscTunnel::new(ring_b, None)), tasks: Arc::new(Mutex::new(JoinSet::new())), + handler_tasks: Arc::new(Mutex::new(JoinSet::new())), + stopped: Arc::new(AtomicBool::new(false)), packet_mergers: Arc::new(DashMap::new()), - stats_manager: None, + metrics: None, } } - pub fn new_with_registry_and_stats_manager( + pub fn new_with_registry_and_stats_manager( registry: Arc, - stats_manager: Arc, + stats_manager: T, + ) -> Self + where + T: RpcMetricsProvider, + { + let mut server = Self::new_with_registry(registry); + server.metrics = stats_manager.into_rpc_metrics(); + server + } + + pub fn new_with_registry_and_metrics( + registry: Arc, + metrics: ArcRpcMetrics, ) -> Self { - let mut ret = Self::new_with_registry(registry); - ret.stats_manager = Some(stats_manager); - ret + let mut server = Self::new_with_registry(registry); + server.metrics = Some(metrics); + server } pub fn registry(&self) -> &ServiceRegistry { @@ -100,26 +116,38 @@ impl Server { } pub fn run(&self) { - let tasks = self.tasks.clone(); - join_joinset_background(tasks.clone(), "rpc server".to_string()); + self.run_with_tunnel_info(None); + } + + pub(crate) fn run_with_tunnel_info(&self, tunnel_info: Option) { + self.stopped.store(false, Ordering::Relaxed); + let handler_tasks = self.handler_tasks.clone(); + self.tasks.lock().unwrap().spawn(reap_joinset_background( + handler_tasks.clone(), + "rpc server handlers", + )); let mpsc = self.mpsc.lock().unwrap().take().unwrap(); let packet_merges = self.packet_mergers.clone(); let reg = self.registry.clone(); - let stats_manager = self.stats_manager.clone(); - let t = Arc::downgrade(&tasks); - let tunnel_info = mpsc.tunnel_info(); - tasks.lock().unwrap().spawn(async move { + let tunnel_info = tunnel_info.or_else(|| mpsc.tunnel_info()); + let metrics = self.metrics.clone(); + let handler_tasks_weak = Arc::downgrade(&handler_tasks); + let stopped = self.stopped.clone(); + self.tasks.lock().unwrap().spawn(async move { let mut mpsc = mpsc; let mut rx = mpsc.get_stream(); while let Some(packet) = rx.next().await { - if let Err(err) = packet { - tracing::error!(?err, "Failed to receive packet"); - continue; - } - let packet = match common::RpcPacket::decode(packet.unwrap().payload()) { + let packet = match packet { + Err(err) => { + tracing::error!(?err, "Failed to receive packet"); + continue; + } + Ok(packet) => packet, + }; + let packet = match common::RpcPacket::decode(packet.payload()) { Err(err) => { tracing::error!(?err, "Failed to decode packet"); continue; @@ -144,30 +172,35 @@ impl Server { match ret { Ok(Some(packet)) => { packet_merges.remove(&key); - let Some(t) = t.upgrade() else { - tracing::error!("tasks is dropped"); + let Some(handler_tasks) = handler_tasks_weak.upgrade() else { + tracing::error!("rpc server handler task set is dropped"); return; }; - t.lock().unwrap().spawn(Self::handle_rpc( + let mut handler_tasks = handler_tasks.lock().unwrap(); + if stopped.load(Ordering::Relaxed) { + return; + } + handler_tasks.spawn(Self::handle_rpc( mpsc.get_sink(), packet, reg.clone(), tunnel_info.clone(), - stats_manager.clone(), + metrics.clone(), )); } Ok(None) => {} Err(err) => { - tracing::error!("Failed to feed packet to merger, {}", err.to_string()); + tracing::error!("Failed to feed packet to merger, {}", err); } } } }); let packet_mergers = self.packet_mergers.clone(); - tasks.lock().unwrap().spawn(async move { + self.tasks.lock().unwrap().spawn(async move { loop { - tokio::time::sleep(tokio::time::Duration::from_secs(5)).await; + crate::foundation::time::sleep(crate::foundation::time::Duration::from_secs(5)) + .await; packet_mergers.retain(|_, v| v.last_updated().elapsed().as_secs() < 10); packet_mergers.shrink_to_fit(); } @@ -211,7 +244,7 @@ impl Server { packet: RpcPacket, reg: Arc, tunnel_info: Option, - stats_manager: Option>, + metrics: Option, ) { let from_peer = packet.from_peer; let to_peer = packet.to_peer; @@ -219,22 +252,20 @@ impl Server { let trace_id = packet.trace_id; let desc = packet.descriptor.clone().unwrap(); let method_name = reg.get_method_name(&desc).unwrap_or("".to_owned()); - let labels = LabelSet::new() - .with_label_type(LabelType::NetworkName(desc.domain_name.to_string())) - .with_label_type(LabelType::SrcPeerId(from_peer)) - .with_label_type(LabelType::DstPeerId(to_peer)) - .with_label_type(LabelType::ServiceName(desc.service_name.to_string())) - .with_label_type(LabelType::MethodName(method_name)); + let labels = RpcMetricLabels { + network_name: desc.domain_name.clone(), + src_peer_id: from_peer, + dst_peer_id: to_peer, + service_name: desc.service_name.clone(), + method_name, + }; - // Record RPC server RX stats - if let Some(ref stats_manager) = stats_manager { - stats_manager - .get_counter(MetricName::PeerRpcServerRx, labels.clone()) - .inc(); + if let Some(metrics) = &metrics { + metrics.server_rx(&labels); } let mut resp_msg = RpcResponse::default(); - let now = Instant::now(); + let now = std::time::Instant::now(); let compression_info = packet.compression_info; let resp_bytes = Self::handle_rpc_request(packet, reg, tunnel_info).await; @@ -242,40 +273,18 @@ impl Server { match &resp_bytes { Ok(r) => { resp_msg.response = r.clone().into(); - - // Record successful RPC server TX and duration stats - if let Some(ref stats_manager) = stats_manager { - let labels = labels - .clone() - .with_label_type(LabelType::Status("success".to_string())); - - stats_manager - .get_counter(MetricName::PeerRpcServerTx, labels.clone()) - .inc(); - - let duration_ms = now.elapsed().as_millis() as u64; - stats_manager - .get_counter(MetricName::PeerRpcDuration, labels) - .add(duration_ms); + if let Some(metrics) = &metrics { + metrics.server_tx(&labels, now.elapsed().as_millis() as u64); } } Err(err) => { resp_msg.error = Some(err.into()); - - // Record RPC server error stats - if let Some(ref stats_manager) = stats_manager { - let labels = labels - .clone() - .with_label_type(LabelType::Status("error".to_string())); - - stats_manager - .get_counter(MetricName::PeerRpcErrors, labels.clone()) - .inc(); - - let duration_ms = now.elapsed().as_millis() as u64; - stats_manager - .get_counter(MetricName::PeerRpcDuration, labels) - .add(duration_ms); + if let Some(metrics) = &metrics { + metrics.server_error( + &labels, + Some(format!("{:?}", err)), + now.elapsed().as_millis() as u64, + ); } } }; @@ -308,11 +317,36 @@ impl Server { } } - pub fn inflight_count(&self) -> usize { - self.packet_mergers.len() - } - pub fn close(&self) { self.transport.lock().unwrap().close(); } + + pub async fn stop(&self) { + self.stopped.store(true, Ordering::Relaxed); + self.close(); + let (mut tasks, mut handler_tasks) = { + let mut task_slot = self.tasks.lock().unwrap(); + let mut handler_task_slot = self.handler_tasks.lock().unwrap(); + ( + std::mem::replace(&mut *task_slot, JoinSet::new()), + std::mem::replace(&mut *handler_task_slot, JoinSet::new()), + ) + }; + tasks.abort_all(); + handler_tasks.abort_all(); + while tasks.join_next().await.is_some() {} + while handler_tasks.join_next().await.is_some() {} + } +} + +#[cfg(any(test, feature = "test-utils"))] +mod test_utils { + use super::Server; + + impl Server { + #[doc(hidden)] + pub fn inflight_count(&self) -> usize { + self.packet_mergers.len() + } + } } diff --git a/easytier/src/proto/rpc_impl/service_registry.rs b/easytier-core/src/rpc/service_registry.rs similarity index 100% rename from easytier/src/proto/rpc_impl/service_registry.rs rename to easytier-core/src/rpc/service_registry.rs diff --git a/easytier-core/src/rpc/standalone.rs b/easytier-core/src/rpc/standalone.rs new file mode 100644 index 00000000..5c8a93fe --- /dev/null +++ b/easytier-core/src/rpc/standalone.rs @@ -0,0 +1,733 @@ +use std::{ + sync::{Arc, atomic::AtomicU32}, + time::Duration, +}; + +use anyhow::Context as _; +use tokio::task::JoinSet; + +use crate::{ + connectivity::protocol::raw::TunnelDialer, + proto::{ + common::TunnelInfo, + rpc_types::{__rt::RpcClientFactory, error::Error}, + }, + rpc::{bidirect::BidirectRpcManager, service_registry::ServiceRegistry}, + socket::SocketListener, + tunnel::Tunnel, +}; + +#[async_trait::async_trait] +#[auto_impl::auto_impl(Arc, Box)] +pub trait RpcServerHook: Send + Sync { + async fn on_new_client( + &self, + tunnel_info: Option, + ) -> Result, anyhow::Error> { + Ok(tunnel_info) + } + + async fn on_client_disconnected(&self, _tunnel_info: Option) {} +} + +struct DefaultHook; + +impl RpcServerHook for DefaultHook {} + +struct BoundListener { + listener: L, + // Release protection only after the listener has been dropped. + _guard: G, +} + +pub struct StandAloneServer { + registry: Arc, + listener: Option, + inflight_server: Arc, + tasks: JoinSet<()>, + hook: Option>, + rx_timeout: Option, +} + +impl StandAloneServer +where + L: SocketListener> + 'static, +{ + pub fn new(listener: L) -> Self { + Self { + registry: Arc::new(ServiceRegistry::new()), + listener: Some(listener), + inflight_server: Arc::new(AtomicU32::new(0)), + tasks: JoinSet::new(), + hook: None, + rx_timeout: Some(Duration::from_secs(60)), + } + } + + pub fn set_rx_timeout(&mut self, timeout: Option) { + self.rx_timeout = timeout; + } + + pub fn set_hook(&mut self, hook: Arc) { + self.hook = Some(hook); + } + + pub fn registry(&self) -> &ServiceRegistry { + &self.registry + } + + async fn serve_loop( + listener: &mut L, + inflight: Arc, + registry: Arc, + hook: Arc, + rx_timeout: Option, + ) -> Result<(), Error> { + let mut client_tasks = JoinSet::new(); + + loop { + let accepted = { + let accept = listener.accept(); + tokio::pin!(accept); + loop { + tokio::select! { + accepted = &mut accept => break accepted, + _ = client_tasks.join_next(), if !client_tasks.is_empty() => {} + } + } + }; + let tunnel = accepted?; + let tunnel_info = tunnel.info(); + let registry = registry.clone(); + let inflight_server = inflight.clone(); + let hook = hook.clone(); + + let tunnel_info = match hook.on_new_client(tunnel_info).await { + Ok(info) => info, + Err(error) => { + tracing::warn!(?error, "standalone hook.on_new_client failed"); + continue; + } + }; + + inflight_server.fetch_add(1, std::sync::atomic::Ordering::Relaxed); + client_tasks.spawn(async move { + let server = BidirectRpcManager::new().set_rx_timeout(rx_timeout); + server.rpc_server().registry().replace_registry(®istry); + server.run_with_tunnel_info(tunnel, tunnel_info.clone()); + server.wait().await; + hook.on_client_disconnected(tunnel_info).await; + inflight_server.fetch_sub(1, std::sync::atomic::Ordering::Relaxed); + }); + } + } + + async fn run_bound_listener( + mut bound_listener: BoundListener, + inflight_server: Arc, + registry: Arc, + hook: Arc, + rx_timeout: Option, + ) where + G: Send + 'static, + { + loop { + let ret = Self::serve_loop( + &mut bound_listener.listener, + inflight_server.clone(), + registry.clone(), + hook.clone(), + rx_timeout, + ) + .await; + if let Err(error) = ret { + tracing::error!( + ?error, + url = ?bound_listener.listener.local_url(), + "serve_loop exit unexpectedly" + ); + println!("standalone serve_loop exit unexpectedly: {error:?}"); + } + + crate::foundation::time::sleep(Duration::from_secs(1)).await; + } + } + + pub async fn serve(&mut self) -> Result<(), Error> { + self.serve_with_bound_listener((), |_| Ok(())).await + } + + #[cfg(feature = "management-rpc")] + pub(crate) fn listener_url(&self) -> url::Url { + self.listener + .as_ref() + .expect("standalone listener must be available before serve") + .local_url() + } + + pub(crate) async fn serve_with_bound_listener( + &mut self, + binding_guard: B, + on_bound: F, + ) -> Result<(), Error> + where + F: FnOnce(&url::Url) -> Result, + G: Send + 'static, + { + let mut listener = self.listener.take().unwrap(); + let hook = self.hook.take().unwrap_or_else(|| Arc::new(DefaultHook)); + let rx_timeout = self.rx_timeout; + + if let Err(error) = listener.listen().await.with_context(|| "failed to listen") { + drop(listener); + drop(binding_guard); + return Err(error.into()); + } + let guard = match on_bound(&listener.local_url()) { + Ok(guard) => guard, + Err(error) => { + drop(listener); + drop(binding_guard); + return Err(error); + } + }; + drop(binding_guard); + let bound_listener = BoundListener { + listener, + _guard: guard, + }; + + let registry = self.registry.clone(); + let inflight_server = self.inflight_server.clone(); + + self.tasks.spawn(Self::run_bound_listener( + bound_listener, + inflight_server, + registry, + hook, + rx_timeout, + )); + + Ok(()) + } + + pub fn inflight_server(&self) -> u32 { + self.inflight_server + .load(std::sync::atomic::Ordering::Relaxed) + } +} + +pub struct StandAloneClient { + connector: C, + client: Option, +} + +impl StandAloneClient { + pub fn new(connector: C) -> Self { + Self { + connector, + client: None, + } + } + + async fn connect(&mut self) -> Result, Error> { + Ok(self.connector.connect().await.with_context(|| { + format!( + "failed to connect to server: {:?}", + self.connector.remote_url() + ) + })?) + } + + pub async fn scoped_client( + &mut self, + domain_name: String, + ) -> Result { + let mut client = self.client.take(); + let error = client.as_ref().and_then(BidirectRpcManager::take_error); + if client.is_none() || error.is_some() { + tracing::info!(?error, "reconnect standalone RPC client"); + let tunnel = self.connect().await?; + let manager = BidirectRpcManager::new().set_rx_timeout(Some(Duration::from_secs(60))); + manager.run_with_tunnel(tunnel); + client = Some(manager); + } + + self.client = client; + + Ok(self + .client + .as_ref() + .unwrap() + .rpc_client() + .scoped_client::(1, 1, domain_name)) + } + + pub async fn wait(&mut self) { + if let Some(client) = self.client.take() { + client.wait().await; + } + } +} + +#[cfg(test)] +mod tests { + use std::{ + fmt, + sync::{ + Arc, Mutex, + atomic::{AtomicBool, AtomicU32, Ordering}, + }, + time::Duration, + }; + + use tokio::sync::mpsc; + use url::Url; + + use super::{RpcServerHook, StandAloneClient, StandAloneServer}; + use crate::{ + connectivity::protocol::raw::TunnelDialer, + foundation::time::{sleep, timeout}, + proto::{ + common::TunnelInfo, + peer_rpc::{ + GetGlobalPeerMapRequest, GetGlobalPeerMapResponse, PeerCenterRpc, + PeerCenterRpcClientFactory, PeerCenterRpcServer, ReportPeersRequest, + ReportPeersResponse, + }, + rpc_types::{ + controller::{BaseController, Controller as _}, + error, + }, + }, + socket::SocketListener, + tunnel::{ + Tunnel, + ring::{RING_TUNNEL_CAP, RingTunnel, create_ring_socket_pair, create_ring_tunnel_pair}, + }, + }; + + struct TestListener { + accepted: mpsc::Receiver>, + accept_tracker: Option>, + } + + struct BoundUrlListener { + listening: bool, + drop_order: Option>>>, + binding_guard_alive: Option>, + } + + struct TestProtectionLease(Arc>>); + + struct TestBindingGuard(Arc); + + #[derive(Default)] + struct AcceptTracker { + started: AtomicU32, + cancelled: AtomicU32, + } + + struct AcceptAttempt { + tracker: Arc, + completed: bool, + } + + impl Drop for AcceptAttempt { + fn drop(&mut self) { + if !self.completed { + self.tracker.cancelled.fetch_add(1, Ordering::Relaxed); + } + } + } + + impl fmt::Debug for TestListener { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter.debug_struct("TestListener").finish() + } + } + + impl fmt::Debug for BoundUrlListener { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter.debug_struct("BoundUrlListener").finish() + } + } + + impl Drop for BoundUrlListener { + fn drop(&mut self) { + if let Some(drop_order) = &self.drop_order { + drop_order.lock().unwrap().push("listener"); + } + } + } + + impl Drop for TestProtectionLease { + fn drop(&mut self) { + self.0.lock().unwrap().push("protection"); + } + } + + impl Drop for TestBindingGuard { + fn drop(&mut self) { + self.0.store(false, Ordering::Release); + } + } + + #[async_trait::async_trait] + impl SocketListener for TestListener { + type Accepted = Box; + + async fn listen(&mut self) -> anyhow::Result<()> { + Ok(()) + } + + async fn accept(&mut self) -> anyhow::Result { + let mut attempt = self.accept_tracker.as_ref().map(|tracker| { + tracker.started.fetch_add(1, Ordering::Relaxed); + AcceptAttempt { + tracker: tracker.clone(), + completed: false, + } + }); + let result = self + .accepted + .recv() + .await + .ok_or_else(|| anyhow::anyhow!("test listener closed")); + if let Some(attempt) = attempt.as_mut() { + attempt.completed = true; + } + result + } + + fn local_url(&self) -> Url { + "ring://standalone-rpc".parse().unwrap() + } + } + + #[async_trait::async_trait] + impl SocketListener for BoundUrlListener { + type Accepted = Box; + + async fn listen(&mut self) -> anyhow::Result<()> { + if let Some(alive) = &self.binding_guard_alive { + assert!(alive.load(Ordering::Acquire)); + } + self.listening = true; + Ok(()) + } + + async fn accept(&mut self) -> anyhow::Result { + std::future::pending().await + } + + fn local_url(&self) -> Url { + if self.listening { + "tcp://127.0.0.1:15888".parse().unwrap() + } else { + "tcp://127.0.0.1:0".parse().unwrap() + } + } + } + + struct TestDialer { + accepted: Arc>>>, + connections: Arc, + } + + #[async_trait::async_trait] + impl TunnelDialer for TestDialer { + async fn connect(&self) -> anyhow::Result> { + let (client_socket, accepted_socket) = create_ring_socket_pair(RING_TUNNEL_CAP); + let client = Box::new(RingTunnel::new(client_socket, None)); + let accepted = Box::new(RingTunnel::new( + accepted_socket, + Some(TunnelInfo { + tunnel_type: "tcp".to_owned(), + local_addr: Some("tcp://127.0.0.1:15888".parse::().unwrap().into()), + remote_addr: Some("tcp://127.0.0.1:40000".parse::().unwrap().into()), + resolved_remote_addr: Some( + "tcp://127.0.0.1:40000".parse::().unwrap().into(), + ), + }), + )); + let sender = self.accepted.lock().unwrap().clone(); + sender + .send(accepted) + .await + .map_err(|_| anyhow::anyhow!("test listener closed"))?; + self.connections.fetch_add(1, Ordering::Relaxed); + Ok(client) + } + + fn remote_url(&self) -> Url { + "ring://standalone-rpc".parse().unwrap() + } + } + + #[derive(Clone, Debug)] + struct TestRpcService; + + #[async_trait::async_trait] + impl PeerCenterRpc for TestRpcService { + type Controller = BaseController; + + async fn report_peers( + &self, + _controller: BaseController, + _request: ReportPeersRequest, + ) -> error::Result { + Ok(ReportPeersResponse::default()) + } + + async fn get_global_peer_map( + &self, + controller: BaseController, + _request: GetGlobalPeerMapRequest, + ) -> error::Result { + assert_eq!( + controller + .get_tunnel_info() + .map(|info| info.tunnel_type.as_str()), + Some("tcp") + ); + Ok(GetGlobalPeerMapResponse { + digest: Some(42), + ..Default::default() + }) + } + } + + #[derive(Default)] + struct CountingHook { + connected: AtomicU32, + disconnected: AtomicU32, + } + + #[async_trait::async_trait] + impl RpcServerHook for CountingHook { + async fn on_new_client( + &self, + tunnel_info: Option, + ) -> Result, anyhow::Error> { + self.connected.fetch_add(1, Ordering::Relaxed); + Ok(tunnel_info) + } + + async fn on_client_disconnected(&self, _tunnel_info: Option) { + self.disconnected.fetch_add(1, Ordering::Relaxed); + } + } + + #[tokio::test] + async fn serve_reports_the_actual_bound_listener_url() { + let mut server = StandAloneServer::new(BoundUrlListener { + listening: false, + drop_order: None, + binding_guard_alive: None, + }); + let mut bound_url = None; + + server + .serve_with_bound_listener((), |url| { + bound_url = Some(url.clone()); + Ok(()) + }) + .await + .unwrap(); + + assert_eq!( + bound_url.unwrap(), + "tcp://127.0.0.1:15888".parse::().unwrap() + ); + } + + #[tokio::test] + async fn binding_guard_covers_listen_and_bound_guard_creation() { + let binding_guard_alive = Arc::new(AtomicBool::new(true)); + let mut server = StandAloneServer::new(BoundUrlListener { + listening: false, + drop_order: None, + binding_guard_alive: Some(binding_guard_alive.clone()), + }); + + server + .serve_with_bound_listener(TestBindingGuard(binding_guard_alive.clone()), |_| { + assert!(binding_guard_alive.load(Ordering::Acquire)); + Ok(()) + }) + .await + .unwrap(); + + assert!(!binding_guard_alive.load(Ordering::Acquire)); + } + + #[tokio::test] + async fn listener_drops_before_its_bound_resource_guard() { + let drop_order = Arc::new(Mutex::new(Vec::new())); + let mut server = StandAloneServer::new(BoundUrlListener { + listening: false, + drop_order: Some(drop_order.clone()), + binding_guard_alive: None, + }); + server + .serve_with_bound_listener((), |_| Ok(TestProtectionLease(drop_order.clone()))) + .await + .unwrap(); + + drop(server); + timeout(Duration::from_secs(1), async { + loop { + if drop_order.lock().unwrap().len() == 2 { + break; + } + tokio::task::yield_now().await; + } + }) + .await + .unwrap(); + + assert_eq!(*drop_order.lock().unwrap(), ["listener", "protection"]); + } + + #[tokio::test] + async fn server_owns_accepted_client_lifecycle() { + let (sender, receiver) = mpsc::channel(1); + let hook = Arc::new(CountingHook::default()); + let mut server = StandAloneServer::new(TestListener { + accepted: receiver, + accept_tracker: None, + }); + server.set_hook(hook.clone()); + server.serve().await.unwrap(); + + let (client, accepted) = create_ring_tunnel_pair(); + sender.send(accepted).await.unwrap(); + + timeout(Duration::from_secs(1), async { + while server.inflight_server() != 1 { + sleep(Duration::from_millis(10)).await; + } + }) + .await + .unwrap(); + assert_eq!(hook.connected.load(Ordering::Relaxed), 1); + + drop(client); + timeout(Duration::from_secs(1), async { + while server.inflight_server() != 0 { + sleep(Duration::from_millis(10)).await; + } + }) + .await + .unwrap(); + assert_eq!(hook.disconnected.load(Ordering::Relaxed), 1); + } + + #[tokio::test] + async fn reaping_clients_does_not_cancel_an_in_progress_accept() { + let (sender, receiver) = mpsc::channel(1); + let tracker = Arc::new(AcceptTracker::default()); + let mut server = StandAloneServer::new(TestListener { + accepted: receiver, + accept_tracker: Some(tracker.clone()), + }); + server.serve().await.unwrap(); + + let (client, accepted) = create_ring_tunnel_pair(); + sender.send(accepted).await.unwrap(); + timeout(Duration::from_secs(1), async { + while tracker.started.load(Ordering::Relaxed) < 2 { + sleep(Duration::from_millis(10)).await; + } + }) + .await + .unwrap(); + + drop(client); + timeout(Duration::from_secs(1), async { + while server.inflight_server() != 0 { + sleep(Duration::from_millis(10)).await; + } + sleep(Duration::from_millis(20)).await; + }) + .await + .unwrap(); + + assert_eq!(tracker.started.load(Ordering::Relaxed), 2); + assert_eq!(tracker.cancelled.load(Ordering::Relaxed), 0); + } + + #[tokio::test] + async fn client_reuses_connection_and_reconnects_after_disconnect() { + let (sender, receiver) = mpsc::channel(2); + let accepted = Arc::new(Mutex::new(sender)); + let connections = Arc::new(AtomicU32::new(0)); + let dialer = TestDialer { + accepted: accepted.clone(), + connections: connections.clone(), + }; + let mut server = StandAloneServer::new(TestListener { + accepted: receiver, + accept_tracker: None, + }); + server + .registry() + .register(PeerCenterRpcServer::new(TestRpcService), "test"); + server.serve().await.unwrap(); + + let mut client = StandAloneClient::new(dialer); + let rpc = client + .scoped_client::>("test".to_string()) + .await + .unwrap(); + let response = rpc + .get_global_peer_map( + BaseController::default(), + GetGlobalPeerMapRequest::default(), + ) + .await + .unwrap(); + assert_eq!(response.digest, Some(42)); + + client + .scoped_client::>("test".to_string()) + .await + .unwrap(); + assert_eq!(connections.load(Ordering::Relaxed), 1); + + drop(server); + let (sender, receiver) = mpsc::channel(2); + *accepted.lock().unwrap() = sender; + let mut restarted_server = StandAloneServer::new(TestListener { + accepted: receiver, + accept_tracker: None, + }); + restarted_server + .registry() + .register(PeerCenterRpcServer::new(TestRpcService), "test"); + restarted_server.serve().await.unwrap(); + + let rpc = timeout(Duration::from_secs(1), async { + loop { + let rpc = client + .scoped_client::>("test".to_string()) + .await + .unwrap(); + if connections.load(Ordering::Relaxed) == 2 { + break rpc; + } + sleep(Duration::from_millis(10)).await; + } + }) + .await + .unwrap(); + + rpc.get_global_peer_map( + BaseController::default(), + GetGlobalPeerMapRequest::default(), + ) + .await + .unwrap(); + } +} diff --git a/easytier-core/src/socket/mod.rs b/easytier-core/src/socket/mod.rs new file mode 100644 index 00000000..e651f00d --- /dev/null +++ b/easytier-core/src/socket/mod.rs @@ -0,0 +1,104 @@ +//! Core-visible socket primitives. +//! +//! This Module is below [`crate::tunnel`]. Sockets represent established or +//! bindable communication endpoints; tunnels are produced later by runtime +//! upgraders and can be handed to peers. Host capability seams (DNS, packet +//! egress, environment facts, and the WASI mechanism backend) live in +//! [`crate::host`]. + +pub mod ring; +pub mod tcp; +pub mod udp; + +use std::{fmt::Debug, sync::Arc}; + +use async_trait::async_trait; +use serde::{Deserialize, Serialize}; +use url::Url; + +pub trait ListenerConnectionCounter: Debug + Send + Sync { + fn get(&self) -> Option; +} + +#[derive(Debug)] +struct EmptyConnectionCounter; + +impl ListenerConnectionCounter for EmptyConnectionCounter { + fn get(&self) -> Option { + None + } +} + +#[async_trait] +#[auto_impl::auto_impl(Box)] +pub trait SocketListener: Debug + Send { + type Accepted: Send + 'static; + + async fn listen(&mut self) -> anyhow::Result<()>; + + async fn accept(&mut self) -> anyhow::Result; + + fn local_url(&self) -> Url; + + fn connection_counter(&self) -> Arc { + Arc::new(EmptyConnectionCounter) + } +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)] +pub enum IpVersion { + V4, + V6, + Both, +} + +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +pub struct NetNamespace(String); + +impl NetNamespace { + pub fn new(token: impl Into) -> Self { + Self(token.into()) + } + + pub fn token(&self) -> &str { + &self.0 + } +} + +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +pub struct SocketContext { + pub ip_version: IpVersion, + pub socket_mark: Option, + pub netns: Option, +} + +impl SocketContext { + pub fn new() -> Self { + Self::default() + } + + pub fn with_ip_version(mut self, ip_version: IpVersion) -> Self { + self.ip_version = ip_version; + self + } + + pub fn with_socket_mark(mut self, socket_mark: Option) -> Self { + self.socket_mark = socket_mark; + self + } + + pub fn with_netns(mut self, netns: Option) -> Self { + self.netns = netns; + self + } +} + +impl Default for SocketContext { + fn default() -> Self { + Self { + ip_version: IpVersion::Both, + socket_mark: None, + netns: None, + } + } +} diff --git a/easytier-core/src/socket/ring.rs b/easytier-core/src/socket/ring.rs new file mode 100644 index 00000000..1cf00ad2 --- /dev/null +++ b/easytier-core/src/socket/ring.rs @@ -0,0 +1,299 @@ +use std::{ + fmt::Debug, + pin::Pin, + sync::{Arc, Mutex}, + task::{Context, Poll, ready}, +}; + +use async_ringbuf::{AsyncHeapCons, AsyncHeapProd, AsyncHeapRb, traits::*}; +use futures::{Sink, SinkExt, Stream, StreamExt}; +use uuid::Uuid; + +pub const RING_SOCKET_CAPACITY: usize = 128; +const RING_SOCKET_RESERVED_CAPACITY: usize = 4; + +pub type RingSocketId = Uuid; + +#[derive(Debug, thiserror::Error, Clone, Copy, PartialEq, Eq)] +pub enum RingSocketError { + #[error("ring socket already split")] + AlreadySplit, + #[error("ring socket closed")] + Closed, + #[error("ring socket full")] + Full, +} + +#[derive(Debug, thiserror::Error, Clone, Copy, PartialEq, Eq)] +pub enum RingSocketSendError { + #[error("ring socket closed")] + Closed(T), + #[error("ring socket full")] + Full(T), +} + +pub type RingSocketStreamItem = Result; + +/// An in-process socket primitive. +/// +/// `RingSocket` is intentionally below `Tunnel`: it contains no tunnel schema, +/// no peer metadata, and no `TunnelInfo`. The core ring tunnel Module wraps it +/// into a `Tunnel` when a peer connection needs one. +pub struct RingSocket { + id: RingSocketId, + parts: Mutex>>, +} + +struct RingSocketParts { + recv: AsyncHeapCons, + send: AsyncHeapProd, +} + +impl RingSocket { + pub fn pair(capacity: usize) -> (Arc, Arc) { + Self::pair_with_ids(Uuid::new_v4(), Uuid::new_v4(), capacity) + } + + pub fn pair_with_ids( + first_id: RingSocketId, + second_id: RingSocketId, + capacity: usize, + ) -> (Arc, Arc) { + let capacity = std::cmp::max(RING_SOCKET_RESERVED_CAPACITY * 2, capacity); + let first_to_second = AsyncHeapRb::new(capacity); + let second_to_first = AsyncHeapRb::new(capacity); + let (first_to_second_send, first_to_second_recv) = first_to_second.split(); + let (second_to_first_send, second_to_first_recv) = second_to_first.split(); + + ( + Arc::new(Self { + id: first_id, + parts: Mutex::new(Some(RingSocketParts { + recv: second_to_first_recv, + send: first_to_second_send, + })), + }), + Arc::new(Self { + id: second_id, + parts: Mutex::new(Some(RingSocketParts { + recv: first_to_second_recv, + send: second_to_first_send, + })), + }), + ) + } + + pub fn id(&self) -> RingSocketId { + self.id + } + + pub fn split(&self) -> (RingSocketReceiver, RingSocketSender) { + self.try_split().expect("RingSocket can only be split once") + } + + pub fn try_split( + &self, + ) -> Result<(RingSocketReceiver, RingSocketSender), RingSocketError> { + let parts = self + .parts + .lock() + .unwrap() + .take() + .ok_or(RingSocketError::AlreadySplit)?; + + Ok(( + RingSocketReceiver { + id: self.id, + recv: parts.recv, + }, + RingSocketSender { + id: self.id, + send: parts.send, + }, + )) + } +} + +impl Debug for RingSocket { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("RingSocket") + .field("id", &self.id) + .finish_non_exhaustive() + } +} + +pub struct RingSocketReceiver { + id: RingSocketId, + recv: AsyncHeapCons, +} + +impl Stream for RingSocketReceiver { + type Item = RingSocketStreamItem; + + fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + match ready!(self.get_mut().recv.poll_next_unpin(cx)) { + Some(item) => Poll::Ready(Some(Ok(item))), + None => Poll::Ready(None), + } + } +} + +impl Debug for RingSocketReceiver { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("RingSocketReceiver") + .field("id", &self.id) + .field("len", &self.recv.base().occupied_len()) + .field("cap", &self.recv.base().capacity()) + .finish() + } +} + +pub struct RingSocketSender { + id: RingSocketId, + send: AsyncHeapProd, +} + +impl RingSocketSender { + pub fn try_send(&mut self, item: T) -> Result<(), RingSocketSendError> { + if self.send.is_closed() { + return Err(RingSocketSendError::Closed(item)); + } + + let base = self.send.base(); + if base.occupied_len() >= base.capacity().get() - RING_SOCKET_RESERVED_CAPACITY { + return Err(RingSocketSendError::Full(item)); + } + + self.send.try_push(item).map_err(|item| { + if self.send.is_closed() { + RingSocketSendError::Closed(item) + } else { + RingSocketSendError::Full(item) + } + }) + } + + pub fn force_send(&mut self, item: T) -> Result<(), RingSocketSendError> { + if self.send.is_closed() { + return Err(RingSocketSendError::Closed(item)); + } + + self.send.try_push(item).map_err(|item| { + if self.send.is_closed() { + RingSocketSendError::Closed(item) + } else { + RingSocketSendError::Full(item) + } + }) + } +} + +impl Sink for RingSocketSender { + type Error = RingSocketError; + + fn poll_ready(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + ready!(self.get_mut().send.poll_ready_unpin(cx)).map_err(|_| RingSocketError::Closed)?; + Poll::Ready(Ok(())) + } + + fn start_send(self: Pin<&mut Self>, item: T) -> Result<(), Self::Error> { + self.get_mut() + .force_send(item) + .map_err(|error| match error { + RingSocketSendError::Closed(_) => RingSocketError::Closed, + RingSocketSendError::Full(_) => RingSocketError::Full, + }) + } + + fn poll_flush(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + ready!(self.get_mut().send.poll_flush_unpin(cx)).map_err(|_| RingSocketError::Closed)?; + Poll::Ready(Ok(())) + } + + fn poll_close(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + ready!(self.get_mut().send.poll_close_unpin(cx)).map_err(|_| RingSocketError::Closed)?; + Poll::Ready(Ok(())) + } +} + +impl Debug for RingSocketSender { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("RingSocketSender") + .field("id", &self.id) + .field("len", &self.send.base().occupied_len()) + .field("cap", &self.send.base().capacity()) + .finish() + } +} + +#[cfg(test)] +mod tests { + use futures::{SinkExt, StreamExt}; + + use crate::packet::ZCPacket; + + use super::*; + + #[tokio::test] + async fn ring_socket_pair_transfers_packets() { + let (left, right) = RingSocket::::pair(8); + let (_left_recv, mut left_send) = left.split(); + let (mut right_recv, _right_send) = right.split(); + + let packet = ZCPacket::new_with_payload(&[1, 2, 3]); + left_send.send(packet.clone()).await.unwrap(); + + let received = right_recv.next().await.unwrap().unwrap(); + assert_eq!(received.payload(), packet.payload()); + } + + #[test] + fn ring_socket_split_is_single_use() { + let (left, _right) = RingSocket::::pair(8); + + let _first = left.try_split().unwrap(); + assert_eq!(left.try_split().unwrap_err(), RingSocketError::AlreadySplit); + } + + #[test] + fn ring_socket_try_send_reserves_capacity() { + let (left, _right) = RingSocket::::pair(8); + let (_left_recv, mut left_send) = left.split(); + + for _ in 0..4 { + left_send + .try_send(ZCPacket::new_with_payload(&[1])) + .unwrap(); + } + + assert!( + left_send + .try_send(ZCPacket::new_with_payload(&[1])) + .is_err_and(|error| matches!(error, RingSocketSendError::Full(_))) + ); + assert!( + left_send + .force_send(ZCPacket::new_with_payload(&[1])) + .is_ok() + ); + } + + #[test] + fn ring_socket_sync_send_reports_closed_receiver() { + let (left, right) = RingSocket::::pair(8); + let (_left_recv, mut left_send) = left.split(); + let (right_recv, _right_send) = right.split(); + drop(right_recv); + + assert!( + left_send + .try_send(ZCPacket::new_with_payload(&[1])) + .is_err_and(|error| matches!(error, RingSocketSendError::Closed(_))) + ); + assert!( + left_send + .force_send(ZCPacket::new_with_payload(&[1])) + .is_err_and(|error| matches!(error, RingSocketSendError::Closed(_))) + ); + } +} diff --git a/easytier-core/src/socket/tcp.rs b/easytier-core/src/socket/tcp.rs new file mode 100644 index 00000000..52588b12 --- /dev/null +++ b/easytier-core/src/socket/tcp.rs @@ -0,0 +1,683 @@ +use std::{fmt, io, net::SocketAddr, sync::Arc}; + +use async_trait::async_trait; +use serde::{Deserialize, Serialize}; +use tokio::io::{AsyncRead, AsyncWrite}; + +use crate::socket::{IpVersion, SocketContext, SocketListener}; + +/// A core-visible TCP stream endpoint. +/// +/// Implementations are runtime adapters over concrete TCP stream types. This +/// trait deliberately stays below tunnel framing: it only exposes stream I/O and +/// socket addresses. +pub trait VirtualTcpSocket: AsyncRead + AsyncWrite + Unpin + Send + 'static { + fn local_addr(&self) -> io::Result; + + fn peer_addr(&self) -> io::Result; + + /// Optional host transport label retained in tunnel management metadata. + fn transport_label(&self) -> Option<&str> { + None + } +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum TcpSocketPurpose { + DirectConnect, + FakeTcp, + HolePunch, + ManualConnect, + ProxyNat, + StunProbe, + Socks5, + PortForward, + DataPlane, +} + +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +pub struct TcpBindOptions { + #[serde(default)] + pub context: SocketContext, + pub local_addr: Option, + pub bind_device: Option, + /// `None` delegates the platform default to the host socket adapter. + pub reuse_addr: Option, + pub reuse_port: bool, + pub only_v6: bool, +} + +impl TcpBindOptions { + pub fn new() -> Self { + Self { + context: SocketContext::default(), + local_addr: None, + bind_device: None, + reuse_addr: None, + reuse_port: false, + only_v6: false, + } + } + + pub fn with_local_addr(mut self, local_addr: Option) -> Self { + self.local_addr = local_addr; + self + } + + pub fn with_socket_mark(mut self, socket_mark: Option) -> Self { + self.context.socket_mark = socket_mark; + self + } + + pub fn with_context(mut self, context: SocketContext) -> Self { + self.context = context; + self + } + + pub fn with_ip_version(mut self, ip_version: IpVersion) -> Self { + self.context.ip_version = ip_version; + self + } + + pub fn with_bind_device(mut self, bind_device: Option) -> Self { + self.bind_device = bind_device; + self + } + + pub fn with_reuse_addr(mut self, reuse_addr: bool) -> Self { + self.reuse_addr = Some(reuse_addr); + self + } + + pub fn with_reuse_port(mut self, reuse_port: bool) -> Self { + self.reuse_port = reuse_port; + self + } + + pub fn with_only_v6(mut self, only_v6: bool) -> Self { + self.only_v6 = only_v6; + self + } +} + +impl Default for TcpBindOptions { + fn default() -> Self { + Self::new() + } +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct TcpConnectOptions { + pub remote_addr: SocketAddr, + pub bind: TcpBindOptions, + pub purpose: TcpSocketPurpose, +} + +impl TcpConnectOptions { + pub fn direct_connect(remote_addr: SocketAddr) -> Self { + Self { + remote_addr, + bind: TcpBindOptions::default(), + purpose: TcpSocketPurpose::DirectConnect, + } + } + + pub fn with_purpose(mut self, purpose: TcpSocketPurpose) -> Self { + self.purpose = purpose; + self + } + + pub fn hole_punch(remote_addr: SocketAddr, local_addr: Option) -> Self { + Self { + remote_addr, + bind: TcpBindOptions::default().with_local_addr(local_addr), + purpose: TcpSocketPurpose::HolePunch, + } + } + + pub fn manual_connect(remote_addr: SocketAddr, local_addr: Option) -> Self { + Self { + remote_addr, + bind: TcpBindOptions::default().with_local_addr(local_addr), + purpose: TcpSocketPurpose::ManualConnect, + } + } + + pub fn proxy_nat(remote_addr: SocketAddr) -> Self { + Self { + remote_addr, + bind: TcpBindOptions::default(), + purpose: TcpSocketPurpose::ProxyNat, + } + } + + pub fn stun_probe(remote_addr: SocketAddr, local_addr: SocketAddr) -> Self { + Self { + remote_addr, + bind: TcpBindOptions::default().with_local_addr(Some(local_addr)), + purpose: TcpSocketPurpose::StunProbe, + } + } + + pub fn socks5(remote_addr: SocketAddr) -> Self { + Self::direct_connect(remote_addr).with_purpose(TcpSocketPurpose::Socks5) + } + + pub fn port_forward(remote_addr: SocketAddr) -> Self { + Self::direct_connect(remote_addr).with_purpose(TcpSocketPurpose::PortForward) + } + + pub fn data_plane(remote_addr: SocketAddr) -> Self { + Self::direct_connect(remote_addr).with_purpose(TcpSocketPurpose::DataPlane) + } + + pub fn with_bind(mut self, bind: TcpBindOptions) -> Self { + self.bind = bind; + self + } +} + +#[async_trait] +pub trait VirtualTcpSocketFactory: Send + Sync + 'static { + type Socket: VirtualTcpSocket; + + async fn connect_tcp(&self, options: TcpConnectOptions) -> anyhow::Result; +} + +#[async_trait] +pub trait VirtualTcpListener: Send + Sync + 'static { + type Socket: VirtualTcpSocket; + + fn local_addr(&self) -> io::Result; + + async fn accept(&self) -> io::Result<(Self::Socket, SocketAddr)>; +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] +pub enum TcpListenPurpose { + DirectConnect, + HolePunch, + ManualConnect, + ProxyNat, + Socks5, + PortForward, + PortLease, +} + +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +pub struct TcpListenOptions { + pub bind: TcpBindOptions, + pub purpose: TcpListenPurpose, +} + +impl TcpListenOptions { + pub fn direct_connect(local_addr: SocketAddr) -> Self { + Self { + bind: TcpBindOptions::default().with_local_addr(Some(local_addr)), + purpose: TcpListenPurpose::DirectConnect, + } + } + + pub fn hole_punch(local_addr: SocketAddr) -> Self { + Self { + bind: TcpBindOptions::default().with_local_addr(Some(local_addr)), + purpose: TcpListenPurpose::HolePunch, + } + } + + pub fn manual_connect(local_addr: SocketAddr) -> Self { + Self { + bind: TcpBindOptions::default().with_local_addr(Some(local_addr)), + purpose: TcpListenPurpose::ManualConnect, + } + } + + pub fn proxy_nat(local_addr: SocketAddr) -> Self { + Self { + bind: TcpBindOptions::default().with_local_addr(Some(local_addr)), + purpose: TcpListenPurpose::ProxyNat, + } + } + + pub fn socks5(local_addr: SocketAddr) -> Self { + Self { + bind: TcpBindOptions::default().with_local_addr(Some(local_addr)), + purpose: TcpListenPurpose::Socks5, + } + } + + pub fn port_forward(local_addr: SocketAddr) -> Self { + Self { + bind: TcpBindOptions::default().with_local_addr(Some(local_addr)), + purpose: TcpListenPurpose::PortForward, + } + } + + pub fn port_lease(local_addr: SocketAddr) -> Self { + Self { + bind: TcpBindOptions::default().with_local_addr(Some(local_addr)), + purpose: TcpListenPurpose::PortLease, + } + } + + pub fn with_bind(mut self, bind: TcpBindOptions) -> Self { + self.bind = bind; + self + } +} + +#[async_trait] +pub trait VirtualTcpListenerFactory: Send + Sync + 'static { + type Listener: VirtualTcpListener; + + async fn bind_tcp(&self, options: TcpListenOptions) -> anyhow::Result>; +} + +type AcceptedTcpSocket = + <::Listener as VirtualTcpListener>::Socket; + +pub struct TcpSocketListener +where + F: VirtualTcpListenerFactory, +{ + url: url::Url, + options: TcpListenOptions, + factory: Arc, + listener: Option>, +} + +impl TcpSocketListener +where + F: VirtualTcpListenerFactory, +{ + pub fn new_with_options(url: url::Url, options: TcpListenOptions, factory: Arc) -> Self { + Self { + url, + options, + factory, + listener: None, + } + } + + fn listener(&self) -> anyhow::Result> { + self.listener + .clone() + .ok_or_else(|| anyhow::anyhow!("tcp socket listener is not started")) + } +} + +impl fmt::Debug for TcpSocketListener +where + F: VirtualTcpListenerFactory, +{ + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.debug_struct("TcpSocketListener") + .field("url", &self.url) + .field("options", &self.options) + .field("listening", &self.listener.is_some()) + .finish() + } +} + +#[async_trait] +impl SocketListener for TcpSocketListener +where + F: VirtualTcpListenerFactory, +{ + type Accepted = AcceptedTcpSocket; + + async fn listen(&mut self) -> anyhow::Result<()> { + if self.listener.is_some() { + return Ok(()); + } + + let listener = self.factory.bind_tcp(self.options.clone()).await?; + let local_addr = listener.local_addr()?; + self.url + .set_port(Some(local_addr.port())) + .map_err(|_| anyhow::anyhow!("failed to update tcp listener port for {}", self.url))?; + self.listener = Some(listener); + Ok(()) + } + + async fn accept(&mut self) -> anyhow::Result { + loop { + let listener = self.listener()?; + match listener.accept().await { + Ok((socket, _)) => return Ok(socket), + Err(error) if is_retryable_tcp_accept_error(&error) => { + tracing::warn!(?error, "tcp accept failed with retryable error"); + } + Err(error) => { + tracing::warn!(?error, "tcp accept failed"); + return Err(error.into()); + } + } + } + } + + fn local_url(&self) -> url::Url { + self.url.clone() + } +} + +fn is_retryable_tcp_accept_error(error: &io::Error) -> bool { + use io::ErrorKind::*; + matches!( + error.kind(), + NotConnected | ConnectionAborted | ConnectionRefused | ConnectionReset + ) +} + +#[cfg(test)] +mod tests { + use std::{ + collections::VecDeque, + pin::Pin, + sync::Mutex, + task::{Context, Poll}, + }; + + use tokio::io::{DuplexStream, ReadBuf}; + + use super::*; + + struct MockTcpSocket { + stream: DuplexStream, + local_addr: SocketAddr, + peer_addr: SocketAddr, + } + + impl MockTcpSocket { + fn new(local_addr: SocketAddr, peer_addr: SocketAddr) -> Self { + let (stream, _) = tokio::io::duplex(64); + Self { + stream, + local_addr, + peer_addr, + } + } + } + + impl AsyncRead for MockTcpSocket { + fn poll_read( + mut self: Pin<&mut Self>, + cx: &mut Context<'_>, + buf: &mut ReadBuf<'_>, + ) -> Poll> { + Pin::new(&mut self.stream).poll_read(cx, buf) + } + } + + impl AsyncWrite for MockTcpSocket { + fn poll_write( + mut self: Pin<&mut Self>, + cx: &mut Context<'_>, + buf: &[u8], + ) -> Poll> { + Pin::new(&mut self.stream).poll_write(cx, buf) + } + + fn poll_flush(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + Pin::new(&mut self.stream).poll_flush(cx) + } + + fn poll_shutdown(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + Pin::new(&mut self.stream).poll_shutdown(cx) + } + } + + impl VirtualTcpSocket for MockTcpSocket { + fn local_addr(&self) -> io::Result { + Ok(self.local_addr) + } + + fn peer_addr(&self) -> io::Result { + Ok(self.peer_addr) + } + } + + struct MockTcpListener { + local_addr: SocketAddr, + accepts: Mutex>>, + } + + impl MockTcpListener { + fn new(local_addr: SocketAddr, accepts: Vec>) -> Self { + Self { + local_addr, + accepts: Mutex::new(accepts.into()), + } + } + } + + #[async_trait::async_trait] + impl VirtualTcpListener for MockTcpListener { + type Socket = MockTcpSocket; + + fn local_addr(&self) -> io::Result { + Ok(self.local_addr) + } + + async fn accept(&self) -> io::Result<(Self::Socket, SocketAddr)> { + let result = self.accepts.lock().unwrap().pop_front(); + match result { + Some(Ok(socket)) => { + let peer_addr = socket.peer_addr()?; + Ok((socket, peer_addr)) + } + Some(Err(error)) => Err(error), + None => std::future::pending().await, + } + } + } + + struct MockTcpListenerFactory { + listener: Arc, + binds: Mutex>, + } + + impl MockTcpListenerFactory { + fn new(listener: Arc) -> Self { + Self { + listener, + binds: Mutex::new(Vec::new()), + } + } + } + + #[async_trait::async_trait] + impl VirtualTcpListenerFactory for MockTcpListenerFactory { + type Listener = MockTcpListener; + + async fn bind_tcp(&self, options: TcpListenOptions) -> anyhow::Result> { + self.binds.lock().unwrap().push(options); + Ok(self.listener.clone()) + } + } + + #[test] + fn tcp_connect_options_preserve_socket_purpose() { + let remote_addr = SocketAddr::from(([127, 0, 0, 1], 11010)); + let local_addr = SocketAddr::from(([0, 0, 0, 0], 0)); + + assert_eq!( + TcpConnectOptions::direct_connect(remote_addr), + TcpConnectOptions { + remote_addr, + bind: TcpBindOptions::default(), + purpose: TcpSocketPurpose::DirectConnect, + } + ); + assert_eq!( + TcpConnectOptions::hole_punch(remote_addr, Some(local_addr)), + TcpConnectOptions { + remote_addr, + bind: TcpBindOptions::default().with_local_addr(Some(local_addr)), + purpose: TcpSocketPurpose::HolePunch, + } + ); + assert_eq!( + TcpConnectOptions::manual_connect(remote_addr, Some(local_addr)), + TcpConnectOptions { + remote_addr, + bind: TcpBindOptions::default().with_local_addr(Some(local_addr)), + purpose: TcpSocketPurpose::ManualConnect, + } + ); + assert_eq!( + TcpConnectOptions::proxy_nat(remote_addr), + TcpConnectOptions { + remote_addr, + bind: TcpBindOptions::default(), + purpose: TcpSocketPurpose::ProxyNat, + } + ); + assert_eq!( + TcpConnectOptions::stun_probe(remote_addr, local_addr), + TcpConnectOptions { + remote_addr, + bind: TcpBindOptions::default().with_local_addr(Some(local_addr)), + purpose: TcpSocketPurpose::StunProbe, + } + ); + assert_eq!( + TcpConnectOptions::socks5(remote_addr).purpose, + TcpSocketPurpose::Socks5 + ); + assert_eq!( + TcpConnectOptions::port_forward(remote_addr).purpose, + TcpSocketPurpose::PortForward + ); + assert_eq!( + TcpConnectOptions::data_plane(remote_addr).purpose, + TcpSocketPurpose::DataPlane + ); + } + + #[test] + fn tcp_listen_options_preserve_socket_purpose() { + let local_addr = SocketAddr::from(([0, 0, 0, 0], 11010)); + + assert_eq!( + TcpListenOptions::socks5(local_addr).purpose, + TcpListenPurpose::Socks5 + ); + assert_eq!( + TcpListenOptions::port_forward(local_addr).purpose, + TcpListenPurpose::PortForward + ); + assert_eq!( + TcpListenOptions::port_lease(local_addr).purpose, + TcpListenPurpose::PortLease + ); + + assert_eq!( + TcpListenOptions::direct_connect(local_addr), + TcpListenOptions { + bind: TcpBindOptions::default().with_local_addr(Some(local_addr)), + purpose: TcpListenPurpose::DirectConnect, + } + ); + assert_eq!( + TcpListenOptions::hole_punch(local_addr), + TcpListenOptions { + bind: TcpBindOptions::default().with_local_addr(Some(local_addr)), + purpose: TcpListenPurpose::HolePunch, + } + ); + assert_eq!( + TcpListenOptions::manual_connect(local_addr), + TcpListenOptions { + bind: TcpBindOptions::default().with_local_addr(Some(local_addr)), + purpose: TcpListenPurpose::ManualConnect, + } + ); + assert_eq!( + TcpListenOptions::proxy_nat(local_addr), + TcpListenOptions { + bind: TcpBindOptions::default().with_local_addr(Some(local_addr)), + purpose: TcpListenPurpose::ProxyNat, + } + ); + } + + #[test] + fn tcp_bind_options_preserve_socket_configuration() { + let local_addr = SocketAddr::from(([0, 0, 0, 0], 0)); + let options = TcpBindOptions::default() + .with_local_addr(Some(local_addr)) + .with_socket_mark(Some(7)) + .with_bind_device(Some("eth0".to_owned())) + .with_reuse_addr(true) + .with_reuse_port(true) + .with_only_v6(true); + + assert_eq!( + options, + TcpBindOptions { + context: SocketContext::default().with_socket_mark(Some(7)), + local_addr: Some(local_addr), + bind_device: Some("eth0".to_owned()), + reuse_addr: Some(true), + reuse_port: true, + only_v6: true, + } + ); + } + + #[test] + fn tcp_bind_default_delegates_reuse_addr_policy_to_host() { + assert_eq!(TcpBindOptions::default().reuse_addr, None); + } + + #[tokio::test] + async fn tcp_socket_listener_binds_and_accepts_socket() { + let requested_addr = SocketAddr::from(([127, 0, 0, 1], 0)); + let bound_addr = SocketAddr::from(([127, 0, 0, 1], 12000)); + let peer_addr = SocketAddr::from(([127, 0, 0, 1], 12001)); + let options = TcpListenOptions::direct_connect(requested_addr); + let listener = Arc::new(MockTcpListener::new( + bound_addr, + vec![Ok(MockTcpSocket::new(bound_addr, peer_addr))], + )); + let factory = Arc::new(MockTcpListenerFactory::new(listener)); + let mut socket_listener = TcpSocketListener::new_with_options( + "tcp://127.0.0.1:0".parse().unwrap(), + options.clone(), + factory.clone(), + ); + + socket_listener.listen().await.unwrap(); + let accepted = socket_listener.accept().await.unwrap(); + + assert_eq!(socket_listener.local_url().port(), Some(bound_addr.port())); + assert_eq!(accepted.peer_addr().unwrap(), peer_addr); + assert_eq!(factory.binds.lock().unwrap().as_slice(), &[options]); + } + + #[tokio::test] + async fn tcp_socket_listener_retries_retryable_accept_error() { + let requested_addr = SocketAddr::from(([127, 0, 0, 1], 0)); + let bound_addr = SocketAddr::from(([127, 0, 0, 1], 12010)); + let peer_addr = SocketAddr::from(([127, 0, 0, 1], 12011)); + let listener = Arc::new(MockTcpListener::new( + bound_addr, + vec![ + Err(io::Error::new(io::ErrorKind::ConnectionReset, "reset")), + Ok(MockTcpSocket::new(bound_addr, peer_addr)), + ], + )); + let factory = Arc::new(MockTcpListenerFactory::new(listener)); + let mut socket_listener = TcpSocketListener::new_with_options( + "tcp://127.0.0.1:0".parse().unwrap(), + TcpListenOptions::direct_connect(requested_addr), + factory, + ); + + socket_listener.listen().await.unwrap(); + let accepted = socket_listener.accept().await.unwrap(); + + assert_eq!(accepted.peer_addr().unwrap(), peer_addr); + } +} diff --git a/easytier-core/src/socket/udp/layer.rs b/easytier-core/src/socket/udp/layer.rs new file mode 100644 index 00000000..9aec1ea6 --- /dev/null +++ b/easytier-core/src/socket/udp/layer.rs @@ -0,0 +1,1144 @@ +use std::{ + io, + net::{SocketAddr, SocketAddrV4, SocketAddrV6}, + sync::{Arc, Mutex as StdMutex, atomic::Ordering}, +}; + +use async_trait::async_trait; +use bytes::BytesMut; +use dashmap::DashMap; +use tokio::{ + sync::{Mutex as TokioMutex, Semaphore, mpsc, watch}, + task::JoinHandle, +}; + +use crate::packet::{ZCPacket, new_hole_punch_packet}; + +use super::{ + UDP_SESSION_CONNECT_TIMEOUT, UDP_SESSION_QUEUE_CAPACITY, UDP_SESSION_RESEND_INTERVAL, + packet::{ + EasyTierUdpPacketKind, UdpDatagramClassification, UdpSessionPacketKind, + classify_udp_datagram, extract_dst_addr_from_v4_hole_punch_packet, + extract_v6_hole_punch_packet, new_sack_packet, new_syn_packet, + }, + session::{ + ClassifiedUdpSessionAccept, ClassifiedUdpSessionAccepts, ClassifiedUdpSessionKey, + ClassifiedUdpSessionRegistry, PendingUdpSessionConnect, PendingUdpSessionConnects, + UdpConnectControl, UdpSession, UdpSessionClose, UdpSessionCodec, UdpSessionConnectError, + UdpSessionConnectRequest, UdpSessionConnector, UdpSessionDatagram, UdpSessionEnqueuePolicy, + UdpSessionKey, UdpSessionKind, UdpSessionLayerControl, UdpSessionProtocol, + UdpSessionRegistry, close_all_classified_udp_sessions, close_all_udp_sessions, + close_classified_udp_session, close_udp_session, create_udp_session_rings, + dispatch_payload_to_session, udp_session_registry_entry, + }, + virtual_socket::{ + NoopUdpSessionStunResponder, PreferredIpv6Source, UdpSessionStunResponder, + UdpSocketRecvMeta, UdpSocketSendMeta, VirtualUdpSocket, VirtualUdpSocketFactory, + }, +}; + +pub(super) const UDP_SESSION_HOLE_PUNCH_PACKET_BODY_LEN: u16 = 32; + +#[derive(Debug)] +pub struct UdpSessionLayer { + socket: Arc, + _stun_responder: Arc, + pub(super) sessions: Arc, + classified_sessions: Arc, + pub(super) classified_accepts: Arc, + pub(super) pending_connects: Arc, + mux_accepted_rx: TokioMutex>, + _control_rx: TokioMutex>, + session_shutdown_tx: watch::Sender, + recv_task: JoinHandle<()>, +} + +pub(super) fn create_classified_udp_session_accepts() -> Arc { + let accepts = Arc::new(DashMap::new()); + for protocol in [UdpSessionProtocol::WireGuard, UdpSessionProtocol::Quic] { + let (accepted, accepted_rx) = mpsc::channel(UDP_SESSION_QUEUE_CAPACITY); + accepts.insert( + protocol, + Arc::new(ClassifiedUdpSessionAccept { + accepted, + accepted_rx: TokioMutex::new(accepted_rx), + accept_enabled: std::sync::atomic::AtomicBool::new(false), + }), + ); + } + accepts +} + +impl UdpSessionLayer +where + S: VirtualUdpSocket, +{ + pub fn new(socket: Arc) -> Self { + Self::new_with_stun_responder(socket, Arc::new(NoopUdpSessionStunResponder)) + } +} + +impl UdpSessionLayer +where + S: VirtualUdpSocket, + R: UdpSessionStunResponder, +{ + pub fn new_with_stun_responder(socket: Arc, stun_responder: Arc) -> Self { + let sessions = Arc::new(DashMap::new()); + let classified_sessions = Arc::new(DashMap::new()); + let classified_accepts = create_classified_udp_session_accepts(); + let pending_connects = Arc::new(DashMap::new()); + let (mux_accepted_tx, mux_accepted_rx) = mpsc::channel(UDP_SESSION_QUEUE_CAPACITY); + let (control_tx, control_rx) = mpsc::channel(UDP_SESSION_QUEUE_CAPACITY); + let (session_shutdown_tx, _) = watch::channel(false); + let recv_task = tokio::spawn(udp_session_layer_recv_task( + socket.clone(), + sessions.clone(), + classified_sessions.clone(), + classified_accepts.clone(), + pending_connects.clone(), + mux_accepted_tx, + control_tx, + stun_responder.clone(), + session_shutdown_tx.clone(), + )); + + Self { + socket, + _stun_responder: stun_responder, + sessions, + classified_sessions, + classified_accepts, + pending_connects, + mux_accepted_rx: TokioMutex::new(mux_accepted_rx), + _control_rx: TokioMutex::new(control_rx), + session_shutdown_tx, + recv_task, + } + } + + pub fn local_addr(&self) -> io::Result { + self.socket.local_addr() + } + + pub fn active_session_count(&self) -> usize { + self.sessions.len() + } + + pub fn active_classified_session_count(&self) -> usize { + self.classified_sessions.len() + } + + pub fn open_classified_session( + &self, + protocol: UdpSessionProtocol, + remote_addr: SocketAddr, + ) -> io::Result { + let local_addr = self.socket.local_addr()?; + let key = ClassifiedUdpSessionKey::new(protocol, remote_addr); + let rings = create_udp_session_rings(); + match self.classified_sessions.entry(key) { + dashmap::mapref::entry::Entry::Vacant(entry) => { + entry.insert(udp_session_registry_entry(&rings)); + } + dashmap::mapref::entry::Entry::Occupied(_) => { + return Err(io::Error::new( + io::ErrorKind::AddrInUse, + format!("{protocol:?} udp session already exists for {remote_addr}"), + )); + } + } + + let close = UdpSessionClose::classified( + key, + rings.close_tx.clone(), + self.classified_sessions.clone(), + ); + Ok(UdpSession::new( + self.socket.clone(), + local_addr, + remote_addr, + protocol.session_kind(), + UdpSessionCodec::Identity, + rings, + close, + self.session_shutdown_tx.subscribe(), + )) + } + + pub async fn connect( + &self, + remote_addr: SocketAddr, + ) -> Result { + let local_addr = self.socket.local_addr()?; + let magic = rand::random(); + let (control_tx, mut control_rx) = mpsc::channel(UDP_SESSION_QUEUE_CAPACITY); + let (sack_tx, mut sack_rx) = watch::channel(None); + let rings = create_udp_session_rings(); + let session_key = Arc::new(StdMutex::new(None)); + let conn_id = loop { + let conn_id = rand::random(); + if self + .sessions + .contains_key(&UdpSessionKey::new(remote_addr, conn_id)) + { + continue; + } + + let pending = PendingUdpSessionConnect { + expected_addr: remote_addr, + magic, + session_key: session_key.clone(), + entry: udp_session_registry_entry(&rings), + control: control_tx.clone(), + sack: sack_tx.clone(), + }; + if let dashmap::mapref::entry::Entry::Vacant(entry) = + self.pending_connects.entry(conn_id) + { + entry.insert(pending); + break conn_id; + } + }; + let mut cleanup_guard = PendingUdpSessionGuard::new( + self.sessions.clone(), + self.pending_connects.clone(), + session_key, + conn_id, + ); + + let result = self + .connect_with_registered_attempt( + remote_addr, + conn_id, + magic, + &mut control_rx, + &mut sack_rx, + ) + .await; + + match result { + Ok(recv_addr) => { + let key = UdpSessionKey::new(recv_addr, conn_id); + if cleanup_guard.session_key() != Some(key) { + return Err(UdpSessionConnectError::InvalidPacket(format!( + "udp session was not registered: {key:?}" + ))); + } + cleanup_guard.set_session_key(key); + cleanup_guard.disarm_keep_session(); + + let close = + UdpSessionClose::easy_tier(key, rings.close_tx.clone(), self.sessions.clone()); + Ok(UdpSession::new( + self.socket.clone(), + local_addr, + key.peer_addr, + UdpSessionKind::EasyTierMux, + UdpSessionCodec::EasyTierData { conn_id }, + rings, + close, + self.session_shutdown_tx.subscribe(), + )) + } + Err(err) => Err(err), + } + } + + pub async fn accept(&self) -> io::Result { + let mut mux_accepted_rx = self.mux_accepted_rx.lock().await; + mux_accepted_rx + .recv() + .await + .ok_or_else(|| io::Error::new(io::ErrorKind::UnexpectedEof, "udp listener closed")) + } + + fn classified_accept( + &self, + protocol: UdpSessionProtocol, + ) -> io::Result> { + self.classified_accepts + .get(&protocol) + .map(|entry| entry.value().clone()) + .ok_or_else(|| { + io::Error::new( + io::ErrorKind::InvalidInput, + format!("{protocol:?} udp listener is not registered"), + ) + }) + } + + pub fn enable_classified_accept(&self, protocol: UdpSessionProtocol) -> io::Result<()> { + let accept = self.classified_accept(protocol)?; + accept.accept_enabled.store(true, Ordering::Relaxed); + Ok(()) + } + + pub async fn accept_classified_session( + &self, + protocol: UdpSessionProtocol, + ) -> io::Result { + let accept = self.classified_accept(protocol)?; + accept.accept_enabled.store(true, Ordering::Relaxed); + let mut accepted_rx = accept.accepted_rx.lock().await; + accepted_rx.recv().await.ok_or_else(|| { + io::Error::new( + io::ErrorKind::UnexpectedEof, + format!("{protocol:?} udp listener closed"), + ) + }) + } + + async fn connect_with_registered_attempt( + &self, + remote_addr: SocketAddr, + conn_id: u32, + magic: u64, + control_rx: &mut mpsc::Receiver, + sack_rx: &mut watch::Receiver>, + ) -> Result { + let syn_packet = new_syn_packet(conn_id, magic).into_bytes(); + self.socket.send_to(&syn_packet, remote_addr).await?; + + let timeout = crate::foundation::time::sleep(UDP_SESSION_CONNECT_TIMEOUT); + let resend_sleep = crate::foundation::time::sleep(UDP_SESSION_RESEND_INTERVAL); + tokio::pin!(timeout); + tokio::pin!(resend_sleep); + + loop { + if let Some(recv_addr) = *sack_rx.borrow_and_update() { + return Ok(recv_addr); + } + + tokio::select! { + biased; + sack = sack_rx.changed() => { + if sack.is_err() { + return Err(UdpSessionConnectError::InvalidPacket( + "udp sack channel closed".to_owned(), + )); + } + if let Some(recv_addr) = *sack_rx.borrow_and_update() { + return Ok(recv_addr); + } + } + _ = &mut timeout => return Err(UdpSessionConnectError::Timeout), + _ = &mut resend_sleep => { + self.socket.send_to(&syn_packet, remote_addr).await?; + resend_sleep + .as_mut() + .reset(crate::foundation::time::Instant::now() + UDP_SESSION_RESEND_INTERVAL); + } + control = control_rx.recv() => { + match control { + Some(UdpConnectControl::HolePunch { recv_addr }) => { + self.socket.send_to(&syn_packet, recv_addr).await?; + } + Some(UdpConnectControl::InvalidPacket(reason)) => { + tracing::debug!(?reason, "udp wait sack error"); + } + None => { + return Err(UdpSessionConnectError::InvalidPacket( + "udp connect control channel closed".to_owned(), + )); + } + } + } + } + } + } +} + +impl Drop for UdpSessionLayer { + fn drop(&mut self) { + let _ = self.session_shutdown_tx.send(true); + self.pending_connects.clear(); + close_all_udp_sessions(&self.sessions); + close_all_classified_udp_sessions(&self.classified_sessions); + self.recv_task.abort(); + } +} + +struct PendingUdpSessionGuard { + sessions: Arc, + pending_connects: Arc, + session_key: Arc>>, + conn_id: u32, + active: bool, +} + +impl PendingUdpSessionGuard { + fn new( + sessions: Arc, + pending_connects: Arc, + session_key: Arc>>, + conn_id: u32, + ) -> Self { + Self { + sessions, + pending_connects, + session_key, + conn_id, + active: true, + } + } + + fn session_key(&self) -> Option { + *self.session_key.lock().unwrap() + } + + fn set_session_key(&mut self, session_key: UdpSessionKey) { + *self.session_key.lock().unwrap() = Some(session_key); + } + + fn disarm_keep_session(mut self) { + self.pending_connects.remove(&self.conn_id); + self.active = false; + } +} + +impl Drop for PendingUdpSessionGuard { + fn drop(&mut self) { + if self.active { + self.pending_connects.remove(&self.conn_id); + if let Some(session_key) = self.session_key() { + close_udp_session(&self.sessions, session_key); + } + } + } +} + +fn move_pending_udp_session_sender( + sessions: &UdpSessionRegistry, + pending: &PendingUdpSessionConnect, + new_key: UdpSessionKey, +) -> bool { + let mut current_key = pending.session_key.lock().unwrap(); + if let Some(current_key) = *current_key { + return current_key == new_key; + } + + match sessions.entry(new_key) { + dashmap::mapref::entry::Entry::Vacant(entry) => { + entry.insert(pending.entry.clone()); + *current_key = Some(new_key); + true + } + dashmap::mapref::entry::Entry::Occupied(_) => false, + } +} + +#[allow(clippy::too_many_arguments)] +pub(super) async fn udp_session_layer_recv_task( + socket: Arc, + sessions: Arc, + classified_sessions: Arc, + classified_accepts: Arc, + pending_connects: Arc, + mux_accepted: mpsc::Sender, + control: mpsc::Sender, + stun_responder: Arc, + session_shutdown_tx: watch::Sender, +) where + S: VirtualUdpSocket, + R: UdpSessionStunResponder, +{ + let mut buf = [0u8; 65535]; + let control_permits = Arc::new(Semaphore::new(UDP_SESSION_QUEUE_CAPACITY)); + loop { + let (len, remote_addr, recv_meta) = match socket.recv_from_with_meta(&mut buf).await { + Ok(ret) => ret, + Err(err) => { + tracing::debug!(?err, "udp session recv loop stopped"); + let _ = session_shutdown_tx.send(true); + pending_connects.clear(); + close_all_udp_sessions(&sessions); + close_all_classified_udp_sessions(&classified_sessions); + break; + } + }; + + let payload = BytesMut::from(&buf[..len]); + let datagram = UdpSessionDatagram::new(payload.clone(), recv_meta); + let quic_key = ClassifiedUdpSessionKey::new(UdpSessionProtocol::Quic, remote_addr); + if classified_sessions.contains_key(&quic_key) { + dispatch_existing_classified_udp_datagram(&classified_sessions, quic_key, datagram); + continue; + } + match classify_udp_datagram(payload) { + UdpDatagramClassification::Stun(datagram_payload) => { + spawn_stun_control_handler( + socket.clone(), + stun_responder.clone(), + control_permits.clone(), + datagram_payload.clone(), + remote_addr, + ); + dispatch_control_packet( + &control, + UdpSessionLayerControl::Stun { + remote_addr, + datagram: datagram_payload, + }, + ); + } + UdpDatagramClassification::SessionPacket { + kind, + datagram: datagram_payload, + } => { + dispatch_session_udp_datagram( + socket.clone(), + &classified_sessions, + &classified_accepts, + session_shutdown_tx.subscribe(), + remote_addr, + kind, + UdpSessionDatagram::new(datagram_payload, recv_meta), + ); + } + UdpDatagramClassification::EasyTier { + kind, + conn_id, + packet, + fallback, + } => { + let consumed = dispatch_easy_tier_udp_datagram( + socket.clone(), + &sessions, + &pending_connects, + &mux_accepted, + &control, + control_permits.clone(), + remote_addr, + kind, + conn_id, + &packet, + recv_meta, + session_shutdown_tx.subscribe(), + ); + if !consumed { + dispatch_session_udp_datagram( + socket.clone(), + &classified_sessions, + &classified_accepts, + session_shutdown_tx.subscribe(), + remote_addr, + fallback, + UdpSessionDatagram::new(packet.into_bytes().into(), recv_meta), + ); + } + } + } + } +} + +fn dispatch_existing_classified_udp_datagram( + classified_sessions: &Arc, + key: ClassifiedUdpSessionKey, + datagram: UdpSessionDatagram, +) { + let Some(entry) = classified_sessions + .get(&key) + .map(|entry| entry.value().clone()) + else { + return; + }; + + if !dispatch_payload_to_session(&entry.incoming, datagram, UdpSessionEnqueuePolicy::Reliable) { + close_classified_udp_session(classified_sessions, key); + tracing::debug!(?key, "classified udp session data queue closed"); + } +} + +#[allow(clippy::too_many_arguments)] +fn dispatch_easy_tier_udp_datagram( + socket: Arc, + sessions: &Arc, + pending_connects: &Arc, + mux_accepted: &mpsc::Sender, + control: &mpsc::Sender, + control_permits: Arc, + remote_addr: SocketAddr, + kind: EasyTierUdpPacketKind, + conn_id: u32, + packet: &ZCPacket, + recv_meta: UdpSocketRecvMeta, + session_shutdown: watch::Receiver, +) -> bool +where + S: VirtualUdpSocket, +{ + match kind { + EasyTierUdpPacketKind::Data => { + dispatch_data_packet(sessions, remote_addr, conn_id, packet, recv_meta) + } + EasyTierUdpPacketKind::Syn => handle_new_easy_tier_mux_connect( + socket, + sessions.clone(), + mux_accepted.clone(), + remote_addr, + conn_id, + packet, + session_shutdown, + ), + EasyTierUdpPacketKind::Sack => { + dispatch_sack_packet(sessions, pending_connects, remote_addr, conn_id, packet) + } + EasyTierUdpPacketKind::HolePunch => { + dispatch_hole_punch_packet(pending_connects, remote_addr) + } + EasyTierUdpPacketKind::V4HolePunch => { + dispatch_v4_hole_punch_control(socket, control_permits, control, remote_addr, packet) + } + EasyTierUdpPacketKind::V6HolePunch => { + dispatch_v6_hole_punch_control(socket, control_permits, control, remote_addr, packet) + } + } +} + +pub(super) fn dispatch_data_packet( + sessions: &UdpSessionRegistry, + peer_addr: SocketAddr, + conn_id: u32, + packet: &ZCPacket, + recv_meta: UdpSocketRecvMeta, +) -> bool { + let key = UdpSessionKey::new(peer_addr, conn_id); + let Some(entry) = sessions.get(&key).map(|entry| entry.value().clone()) else { + return false; + }; + + let payload = UdpSessionDatagram::new(BytesMut::from(packet.udp_payload()), recv_meta); + let policy = if packet.is_lossy() { + UdpSessionEnqueuePolicy::Lossy + } else { + UdpSessionEnqueuePolicy::Reliable + }; + if !dispatch_payload_to_session(&entry.incoming, payload, policy) { + close_udp_session(sessions, key); + tracing::debug!(?key, "udp session data queue closed"); + } + true +} + +fn dispatch_session_udp_datagram( + socket: Arc, + classified_sessions: &Arc, + classified_accepts: &Arc, + session_shutdown: watch::Receiver, + remote_addr: SocketAddr, + kind: UdpSessionPacketKind, + datagram: UdpSessionDatagram, +) where + S: VirtualUdpSocket, +{ + match kind { + UdpSessionPacketKind::Classified(protocol) => dispatch_classified_udp_datagram( + socket, + classified_sessions, + classified_accepts, + protocol, + session_shutdown, + remote_addr, + datagram, + ), + UdpSessionPacketKind::Unknown => { + tracing::trace!(?remote_addr, "unknown udp packet has no session route"); + } + } +} + +fn dispatch_classified_udp_datagram( + socket: Arc, + classified_sessions: &Arc, + classified_accepts: &Arc, + protocol: UdpSessionProtocol, + session_shutdown: watch::Receiver, + remote_addr: SocketAddr, + datagram: UdpSessionDatagram, +) where + S: VirtualUdpSocket, +{ + let key = ClassifiedUdpSessionKey::new(protocol, remote_addr); + if let Some(entry) = classified_sessions + .get(&key) + .map(|entry| entry.value().clone()) + { + if !dispatch_payload_to_session( + &entry.incoming, + datagram, + UdpSessionEnqueuePolicy::Reliable, + ) { + close_classified_udp_session(classified_sessions, key); + tracing::debug!(?key, "classified udp session data queue closed"); + } + return; + } + + let Some(accept) = classified_accepts + .get(&protocol) + .map(|entry| entry.value().clone()) + else { + tracing::trace!( + ?protocol, + ?remote_addr, + "classified udp accept is not registered" + ); + return; + }; + + if !accept.accept_enabled.load(Ordering::Relaxed) { + return; + } + + let accept_permit = match accept.accepted.clone().try_reserve_owned() { + Ok(permit) => permit, + Err(err) => { + tracing::debug!(?err, ?key, "classified udp accept queue unavailable"); + return; + } + }; + let local_addr = match socket.local_addr() { + Ok(addr) => addr, + Err(err) => { + tracing::debug!(?err, ?key, "classified udp get local addr error"); + return; + } + }; + let rings = create_udp_session_rings(); + match classified_sessions.entry(key) { + dashmap::mapref::entry::Entry::Vacant(entry) => { + entry.insert(udp_session_registry_entry(&rings)); + } + dashmap::mapref::entry::Entry::Occupied(entry) => { + let entry = entry.get().clone(); + if !dispatch_payload_to_session( + &entry.incoming, + datagram, + UdpSessionEnqueuePolicy::Reliable, + ) { + close_classified_udp_session(classified_sessions, key); + tracing::debug!(?key, "classified udp session data queue closed"); + } + return; + } + } + if !dispatch_payload_to_session( + &rings.session_recv_tx, + datagram, + UdpSessionEnqueuePolicy::Reliable, + ) { + close_classified_udp_session(classified_sessions, key); + tracing::debug!(?key, "classified udp session data queue closed"); + return; + } + let close = + UdpSessionClose::classified(key, rings.close_tx.clone(), classified_sessions.clone()); + let session = UdpSession::new( + socket, + local_addr, + remote_addr, + protocol.session_kind(), + UdpSessionCodec::Identity, + rings, + close, + session_shutdown, + ); + accept_permit.send(session); +} + +pub(super) fn handle_new_easy_tier_mux_connect( + socket: Arc, + sessions: Arc, + mux_accepted: mpsc::Sender, + remote_addr: SocketAddr, + conn_id: u32, + packet: &ZCPacket, + session_shutdown: watch::Receiver, +) -> bool +where + S: VirtualUdpSocket, +{ + let payload = packet.udp_payload(); + if payload.len() != 8 { + tracing::warn!( + payload_len = payload.len(), + ?remote_addr, + ?conn_id, + "udp syn packet payload len not match", + ); + return false; + } + + let magic = u64::from_le_bytes(payload[..8].try_into().unwrap()); + let key = UdpSessionKey::new(remote_addr, conn_id); + let sack_packet = new_sack_packet(conn_id, magic).into_bytes(); + if sessions.contains_key(&key) { + let sessions = sessions.clone(); + tokio::spawn(async move { + if let Err(err) = socket.send_to(&sack_packet, remote_addr).await { + tracing::debug!(?err, ?key, "udp resend sack packet error"); + close_udp_session(&sessions, key); + } + }); + return true; + } + + let accept_permit = match mux_accepted.clone().try_reserve_owned() { + Ok(permit) => permit, + Err(err) => { + tracing::debug!(?err, ?key, "udp accept queue unavailable"); + return true; + } + }; + let local_addr = match socket.local_addr() { + Ok(addr) => addr, + Err(err) => { + tracing::debug!(?err, ?key, "udp get local addr for accepted session error"); + return true; + } + }; + let rings = create_udp_session_rings(); + sessions.insert(key, udp_session_registry_entry(&rings)); + let close = UdpSessionClose::easy_tier(key, rings.close_tx.clone(), sessions.clone()); + let session = UdpSession::new( + socket.clone(), + local_addr, + key.peer_addr, + UdpSessionKind::EasyTierMux, + UdpSessionCodec::EasyTierData { conn_id }, + rings, + close, + session_shutdown, + ); + tokio::spawn(async move { + if let Err(err) = socket.send_to(&sack_packet, remote_addr).await { + close_udp_session(&sessions, key); + tracing::debug!(?err, ?key, "udp send sack packet error"); + return; + } + + accept_permit.send(session); + }); + true +} + +pub(super) fn dispatch_sack_packet( + sessions: &UdpSessionRegistry, + pending_connects: &PendingUdpSessionConnects, + recv_addr: SocketAddr, + conn_id: u32, + packet: &ZCPacket, +) -> bool { + let payload = packet.udp_payload(); + if payload.len() != 8 { + if let Some(pending) = pending_connects + .get(&conn_id) + .map(|entry| entry.value().control.clone()) + { + let _ = pending.try_send(UdpConnectControl::InvalidPacket( + "udp sack packet payload len not match".to_owned(), + )); + return true; + } + return false; + } + + let magic = u64::from_le_bytes(payload[..8].try_into().unwrap()); + let Some((_, pending)) = pending_connects.remove_if(&conn_id, |_, pending| { + pending.magic == magic && (*pending.session_key.lock().unwrap()).is_none() + }) else { + if let Some(pending) = pending_connects + .get(&conn_id) + .map(|entry| entry.value().control.clone()) + { + let _ = pending.try_send(UdpConnectControl::InvalidPacket( + "udp sack magic not match".to_owned(), + )); + return true; + } + return false; + }; + + let new_key = UdpSessionKey::new(recv_addr, conn_id); + if !move_pending_udp_session_sender(sessions, &pending, new_key) { + let _ = pending.control.try_send(UdpConnectControl::InvalidPacket( + "udp session already exists".to_owned(), + )); + return true; + } + if pending.sack.send(Some(recv_addr)).is_err() { + close_udp_session(sessions, new_key); + } + true +} + +fn dispatch_hole_punch_packet( + pending_connects: &PendingUdpSessionConnects, + recv_addr: SocketAddr, +) -> bool { + let controls = pending_connects + .iter() + .filter(|entry| entry.value().expected_addr == recv_addr) + .map(|entry| entry.value().control.clone()) + .collect::>(); + if controls.is_empty() { + return false; + } + + for control in controls { + let _ = control.try_send(UdpConnectControl::HolePunch { recv_addr }); + } + true +} + +fn spawn_stun_control_handler( + socket: Arc, + stun_responder: Arc, + permits: Arc, + datagram: BytesMut, + remote_addr: SocketAddr, +) where + S: VirtualUdpSocket, + H: UdpSessionStunResponder, +{ + let Ok(permit) = permits.try_acquire_owned() else { + tracing::debug!(?remote_addr, "udp stun responder queue full"); + return; + }; + tokio::spawn(async move { + let _permit = permit; + if let Err(err) = stun_responder + .respond_stun(socket, &datagram, remote_addr) + .await + { + tracing::debug!(?err, ?remote_addr, "udp respond stun packet error"); + } + }); +} + +fn spawn_v4_hole_punch_control_handler( + socket: Arc, + permits: Arc, + remote_addr: SocketAddr, + dst_addr: SocketAddrV4, +) where + S: VirtualUdpSocket, +{ + let Ok(permit) = permits.try_acquire_owned() else { + tracing::debug!(?remote_addr, ?dst_addr, "udp control handler queue full"); + return; + }; + tokio::spawn(async move { + let _permit = permit; + let packet = new_hole_punch_packet(1, UDP_SESSION_HOLE_PUNCH_PACKET_BODY_LEN).into_bytes(); + if let Err(err) = socket + .send_to_with_meta( + &packet, + SocketAddr::V4(dst_addr), + UdpSocketSendMeta::default(), + ) + .await + { + tracing::debug!( + ?err, + ?remote_addr, + ?dst_addr, + "udp send v4 hole punch packet error" + ); + } + }); +} + +fn spawn_v6_hole_punch_control_handler( + socket: Arc, + permits: Arc, + remote_addr: SocketAddr, + dst_addr: SocketAddrV6, + preferred_src: Option, +) where + S: VirtualUdpSocket, +{ + let Ok(permit) = permits.try_acquire_owned() else { + tracing::debug!(?remote_addr, ?dst_addr, "udp control handler queue full"); + return; + }; + tokio::spawn(async move { + let _permit = permit; + let packet = new_hole_punch_packet(1, UDP_SESSION_HOLE_PUNCH_PACKET_BODY_LEN).into_bytes(); + if let Some(source) = preferred_src { + match socket + .send_to_with_meta( + &packet, + SocketAddr::V6(dst_addr), + UdpSocketSendMeta { + src_ip: Some(source.ip.into()), + src_ifindex: Some(source.ifindex), + }, + ) + .await + { + Ok(_) => return, + Err(error) => tracing::debug!( + ?source, + ?dst_addr, + ?error, + "udp preferred v6 source failed, falling back" + ), + } + } + if let Err(err) = socket + .send_to_with_meta( + &packet, + SocketAddr::V6(dst_addr), + UdpSocketSendMeta::default(), + ) + .await + { + tracing::debug!( + ?err, + ?remote_addr, + ?dst_addr, + ?preferred_src, + "udp send v6 hole punch packet error" + ); + } + }); +} + +pub(super) fn dispatch_v4_hole_punch_control( + socket: Arc, + permits: Arc, + control: &mpsc::Sender, + remote_addr: SocketAddr, + packet: &ZCPacket, +) -> bool +where + S: VirtualUdpSocket, +{ + if !remote_addr.ip().is_loopback() { + tracing::warn!(?remote_addr, "v4 hole punch packet should be from loopback"); + return false; + } + if !remote_addr.ip().is_ipv4() { + tracing::warn!( + ?remote_addr, + "v4 hole punch packet should be sent from ipv4" + ); + return false; + } + let Some(dst_addr) = extract_dst_addr_from_v4_hole_punch_packet(packet.udp_payload()) else { + tracing::debug!(?remote_addr, "invalid v4 hole punch packet"); + return false; + }; + spawn_v4_hole_punch_control_handler(socket, permits, remote_addr, dst_addr); + dispatch_control_packet( + control, + UdpSessionLayerControl::V4HolePunch { + remote_addr, + dst_addr, + }, + ); + true +} + +fn dispatch_v6_hole_punch_control( + socket: Arc, + permits: Arc, + control: &mpsc::Sender, + remote_addr: SocketAddr, + packet: &ZCPacket, +) -> bool +where + S: VirtualUdpSocket, +{ + if !remote_addr.ip().is_loopback() { + tracing::warn!(?remote_addr, "v6 hole punch packet should be from loopback"); + return false; + } + if !remote_addr.ip().is_ipv6() { + tracing::warn!( + ?remote_addr, + "v6 hole punch packet should be sent from ipv6" + ); + return false; + } + let Some((dst_addr, preferred_src)) = extract_v6_hole_punch_packet(packet.udp_payload()) else { + tracing::debug!(?remote_addr, "invalid v6 hole punch packet"); + return false; + }; + spawn_v6_hole_punch_control_handler(socket, permits, remote_addr, dst_addr, preferred_src); + dispatch_control_packet( + control, + UdpSessionLayerControl::V6HolePunch { + remote_addr, + dst_addr, + preferred_src, + }, + ); + true +} + +fn dispatch_control_packet( + control: &mpsc::Sender, + packet: UdpSessionLayerControl, +) { + if let Err(err) = control.try_send(packet) { + tracing::debug!(?err, "udp session control queue full"); + } +} + +#[derive(Debug)] +pub struct UdpSessionDialer { + factory: Arc, +} + +impl UdpSessionDialer +where + F: VirtualUdpSocketFactory, +{ + pub fn new(factory: Arc) -> Self { + Self { factory } + } +} + +#[async_trait] +impl UdpSessionConnector for UdpSessionDialer +where + F: VirtualUdpSocketFactory, +{ + type Session = UdpSession; + + async fn connect( + &mut self, + request: UdpSessionConnectRequest, + ) -> anyhow::Result { + let socket = self.factory.bind_udp(request.bind).await?; + let layer = Arc::new(UdpSessionLayer::new_with_stun_responder( + socket, + self.factory.clone(), + )); + let mut session = layer.open_classified_session(request.protocol, request.remote_addr)?; + session._cleanup.layer_guard = Some(Box::new(layer)); + Ok(session) + } +} + +#[cfg(test)] +mod tests { + use super::*; + + impl UdpSessionLayer + where + S: VirtualUdpSocket, + R: UdpSessionStunResponder, + { + pub(crate) async fn recv_control(&self) -> io::Result { + let mut control_rx = self._control_rx.lock().await; + control_rx + .recv() + .await + .ok_or_else(|| io::Error::new(io::ErrorKind::UnexpectedEof, "udp listener closed")) + } + } +} diff --git a/easytier-core/src/socket/udp/listener.rs b/easytier-core/src/socket/udp/listener.rs new file mode 100644 index 00000000..27c4e634 --- /dev/null +++ b/easytier-core/src/socket/udp/listener.rs @@ -0,0 +1,205 @@ +use std::{ + fmt, io, + net::SocketAddr, + sync::{Arc, Mutex as StdMutex, Weak}, +}; + +use async_trait::async_trait; +use serde::{Deserialize, Serialize}; +use url::Url; + +use crate::socket::{ListenerConnectionCounter, SocketListener}; + +use super::{ + UdpBindOptions, UdpSession, UdpSessionLayer, UdpSessionListenRequest, UdpSessionProtocol, + UdpSessionStunResponder, VirtualUdpSocket, VirtualUdpSocketFactory, +}; + +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] +pub enum UdpSessionAcceptKind { + EasyTierMux, + Classified(UdpSessionProtocol), +} + +pub async fn accept_udp_session( + layer: &Arc>, + accept_kind: UdpSessionAcceptKind, +) -> io::Result +where + S: VirtualUdpSocket, + R: UdpSessionStunResponder, +{ + match accept_kind { + UdpSessionAcceptKind::EasyTierMux => layer.accept().await, + UdpSessionAcceptKind::Classified(protocol) => { + layer.accept_classified_session(protocol).await + } + } +} + +type Layer = UdpSessionLayer<::Socket, F>; + +pub struct UdpSessionSocketListener +where + F: VirtualUdpSocketFactory, +{ + url: Url, + request: UdpSessionListenRequest, + accept_kind: UdpSessionAcceptKind, + factory: Arc, + socket: Option>, + layer: Option>>, + layer_ref: Arc>>>>, +} + +impl UdpSessionSocketListener +where + F: VirtualUdpSocketFactory, +{ + pub fn new(url: Url, local_addr: SocketAddr, factory: Arc) -> Self { + let request = UdpSessionListenRequest::new( + UdpBindOptions::port_bound_listener(local_addr).with_only_v6(true), + ); + Self::new_with_request(url, request, UdpSessionAcceptKind::EasyTierMux, factory) + } + + pub fn new_with_request( + url: Url, + request: UdpSessionListenRequest, + accept_kind: UdpSessionAcceptKind, + factory: Arc, + ) -> Self { + Self { + url, + request, + accept_kind, + factory, + socket: None, + layer: None, + layer_ref: Arc::new(StdMutex::new(None)), + } + } + + fn layer(&self) -> anyhow::Result>> { + self.layer + .clone() + .ok_or_else(|| anyhow::anyhow!("udp session listener is not started")) + } + + pub async fn accept_session(&self) -> anyhow::Result { + let layer = self.layer()?; + let mut session = accept_udp_session(&layer, self.accept_kind).await?; + session.keep_layer_alive(layer); + Ok(session) + } +} + +impl fmt::Debug for UdpSessionSocketListener +where + F: VirtualUdpSocketFactory, +{ + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.debug_struct("UdpSessionSocketListener") + .field("url", &self.url) + .field("request", &self.request) + .field("accept_kind", &self.accept_kind) + .field("listening", &self.socket.is_some()) + .finish() + } +} + +#[async_trait] +impl SocketListener for UdpSessionSocketListener +where + F: VirtualUdpSocketFactory, +{ + type Accepted = UdpSession; + + async fn listen(&mut self) -> anyhow::Result<()> { + if self.layer.is_some() { + return Ok(()); + } + + let socket = self.factory.bind_udp(self.request.bind.clone()).await?; + let local_addr = socket.local_addr()?; + self.url + .set_port(Some(local_addr.port())) + .map_err(|_| anyhow::anyhow!("failed to update udp listener port for {}", self.url))?; + + let layer = Arc::new(UdpSessionLayer::new_with_stun_responder( + socket.clone(), + self.factory.clone(), + )); + if let UdpSessionAcceptKind::Classified(protocol) = self.accept_kind { + layer.enable_classified_accept(protocol)?; + } + + *self.layer_ref.lock().unwrap() = Some(Arc::downgrade(&layer)); + self.socket = Some(socket); + self.layer = Some(layer); + Ok(()) + } + + async fn accept(&mut self) -> anyhow::Result { + self.accept_session().await + } + + fn local_url(&self) -> Url { + self.url.clone() + } + + fn connection_counter(&self) -> Arc { + Arc::new(UdpSessionConnectionCounter { + layer: self.layer_ref.clone(), + }) + } +} + +struct UdpSessionConnectionCounter +where + F: VirtualUdpSocketFactory, +{ + layer: Arc>>>>, +} + +impl fmt::Debug for UdpSessionConnectionCounter +where + F: VirtualUdpSocketFactory, +{ + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.debug_struct("UdpSessionConnectionCounter") + .field("active", &self.get()) + .finish() + } +} + +impl ListenerConnectionCounter for UdpSessionConnectionCounter +where + F: VirtualUdpSocketFactory, +{ + fn get(&self) -> Option { + let layer = self.layer.lock().unwrap(); + let Some(layer) = layer.as_ref().and_then(Weak::upgrade) else { + return Some(0); + }; + let active = layer.active_session_count() + layer.active_classified_session_count(); + Some(active as u32) + } +} + +#[cfg(any(test, feature = "test-utils"))] +mod test_utils { + use super::*; + + impl UdpSessionSocketListener + where + F: VirtualUdpSocketFactory, + { + #[doc(hidden)] + pub fn bound_socket(&self) -> anyhow::Result> { + self.socket + .clone() + .ok_or_else(|| anyhow::anyhow!("udp session listener is not started")) + } + } +} diff --git a/easytier-core/src/socket/udp/mod.rs b/easytier-core/src/socket/udp/mod.rs new file mode 100644 index 00000000..cf181c02 --- /dev/null +++ b/easytier-core/src/socket/udp/mod.rs @@ -0,0 +1,35 @@ +mod layer; +mod listener; +mod packet; +mod session; +mod virtual_socket; + +#[cfg(test)] +mod tests; + +const UDP_SESSION_RESEND_INTERVAL: std::time::Duration = std::time::Duration::from_millis(200); +const UDP_SESSION_CONNECT_TIMEOUT: std::time::Duration = std::time::Duration::from_secs(3); +const UDP_SESSION_QUEUE_CAPACITY: usize = 128; + +pub use layer::{UdpSessionDialer, UdpSessionLayer}; +pub use listener::{UdpSessionAcceptKind, UdpSessionSocketListener, accept_udp_session}; +pub use packet::{ + UdpSessionPacketError, extract_dst_addr_from_v4_hole_punch_packet, + extract_v6_hole_punch_packet, is_stun_packet, new_sack_packet, new_syn_packet, + new_v4_hole_punch_packet, new_v6_hole_punch_packet, parse_quic_initial_dcid, + parse_udp_session_datagram, +}; +pub use session::{ + UdpSession, UdpSessionConnectError, UdpSessionConnectRequest, UdpSessionConnector, + UdpSessionKind, UdpSessionLayerControl, UdpSessionListenRequest, UdpSessionListener, + UdpSessionProtocol, UdpSessionRecvMeta, UdpSessionSocket, +}; +pub(crate) use session::{ + UdpSessionCleanup, UdpSessionCodec, UdpSessionDatagram, UdpSessionOutbound, + UdpSessionTunnelParts, +}; +pub use virtual_socket::{ + NoopUdpSessionStunResponder, PreferredIpv6Source, UdpBindOptions, UdpSessionStunResponder, + UdpSocketPurpose, UdpSocketRecvMeta, UdpSocketSendMeta, VirtualUdpSocket, + VirtualUdpSocketFactory, send_v4_hole_punch_control_packet, send_v6_hole_punch_control_packet, +}; diff --git a/easytier-core/src/socket/udp/packet.rs b/easytier-core/src/socket/udp/packet.rs new file mode 100644 index 00000000..8b101a46 --- /dev/null +++ b/easytier-core/src/socket/udp/packet.rs @@ -0,0 +1,407 @@ +use std::{ + io, + net::{Ipv4Addr, Ipv6Addr, SocketAddrV4, SocketAddrV6}, +}; + +use bytes::BytesMut; +use zerocopy::{AsBytes, FromBytes}; + +use crate::packet::{ + UDP_TUNNEL_HEADER_SIZE, UDPTunnelHeader, UdpPacketType, V4HolePunchPacket, V6HolePunchPacket, + ZCPacket, ZCPacketType, +}; + +use super::{session::UdpSessionProtocol, virtual_socket::PreferredIpv6Source}; + +#[derive(Debug, thiserror::Error)] +pub enum UdpSessionPacketError { + #[error("udp packet size too small: {datagram_size:?}, packet: {packet:?}")] + TooSmall { + datagram_size: usize, + packet: BytesMut, + }, + #[error( + "udp packet payload len not match: header len: {header_len:?}, real len: {datagram_size:?}" + )] + PayloadLenMismatch { + header_len: usize, + datagram_size: usize, + }, +} + +pub(super) fn new_udp_packet(f: F, udp_body: &[u8]) -> ZCPacket +where + F: FnOnce(&mut UDPTunnelHeader), +{ + let mut buf = BytesMut::new(); + buf.resize(UDP_TUNNEL_HEADER_SIZE + udp_body.len(), 0); + buf[UDP_TUNNEL_HEADER_SIZE..].copy_from_slice(udp_body); + + let mut ret = ZCPacket::new_from_buf(buf, ZCPacketType::UDP); + let header = ret.mut_udp_tunnel_header().unwrap(); + f(header); + ret +} + +pub fn new_syn_packet(conn_id: u32, magic: u64) -> ZCPacket { + new_udp_packet( + |header| { + header.msg_type = UdpPacketType::Syn as u8; + header.conn_id.set(conn_id); + header.len.set(8); + }, + &magic.to_le_bytes(), + ) +} + +pub fn new_sack_packet(conn_id: u32, magic: u64) -> ZCPacket { + new_udp_packet( + |header| { + header.msg_type = UdpPacketType::Sack as u8; + header.conn_id.set(conn_id); + header.len.set(8); + }, + &magic.to_le_bytes(), + ) +} + +pub(super) fn new_data_packet(conn_id: u32, payload: &[u8]) -> io::Result { + let len = udp_session_payload_len(payload)?; + + Ok(new_udp_packet( + |header| { + header.msg_type = UdpPacketType::Data as u8; + header.conn_id.set(conn_id); + header.len.set(len); + }, + payload, + )) +} + +pub(super) fn udp_session_payload_len(payload: &[u8]) -> io::Result { + u16::try_from(payload.len()).map_err(|_| { + io::Error::new( + io::ErrorKind::InvalidInput, + format!("udp session payload too large: {}", payload.len()), + ) + }) +} + +pub fn new_v6_hole_punch_packet( + dst: &SocketAddrV6, + preferred_src: Option, +) -> ZCPacket { + let mut body = V6HolePunchPacket::default(); + body.dst_ipv6.copy_from_slice(&dst.ip().octets()); + body.dst_port.set(dst.port()); + if let Some(src) = preferred_src { + body.preferred_src_ipv6.copy_from_slice(&src.ip.octets()); + body.preferred_src_ifindex.set(src.ifindex); + } + new_udp_packet( + |header| { + header.msg_type = UdpPacketType::V6HolePunch as u8; + header.conn_id.set(dst.port() as u32); + header + .len + .set(std::mem::size_of::() as u16); + }, + body.as_bytes(), + ) +} + +pub fn new_v4_hole_punch_packet(dst: &SocketAddrV4) -> ZCPacket { + let mut body = V4HolePunchPacket::default(); + body.dst_ipv4.copy_from_slice(&dst.ip().octets()); + body.dst_port.set(dst.port()); + new_udp_packet( + |header| { + header.msg_type = UdpPacketType::V4HolePunch as u8; + header.conn_id.set(dst.port() as u32); + header + .len + .set(std::mem::size_of::() as u16); + }, + body.as_bytes(), + ) +} + +pub fn extract_dst_addr_from_v4_hole_punch_packet(buf: &[u8]) -> Option { + let body = V4HolePunchPacket::ref_from_prefix(buf)?; + let ip = Ipv4Addr::from(body.dst_ipv4); + Some(SocketAddrV4::new(ip, body.dst_port.get())) +} + +pub fn extract_v6_hole_punch_packet( + buf: &[u8], +) -> Option<(SocketAddrV6, Option)> { + let body = V6HolePunchPacket::ref_from_prefix(buf)?; + let ip = Ipv6Addr::from(body.dst_ipv6); + let preferred_src_ipv6 = Ipv6Addr::from(body.preferred_src_ipv6); + let preferred_src = (!preferred_src_ipv6.is_unspecified()).then_some(PreferredIpv6Source { + ip: preferred_src_ipv6, + ifindex: body.preferred_src_ifindex.get(), + }); + Some(( + SocketAddrV6::new(ip, body.dst_port.get(), 0, 0), + preferred_src, + )) +} + +pub fn is_stun_packet(data: &[u8]) -> bool { + data.len() >= UDP_TUNNEL_HEADER_SIZE + && data[4..8] == [0x21, 0x12, 0xA4, 0x42] + && data[0] & 0xC0 == 0 +} + +#[derive(Debug)] +pub(super) enum UdpDatagramClassification { + Stun(BytesMut), + EasyTier { + kind: EasyTierUdpPacketKind, + conn_id: u32, + packet: ZCPacket, + fallback: UdpSessionPacketKind, + }, + SessionPacket { + kind: UdpSessionPacketKind, + datagram: BytesMut, + }, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub(super) enum UdpSessionPacketKind { + Classified(UdpSessionProtocol), + Unknown, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub(super) enum EasyTierUdpPacketKind { + Data, + Syn, + Sack, + HolePunch, + V4HolePunch, + V6HolePunch, +} + +impl EasyTierUdpPacketKind { + fn from_msg_type(msg_type: u8) -> Option { + match msg_type { + msg_type if msg_type == UdpPacketType::Data as u8 => Some(Self::Data), + msg_type if msg_type == UdpPacketType::Syn as u8 => Some(Self::Syn), + msg_type if msg_type == UdpPacketType::Sack as u8 => Some(Self::Sack), + msg_type if msg_type == UdpPacketType::HolePunch as u8 => Some(Self::HolePunch), + msg_type if msg_type == UdpPacketType::V4HolePunch as u8 => Some(Self::V4HolePunch), + msg_type if msg_type == UdpPacketType::V6HolePunch as u8 => Some(Self::V6HolePunch), + _ => None, + } + } +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub(super) struct EasyTierUdpDatagramInfo { + pub(super) kind: EasyTierUdpPacketKind, + pub(super) conn_id: u32, +} + +#[derive(Debug)] +pub(super) enum EasyTierUdpDatagramInspectError { + TooSmall { + datagram_size: usize, + }, + PayloadLenMismatch { + header_len: usize, + datagram_size: usize, + }, +} + +fn classify_session_udp_datagram(data: &[u8]) -> UdpSessionPacketKind { + if is_wireguard_packet(data) { + UdpSessionPacketKind::Classified(UdpSessionProtocol::WireGuard) + } else if is_quic_packet(data) { + UdpSessionPacketKind::Classified(UdpSessionProtocol::Quic) + } else { + UdpSessionPacketKind::Unknown + } +} + +fn is_wireguard_packet(data: &[u8]) -> bool { + if data.len() < 4 { + return false; + } + + let msg_type = u32::from_le_bytes(data[..4].try_into().unwrap()); + match msg_type { + 1 => data.len() == 148, + 2 => data.len() == 92, + 3 => data.len() == 64, + 4 => data.len() >= 32, + _ => false, + } +} + +fn parse_quic_varint(data: &[u8]) -> Option<(u64, usize)> { + let first = *data.first()?; + let len = 1usize << (first >> 6); + if data.len() < len { + return None; + } + + let mut value = u64::from(first & 0x3f); + for byte in &data[1..len] { + value = (value << 8) | u64::from(*byte); + } + Some((value, len)) +} + +pub fn parse_quic_initial_dcid(data: &[u8]) -> Option> { + const QUIC_INITIAL_HEADER_FORM_AND_FIXED_BIT: u8 = 0xC0; + const QUIC_LONG_PACKET_TYPE_MASK: u8 = 0x30; + const QUIC_MIN_INITIAL_DATAGRAM_LEN: usize = 1200; + const QUIC_MAX_CID_LEN: usize = 20; + + let first = *data.first()?; + if (first & QUIC_INITIAL_HEADER_FORM_AND_FIXED_BIT) != QUIC_INITIAL_HEADER_FORM_AND_FIXED_BIT + || (first & QUIC_LONG_PACKET_TYPE_MASK) != 0 + || data.len() < QUIC_MIN_INITIAL_DATAGRAM_LEN + { + return None; + } + + let version = data.get(1..5)?; + if version == [0, 0, 0, 0] { + return None; + } + + let dcid_len = usize::from(*data.get(5)?); + if dcid_len == 0 || dcid_len > QUIC_MAX_CID_LEN { + return None; + } + let dcid_start = 6; + let dcid_end = dcid_start + dcid_len; + let dcid = data.get(dcid_start..dcid_end)?; + + let scid_len = usize::from(*data.get(dcid_end)?); + if scid_len > QUIC_MAX_CID_LEN { + return None; + } + let token_len_offset = dcid_end + 1 + scid_len; + let (token_len, token_len_size) = parse_quic_varint(data.get(token_len_offset..)?)?; + let packet_len_offset = token_len_offset + token_len_size + usize::try_from(token_len).ok()?; + let (packet_len, packet_len_size) = parse_quic_varint(data.get(packet_len_offset..)?)?; + let packet_offset = packet_len_offset + packet_len_size; + if packet_len == 0 + || data.len().saturating_sub(packet_offset) < usize::try_from(packet_len).ok()? + { + return None; + } + + Some(dcid.to_vec()) +} + +fn is_quic_packet(data: &[u8]) -> bool { + parse_quic_initial_dcid(data).is_some() +} + +pub(super) fn inspect_easytier_udp_datagram( + data: &[u8], +) -> Result, EasyTierUdpDatagramInspectError> { + let datagram_size = data.len(); + if datagram_size < UDP_TUNNEL_HEADER_SIZE { + return Err(EasyTierUdpDatagramInspectError::TooSmall { datagram_size }); + } + + let header = UDPTunnelHeader::ref_from_prefix(data).unwrap(); + let header_len = header.len.get() as usize; + let real_len = datagram_size - UDP_TUNNEL_HEADER_SIZE; + if header_len != real_len { + return Err(EasyTierUdpDatagramInspectError::PayloadLenMismatch { + header_len, + datagram_size, + }); + } + + Ok( + EasyTierUdpPacketKind::from_msg_type(header.msg_type).map(|kind| EasyTierUdpDatagramInfo { + kind, + conn_id: header.conn_id.get(), + }), + ) +} + +pub(super) fn classify_udp_datagram(datagram: BytesMut) -> UdpDatagramClassification { + if is_stun_packet(&datagram) { + return UdpDatagramClassification::Stun(datagram); + } + + let fallback = classify_session_udp_datagram(&datagram); + let easytier = match inspect_easytier_udp_datagram(&datagram) { + Ok(Some(easytier)) => easytier, + Ok(None) => { + return UdpDatagramClassification::SessionPacket { + kind: fallback, + datagram, + }; + } + Err(err) => { + match err { + EasyTierUdpDatagramInspectError::TooSmall { datagram_size } => { + tracing::debug!(datagram_size, "udp session packet too small"); + } + EasyTierUdpDatagramInspectError::PayloadLenMismatch { + header_len, + datagram_size, + } => { + tracing::debug!( + header_len, + datagram_size, + "udp session packet payload len mismatch" + ); + } + } + return UdpDatagramClassification::SessionPacket { + kind: fallback, + datagram, + }; + } + }; + let packet = ZCPacket::new_from_buf(datagram, ZCPacketType::UDP); + + UdpDatagramClassification::EasyTier { + kind: easytier.kind, + conn_id: easytier.conn_id, + packet, + fallback, + } +} + +pub fn parse_udp_session_datagram( + buf: BytesMut, + allow_stun: bool, +) -> Result { + let datagram_size = buf.len(); + if datagram_size < UDP_TUNNEL_HEADER_SIZE { + return Err(UdpSessionPacketError::TooSmall { + datagram_size, + packet: buf, + }); + } + + if allow_stun && is_stun_packet(&buf[..UDP_TUNNEL_HEADER_SIZE]) { + return Ok(ZCPacket::new_from_buf(buf, ZCPacketType::UDP)); + } + + let zc_packet = ZCPacket::new_from_buf(buf, ZCPacketType::UDP); + let header = zc_packet.udp_tunnel_header().unwrap(); + let header_len = header.len.get() as usize; + let real_len = datagram_size - UDP_TUNNEL_HEADER_SIZE; + if header_len != real_len { + return Err(UdpSessionPacketError::PayloadLenMismatch { + header_len, + datagram_size, + }); + } + + Ok(zc_packet) +} diff --git a/easytier-core/src/socket/udp/session.rs b/easytier-core/src/socket/udp/session.rs new file mode 100644 index 00000000..36874ab8 --- /dev/null +++ b/easytier-core/src/socket/udp/session.rs @@ -0,0 +1,785 @@ +use std::{ + io, + net::{IpAddr, SocketAddr, SocketAddrV4, SocketAddrV6}, + sync::{Arc, Mutex as StdMutex}, +}; + +use async_trait::async_trait; +use bytes::BytesMut; +use dashmap::DashMap; +use futures::{SinkExt, StreamExt}; +use serde::{Deserialize, Serialize}; +use tokio::{ + sync::{Mutex as TokioMutex, mpsc, oneshot, watch}, + task::JoinHandle, +}; + +use crate::socket::ring::{RingSocket, RingSocketReceiver, RingSocketSendError, RingSocketSender}; + +use super::{ + UDP_SESSION_QUEUE_CAPACITY, + packet::{new_data_packet, udp_session_payload_len}, + virtual_socket::{PreferredIpv6Source, UdpBindOptions, UdpSocketRecvMeta, VirtualUdpSocket}, +}; + +#[derive(Debug, Clone, PartialEq, Eq)] +pub(crate) struct UdpSessionDatagram { + pub(crate) payload: BytesMut, + pub(crate) dst_ip: Option, +} + +impl UdpSessionDatagram { + pub(crate) fn new(payload: BytesMut, meta: UdpSocketRecvMeta) -> Self { + Self { + payload, + dst_ip: meta.dst_ip, + } + } +} + +impl From for UdpSessionDatagram { + fn from(payload: BytesMut) -> Self { + Self { + payload, + dst_ip: None, + } + } +} + +#[derive(Debug, Default, Clone, Copy, PartialEq, Eq)] +pub struct UdpSessionRecvMeta { + pub dst_ip: Option, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum UdpSessionKind { + EasyTierMux, + WireGuard, + Quic, +} + +#[derive(Debug, thiserror::Error)] +pub enum UdpSessionConnectError { + #[error("io error: {0}")] + Io(#[from] io::Error), + #[error("timeout")] + Timeout, + #[error("invalid packet: {0}")] + InvalidPacket(String), +} + +#[async_trait] +pub trait UdpSessionSocket: Send + Sync + 'static { + fn kind(&self) -> UdpSessionKind; + + fn local_addr(&self) -> std::io::Result; + + fn peer_addr(&self) -> std::io::Result; + + async fn send(&self, data: &[u8]) -> std::io::Result; + + async fn recv(&self, buf: &mut [u8]) -> std::io::Result; + + async fn recv_with_meta(&self, buf: &mut [u8]) -> std::io::Result<(usize, UdpSessionRecvMeta)> { + let len = self.recv(buf).await?; + Ok((len, UdpSessionRecvMeta::default())) + } +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)] +pub enum UdpSessionProtocol { + WireGuard, + Quic, +} + +impl UdpSessionProtocol { + pub(super) fn session_kind(self) -> UdpSessionKind { + match self { + Self::WireGuard => UdpSessionKind::WireGuard, + Self::Quic => UdpSessionKind::Quic, + } + } +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct UdpSessionConnectRequest { + pub remote_addr: SocketAddr, + pub bind: UdpBindOptions, + pub protocol: UdpSessionProtocol, +} + +impl UdpSessionConnectRequest { + pub fn wireguard(remote_addr: SocketAddr) -> Self { + Self { + remote_addr, + bind: UdpBindOptions::direct_connect(), + protocol: UdpSessionProtocol::WireGuard, + } + } + + pub fn with_bind(mut self, bind: UdpBindOptions) -> Self { + self.bind = bind; + self + } +} + +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +pub struct UdpSessionListenRequest { + pub bind: UdpBindOptions, +} + +impl UdpSessionListenRequest { + pub fn new(bind: UdpBindOptions) -> Self { + Self { bind } + } +} + +#[async_trait] +pub trait UdpSessionConnector: Send { + type Session: UdpSessionSocket; + + async fn connect(&mut self, request: UdpSessionConnectRequest) + -> anyhow::Result; +} + +#[async_trait] +pub trait UdpSessionListener: Send { + type Session: UdpSessionSocket; + + async fn listen(&mut self, request: UdpSessionListenRequest) -> anyhow::Result<()>; + + fn local_addr(&self) -> std::io::Result; + + async fn accept(&mut self) -> anyhow::Result; +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] +pub(super) struct UdpSessionKey { + pub(super) peer_addr: SocketAddr, + pub(super) conn_id: u32, +} + +impl UdpSessionKey { + pub(super) fn new(peer_addr: SocketAddr, conn_id: u32) -> Self { + Self { peer_addr, conn_id } + } +} + +pub(super) type UdpSessionRegistry = DashMap; +pub(super) type ClassifiedUdpSessionRegistry = + DashMap; +pub(super) type ClassifiedUdpSessionAccepts = + DashMap>; +pub(super) type PendingUdpSessionConnects = DashMap; + +#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] +pub(super) struct ClassifiedUdpSessionKey { + pub(super) protocol: UdpSessionProtocol, + pub(super) peer_addr: SocketAddr, +} + +impl ClassifiedUdpSessionKey { + pub(super) fn new(protocol: UdpSessionProtocol, peer_addr: SocketAddr) -> Self { + Self { + protocol, + peer_addr, + } + } +} + +#[derive(Debug)] +pub(super) struct ClassifiedUdpSessionAccept { + pub(super) accepted: mpsc::Sender, + pub(super) accepted_rx: TokioMutex>, + pub(super) accept_enabled: std::sync::atomic::AtomicBool, +} + +#[derive(Debug, Clone)] +pub(super) struct UdpSessionRegistryEntry { + pub(super) incoming: Arc>>, + pub(super) close: watch::Sender, +} + +#[derive(Debug, Clone)] +pub(super) struct PendingUdpSessionConnect { + pub(super) expected_addr: SocketAddr, + pub(super) magic: u64, + pub(super) session_key: Arc>>, + pub(super) entry: UdpSessionRegistryEntry, + pub(super) control: mpsc::Sender, + pub(super) sack: watch::Sender>, +} + +#[derive(Debug)] +pub(super) enum UdpConnectControl { + HolePunch { recv_addr: SocketAddr }, + InvalidPacket(String), +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub enum UdpSessionLayerControl { + Stun { + remote_addr: SocketAddr, + datagram: BytesMut, + }, + V4HolePunch { + remote_addr: SocketAddr, + dst_addr: SocketAddrV4, + }, + V6HolePunch { + remote_addr: SocketAddr, + dst_addr: SocketAddrV6, + preferred_src: Option, + }, +} + +#[derive(Debug)] +pub struct UdpSession { + local_addr: SocketAddr, + peer_addr: SocketAddr, + kind: UdpSessionKind, + codec: UdpSessionCodec, + incoming: TokioMutex>, + outgoing: TokioMutex>, + closed: watch::Receiver, + pub(super) _cleanup: UdpSessionCleanup, +} + +pub(crate) struct UdpSessionOutbound { + pub(crate) payload: BytesMut, + pub(crate) completion: oneshot::Sender>, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub(crate) enum UdpSessionCodec { + EasyTierData { conn_id: u32 }, + Identity, +} + +impl UdpSessionCodec { + pub(crate) fn validate_payload(&self, payload: &[u8]) -> io::Result<()> { + if matches!(self, Self::EasyTierData { .. }) { + udp_session_payload_len(payload)?; + } + Ok(()) + } + + fn encode(&self, payload: &[u8]) -> io::Result { + match self { + Self::EasyTierData { conn_id } => { + Ok(new_data_packet(*conn_id, payload)?.into_bytes().into()) + } + Self::Identity => Ok(BytesMut::from(payload)), + } + } +} + +#[derive(Clone)] +enum UdpSessionCloseTarget { + #[cfg(test)] + SignalOnly, + EasyTier { + key: UdpSessionKey, + sessions: Arc, + }, + Classified { + key: ClassifiedUdpSessionKey, + sessions: Arc, + }, +} + +#[derive(Clone)] +pub(super) struct UdpSessionClose { + close: watch::Sender, + target: UdpSessionCloseTarget, +} + +impl UdpSessionClose { + pub(super) fn easy_tier( + key: UdpSessionKey, + close: watch::Sender, + sessions: Arc, + ) -> Self { + Self { + close, + target: UdpSessionCloseTarget::EasyTier { key, sessions }, + } + } + + pub(super) fn classified( + key: ClassifiedUdpSessionKey, + close: watch::Sender, + sessions: Arc, + ) -> Self { + Self { + close, + target: UdpSessionCloseTarget::Classified { key, sessions }, + } + } + + fn close(&self) { + match &self.target { + #[cfg(test)] + UdpSessionCloseTarget::SignalOnly => {} + UdpSessionCloseTarget::EasyTier { key, sessions } => { + close_udp_session(sessions, *key); + } + UdpSessionCloseTarget::Classified { key, sessions } => { + close_classified_udp_session(sessions, *key); + } + } + let _ = self.close.send(true); + } +} + +pub(crate) struct UdpSessionCleanup { + session_close: Option, + shutdown: Option>, + tasks: Vec>, + pub(super) layer_guard: Option>, +} + +impl std::fmt::Debug for UdpSessionCleanup { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("UdpSessionCleanup") + .field("has_session_close", &self.session_close.is_some()) + .field("has_shutdown", &self.shutdown.is_some()) + .field("tasks", &self.tasks.len()) + .field("has_layer_guard", &self.layer_guard.is_some()) + .finish() + } +} + +impl Drop for UdpSessionCleanup { + fn drop(&mut self) { + if let Some(shutdown) = &self.shutdown { + let _ = shutdown.send(true); + } + if let Some(close) = &self.session_close { + close.close(); + } + for task in &self.tasks { + task.abort(); + } + } +} + +impl UdpSession { + #[allow(clippy::too_many_arguments)] + pub(super) fn new( + socket: Arc, + local_addr: SocketAddr, + peer_addr: SocketAddr, + kind: UdpSessionKind, + codec: UdpSessionCodec, + rings: UdpSessionRingParts, + close: UdpSessionClose, + shutdown: watch::Receiver, + ) -> Self + where + S: VirtualUdpSocket, + { + if *shutdown.borrow() { + let _ = rings.close_tx.send(true); + } + let send_task = tokio::spawn(forward_udp_session_to_socket( + socket, + peer_addr, + codec, + rings.session_send_rx, + shutdown, + close.clone(), + )); + + Self { + local_addr, + peer_addr, + kind, + codec, + incoming: TokioMutex::new(rings.session_recv_rx), + outgoing: TokioMutex::new(rings.session_send_tx), + closed: rings.close_rx, + _cleanup: UdpSessionCleanup { + session_close: Some(close), + shutdown: None, + tasks: vec![send_task], + layer_guard: None, + }, + } + } + + pub(super) fn keep_layer_alive(&mut self, layer_guard: T) + where + T: Send + Sync + 'static, + { + self._cleanup.layer_guard = Some(Box::new(layer_guard)); + } + + pub(crate) fn into_tunnel_parts(self) -> UdpSessionTunnelParts { + let Self { + local_addr, + peer_addr, + kind, + codec, + incoming, + outgoing, + closed, + _cleanup, + } = self; + UdpSessionTunnelParts { + local_addr, + peer_addr, + kind, + codec, + session_recv_rx: incoming.into_inner(), + session_send_tx: outgoing.into_inner(), + closed, + cleanup: _cleanup, + } + } +} + +pub(crate) struct UdpSessionTunnelParts { + pub(crate) local_addr: SocketAddr, + pub(crate) peer_addr: SocketAddr, + pub(crate) kind: UdpSessionKind, + pub(crate) codec: UdpSessionCodec, + pub(crate) session_recv_rx: RingSocketReceiver, + pub(crate) session_send_tx: RingSocketSender, + pub(crate) closed: watch::Receiver, + pub(crate) cleanup: UdpSessionCleanup, +} + +#[async_trait] +impl UdpSessionSocket for UdpSession { + fn kind(&self) -> UdpSessionKind { + self.kind + } + + fn local_addr(&self) -> std::io::Result { + Ok(self.local_addr) + } + + fn peer_addr(&self) -> std::io::Result { + Ok(self.peer_addr) + } + + async fn send(&self, data: &[u8]) -> std::io::Result { + self.codec.validate_payload(data)?; + let mut closed = self.closed.clone(); + if *closed.borrow() { + return Err(udp_session_closed_error()); + } + let (completion, sent) = oneshot::channel(); + let outbound = UdpSessionOutbound { + payload: BytesMut::from(data), + completion, + }; + tokio::select! { + biased; + _ = closed.changed() => return Err(udp_session_closed_error()), + ret = async { + let mut outgoing = self.outgoing.lock().await; + outgoing.send(outbound).await + } => ret.map_err(ring_socket_error_to_io)?, + } + + tokio::select! { + biased; + ret = sent => ret.map_err(|_| udp_session_closed_error())?, + _ = closed.changed() => Err(udp_session_closed_error()), + } + } + + async fn recv(&self, buf: &mut [u8]) -> std::io::Result { + self.recv_with_meta(buf).await.map(|(len, _meta)| len) + } + + async fn recv_with_meta(&self, buf: &mut [u8]) -> std::io::Result<(usize, UdpSessionRecvMeta)> { + let mut closed = self.closed.clone(); + if *closed.borrow() { + return Err(udp_session_closed_error()); + } + let mut incoming = self.incoming.lock().await; + let payload = tokio::select! { + biased; + _ = closed.changed() => return Err(udp_session_closed_error()), + payload = incoming.next() => payload + .ok_or_else(udp_session_closed_error)? + .map_err(ring_socket_error_to_io)?, + }; + let len = payload.payload.len().min(buf.len()); + buf[..len].copy_from_slice(&payload.payload[..len]); + Ok(( + len, + UdpSessionRecvMeta { + dst_ip: payload.dst_ip, + }, + )) + } +} + +pub(super) struct UdpSessionRingParts { + pub(super) session_recv_rx: RingSocketReceiver, + pub(super) session_recv_tx: Arc>>, + pub(super) session_send_tx: RingSocketSender, + pub(super) session_send_rx: RingSocketReceiver, + pub(super) close_tx: watch::Sender, + pub(super) close_rx: watch::Receiver, +} + +pub(super) fn create_udp_session_rings() -> UdpSessionRingParts { + let (session_recv_rx_socket, session_recv_tx_socket) = + RingSocket::pair(UDP_SESSION_QUEUE_CAPACITY); + let (session_send_rx_socket, session_send_tx_socket) = + RingSocket::pair(UDP_SESSION_QUEUE_CAPACITY); + let (session_recv_rx, _unused_session_recv_tx) = session_recv_rx_socket.split(); + let (_unused_session_recv_peer_rx, session_recv_tx) = session_recv_tx_socket.split(); + let (session_send_rx, _unused_session_send_peer_tx) = session_send_rx_socket.split(); + let (_unused_session_send_rx, session_send_tx) = session_send_tx_socket.split(); + let (close_tx, close_rx) = watch::channel(false); + UdpSessionRingParts { + session_recv_rx, + session_recv_tx: Arc::new(StdMutex::new(session_recv_tx)), + session_send_tx, + session_send_rx, + close_tx, + close_rx, + } +} + +pub(super) fn udp_session_registry_entry(rings: &UdpSessionRingParts) -> UdpSessionRegistryEntry { + UdpSessionRegistryEntry { + incoming: rings.session_recv_tx.clone(), + close: rings.close_tx.clone(), + } +} + +pub(super) fn close_udp_session(sessions: &UdpSessionRegistry, key: UdpSessionKey) { + if let Some((_, entry)) = sessions.remove(&key) { + let _ = entry.close.send(true); + } +} + +pub(super) fn close_classified_udp_session( + classified_sessions: &ClassifiedUdpSessionRegistry, + key: ClassifiedUdpSessionKey, +) { + if let Some((_, entry)) = classified_sessions.remove(&key) { + let _ = entry.close.send(true); + } +} + +pub(super) fn close_all_udp_sessions(sessions: &UdpSessionRegistry) { + let close_senders = sessions + .iter() + .map(|entry| entry.value().close.clone()) + .collect::>(); + for close in close_senders { + let _ = close.send(true); + } + sessions.clear(); +} + +pub(super) fn close_all_classified_udp_sessions( + classified_sessions: &ClassifiedUdpSessionRegistry, +) { + let close_senders = classified_sessions + .iter() + .map(|entry| entry.value().close.clone()) + .collect::>(); + for close in close_senders { + let _ = close.send(true); + } + classified_sessions.clear(); +} + +fn ring_socket_error_to_io(error: crate::socket::ring::RingSocketError) -> io::Error { + let kind = match error { + crate::socket::ring::RingSocketError::Closed => io::ErrorKind::UnexpectedEof, + crate::socket::ring::RingSocketError::Full => io::ErrorKind::WouldBlock, + crate::socket::ring::RingSocketError::AlreadySplit => io::ErrorKind::Other, + }; + io::Error::new(kind, error.to_string()) +} + +fn udp_session_closed_error() -> io::Error { + io::Error::new(io::ErrorKind::UnexpectedEof, "udp session closed") +} + +#[derive(Debug, Clone, Copy)] +pub(super) enum UdpSessionEnqueuePolicy { + Lossy, + Reliable, +} + +async fn forward_udp_session_to_socket( + socket: Arc, + peer_addr: SocketAddr, + codec: UdpSessionCodec, + mut outgoing: RingSocketReceiver, + mut shutdown: watch::Receiver, + close: UdpSessionClose, +) where + S: VirtualUdpSocket, +{ + loop { + tokio::select! { + biased; + _ = shutdown.changed() => { + close.close(); + break; + } + outbound = outgoing.next() => { + let Some(outbound) = outbound else { + break; + }; + let outbound = match outbound { + Ok(outbound) => outbound, + Err(err) => { + tracing::debug!(?err, ?peer_addr, "udp session outgoing ring closed"); + close.close(); + break; + } + }; + let payload_len = outbound.payload.len(); + let datagram = match codec.encode(&outbound.payload) { + Ok(datagram) => datagram, + Err(err) => { + tracing::debug!(?err, ?peer_addr, ?codec, "udp session datagram encode error"); + let _ = outbound.completion.send(Err(err)); + close.close(); + break; + } + }; + match socket.send_to(&datagram, peer_addr).await { + Ok(_) => { + let _ = outbound.completion.send(Ok(payload_len)); + } + Err(err) => { + tracing::debug!(?err, ?peer_addr, "udp session send error"); + let _ = outbound.completion.send(Err(err)); + close.close(); + break; + } + } + } + } + } +} + +pub(super) fn dispatch_payload_to_session( + incoming: &Arc>>, + payload: impl Into, + policy: UdpSessionEnqueuePolicy, +) -> bool { + let payload = payload.into(); + let result = { + let mut incoming = incoming.lock().unwrap(); + match policy { + UdpSessionEnqueuePolicy::Lossy => incoming.try_send(payload), + UdpSessionEnqueuePolicy::Reliable => incoming.force_send(payload), + } + }; + match result { + Ok(()) => true, + Err(RingSocketSendError::Full(_)) => { + tracing::trace!(?policy, "udp session data queue full"); + true + } + Err(RingSocketSendError::Closed(_)) => false, + } +} + +#[cfg(test)] +mod tests { + use super::*; + + impl UdpSessionClose { + fn signal_only(close: watch::Sender) -> Self { + Self { + close, + target: UdpSessionCloseTarget::SignalOnly, + } + } + } + + impl UdpSession { + pub(crate) fn identity_standalone( + socket: Arc, + peer_addr: SocketAddr, + kind: UdpSessionKind, + ) -> io::Result + where + S: VirtualUdpSocket, + { + let local_addr = socket.local_addr()?; + let rings = create_udp_session_rings(); + let (shutdown_tx, _) = watch::channel(false); + let close = UdpSessionClose::signal_only(rings.close_tx.clone()); + let recv_socket = socket.clone(); + let recv_task = tokio::spawn(forward_identity_socket_to_udp_session( + recv_socket, + peer_addr, + rings.session_recv_tx.clone(), + shutdown_tx.subscribe(), + close.clone(), + )); + let mut session = Self::new( + socket.clone(), + local_addr, + peer_addr, + kind, + UdpSessionCodec::Identity, + rings, + close, + shutdown_tx.subscribe(), + ); + session._cleanup.shutdown = Some(shutdown_tx); + session._cleanup.tasks.push(recv_task); + Ok(session) + } + } + + async fn forward_identity_socket_to_udp_session( + socket: Arc, + peer_addr: SocketAddr, + incoming: Arc>>, + mut shutdown: watch::Receiver, + close: UdpSessionClose, + ) where + S: VirtualUdpSocket, + { + let mut buf = [0u8; 65535]; + loop { + tokio::select! { + biased; + _ = shutdown.changed() => { + close.close(); + break; + } + ret = socket.recv_from(&mut buf) => { + let (len, remote_addr) = match ret { + Ok(ret) => ret, + Err(err) => { + tracing::debug!(?err, ?peer_addr, "identity udp session recv error"); + close.close(); + break; + } + }; + if remote_addr != peer_addr { + continue; + } + if !dispatch_payload_to_session( + &incoming, + BytesMut::from(&buf[..len]), + UdpSessionEnqueuePolicy::Reliable, + ) { + close.close(); + break; + } + } + } + } + } +} diff --git a/easytier-core/src/socket/udp/tests.rs b/easytier-core/src/socket/udp/tests.rs new file mode 100644 index 00000000..694b30a2 --- /dev/null +++ b/easytier-core/src/socket/udp/tests.rs @@ -0,0 +1,2091 @@ +use std::{ + collections::VecDeque, + io, + net::{Ipv4Addr, Ipv6Addr, SocketAddr, SocketAddrV4, SocketAddrV6}, + sync::{ + Arc, Mutex, Mutex as StdMutex, + atomic::{AtomicBool, AtomicU16, Ordering}, + }, + time::Duration, +}; + +use async_trait::async_trait; +use bytecodec::EncodeExt as _; +use bytes::BytesMut; +use dashmap::DashMap; +use futures::StreamExt; +use stun_codec::{Message, MessageClass, MessageEncoder, rfc5389::methods::BINDING}; +use tokio::sync::{Semaphore, mpsc, watch}; + +use crate::{ + packet::stun::{Attribute, ChangeRequest, u32_to_tid}, + packet::{UDP_TUNNEL_HEADER_SIZE, UdpPacketType, hole_punch_packet_tid}, + socket::{IpVersion, NetNamespace, SocketContext, SocketListener}, +}; + +use super::{layer::*, packet::*, session::*, virtual_socket::*, *}; + +#[test] +fn bind_options_constructors_describe_socket_purpose() { + let listener_addr = SocketAddr::from(([0, 0, 0, 0], 12345)); + + assert_eq!( + UdpBindOptions::hole_punch_control(), + UdpBindOptions { + context: SocketContext::default(), + local_addr: None, + bind_device: None, + reuse_addr: false, + reuse_port: false, + only_v6: false, + purpose: UdpSocketPurpose::HolePunchControl, + } + ); + assert_eq!( + UdpBindOptions::hole_punch_candidate(), + UdpBindOptions { + context: SocketContext::default(), + local_addr: None, + bind_device: None, + reuse_addr: false, + reuse_port: false, + only_v6: false, + purpose: UdpSocketPurpose::HolePunchCandidate, + } + ); + assert_eq!( + UdpBindOptions::direct_connect(), + UdpBindOptions { + context: SocketContext::default(), + local_addr: None, + bind_device: None, + reuse_addr: false, + reuse_port: false, + only_v6: false, + purpose: UdpSocketPurpose::DirectConnect, + } + ); + assert_eq!( + UdpBindOptions::port_bound_listener(listener_addr), + UdpBindOptions { + context: SocketContext::default(), + local_addr: Some(listener_addr), + bind_device: None, + reuse_addr: false, + reuse_port: false, + only_v6: false, + purpose: UdpSocketPurpose::PortBoundListener, + } + ); + assert_eq!( + UdpBindOptions::socks5(), + UdpBindOptions { + context: SocketContext::default(), + local_addr: None, + bind_device: None, + reuse_addr: false, + reuse_port: false, + only_v6: false, + purpose: UdpSocketPurpose::Socks5, + } + ); + assert_eq!( + UdpBindOptions::port_forward(listener_addr).purpose, + UdpSocketPurpose::PortForward + ); + assert_eq!( + UdpBindOptions::port_lease(listener_addr).purpose, + UdpSocketPurpose::PortLease + ); + assert_eq!( + UdpBindOptions::default(), + UdpBindOptions::hole_punch_control() + ); +} + +#[test] +fn session_connect_request_keeps_peer_scoped_udp_shape() { + let remote_addr = SocketAddr::from(([192, 0, 2, 1], 11010)); + let bind_addr = SocketAddr::from(([0, 0, 0, 0], 22020)); + + let request = UdpSessionConnectRequest::wireguard(remote_addr) + .with_bind(UdpBindOptions::port_bound_listener(bind_addr)); + + assert_eq!(request.remote_addr, remote_addr); + assert_eq!(request.protocol, UdpSessionProtocol::WireGuard); + assert_eq!( + request.bind, + UdpBindOptions { + context: SocketContext::default(), + local_addr: Some(bind_addr), + bind_device: None, + reuse_addr: false, + reuse_port: false, + only_v6: false, + purpose: UdpSocketPurpose::PortBoundListener, + } + ); +} + +#[test] +fn session_listen_request_keeps_bind_options() { + let bind_addr = SocketAddr::from(([0, 0, 0, 0], 11010)); + let bind = UdpBindOptions::port_bound_listener(bind_addr); + + assert_eq!(UdpSessionListenRequest::new(bind.clone()).bind, bind); +} + +async fn drain_session_payloads(mut rings: UdpSessionRingParts) -> usize { + let mut count = 0; + while let Ok(Some(Ok(_))) = + tokio::time::timeout(Duration::from_millis(10), rings.session_recv_rx.next()).await + { + count += 1; + } + count +} + +#[tokio::test] +async fn lossy_udp_session_enqueue_preserves_reserved_capacity() { + const RING_RESERVED_CAPACITY: usize = 4; + + let rings = create_udp_session_rings(); + for _ in 0..UDP_SESSION_QUEUE_CAPACITY { + assert!(dispatch_payload_to_session( + &rings.session_recv_tx, + BytesMut::from("lossy"), + UdpSessionEnqueuePolicy::Lossy, + )); + } + assert_eq!( + drain_session_payloads(rings).await, + UDP_SESSION_QUEUE_CAPACITY - RING_RESERVED_CAPACITY + ); + + let rings = create_udp_session_rings(); + for _ in 0..UDP_SESSION_QUEUE_CAPACITY { + assert!(dispatch_payload_to_session( + &rings.session_recv_tx, + BytesMut::from("lossy"), + UdpSessionEnqueuePolicy::Lossy, + )); + } + for _ in 0..RING_RESERVED_CAPACITY { + assert!(dispatch_payload_to_session( + &rings.session_recv_tx, + BytesMut::from("reliable"), + UdpSessionEnqueuePolicy::Reliable, + )); + } + assert_eq!( + drain_session_payloads(rings).await, + UDP_SESSION_QUEUE_CAPACITY + ); +} + +struct MockUdpSessionSocket { + kind: UdpSessionKind, + local_addr: SocketAddr, + peer_addr: SocketAddr, + incoming: Mutex>, + sent: Mutex>, +} + +#[async_trait] +impl UdpSessionSocket for MockUdpSessionSocket { + fn kind(&self) -> UdpSessionKind { + self.kind + } + + fn local_addr(&self) -> std::io::Result { + Ok(self.local_addr) + } + + fn peer_addr(&self) -> std::io::Result { + Ok(self.peer_addr) + } + + async fn send(&self, data: &[u8]) -> std::io::Result { + self.sent.lock().unwrap().extend_from_slice(data); + Ok(data.len()) + } + + async fn recv(&self, buf: &mut [u8]) -> std::io::Result { + let incoming = self.incoming.lock().unwrap(); + let len = incoming.len().min(buf.len()); + buf[..len].copy_from_slice(&incoming[..len]); + Ok(len) + } +} + +#[tokio::test] +async fn udp_session_socket_is_peer_scoped() { + let local_addr = SocketAddr::from(([127, 0, 0, 1], 10000)); + let peer_addr = SocketAddr::from(([127, 0, 0, 1], 10001)); + let socket = MockUdpSessionSocket { + kind: UdpSessionKind::WireGuard, + local_addr, + peer_addr, + incoming: Mutex::new(b"pong".to_vec()), + sent: Mutex::new(Vec::new()), + }; + + assert_eq!(socket.kind(), UdpSessionKind::WireGuard); + assert_eq!(socket.local_addr().unwrap(), local_addr); + assert_eq!(socket.peer_addr().unwrap(), peer_addr); + assert_eq!(socket.send(b"ping").await.unwrap(), 4); + + let mut buf = [0; 8]; + let len = socket.recv(&mut buf).await.unwrap(); + + assert_eq!(&buf[..len], b"pong"); + assert_eq!(&*socket.sent.lock().unwrap(), b"ping"); +} + +struct MockUdpSessionListener { + local_addr: SocketAddr, + accepted: Option, +} + +#[async_trait] +impl UdpSessionListener for MockUdpSessionListener { + type Session = MockUdpSessionSocket; + + async fn listen(&mut self, request: UdpSessionListenRequest) -> anyhow::Result<()> { + if let Some(local_addr) = request.bind.local_addr { + self.local_addr = local_addr; + } + Ok(()) + } + + fn local_addr(&self) -> std::io::Result { + Ok(self.local_addr) + } + + async fn accept(&mut self) -> anyhow::Result { + self.accepted + .take() + .ok_or_else(|| anyhow::anyhow!("no accepted session")) + } +} + +#[tokio::test] +async fn udp_session_listener_reports_bound_local_addr_before_accept() { + let local_addr = SocketAddr::from(([127, 0, 0, 1], 10000)); + let peer_addr = SocketAddr::from(([127, 0, 0, 1], 10001)); + let mut listener = MockUdpSessionListener { + local_addr: SocketAddr::from(([0, 0, 0, 0], 0)), + accepted: Some(MockUdpSessionSocket { + kind: UdpSessionKind::EasyTierMux, + local_addr, + peer_addr, + incoming: Mutex::new(Vec::new()), + sent: Mutex::new(Vec::new()), + }), + }; + + listener + .listen(UdpSessionListenRequest::new( + UdpBindOptions::port_bound_listener(local_addr), + )) + .await + .unwrap(); + + assert_eq!(listener.local_addr().unwrap(), local_addr); + assert_eq!( + listener.accept().await.unwrap().peer_addr().unwrap(), + peer_addr + ); +} + +struct MockVirtualUdpSocket { + local_addr: SocketAddr, + incoming: Mutex, SocketAddr)>>, + sent: Mutex, SocketAddr)>>, + send_attempts: Mutex, SocketAddr, UdpSocketSendMeta)>>, + reject_preferred_source: AtomicBool, +} + +impl MockVirtualUdpSocket { + fn new(local_addr: SocketAddr, incoming: Vec<(Vec, SocketAddr)>) -> Self { + Self { + local_addr, + incoming: Mutex::new(incoming.into()), + sent: Mutex::new(Vec::new()), + send_attempts: Mutex::new(Vec::new()), + reject_preferred_source: AtomicBool::new(false), + } + } + + fn sent(&self) -> Vec<(Vec, SocketAddr)> { + self.sent.lock().unwrap().clone() + } + + fn send_attempts(&self) -> Vec<(Vec, SocketAddr, UdpSocketSendMeta)> { + self.send_attempts.lock().unwrap().clone() + } +} + +#[async_trait] +impl VirtualUdpSocket for MockVirtualUdpSocket { + fn local_addr(&self) -> io::Result { + Ok(self.local_addr) + } + + async fn send_to(&self, data: &[u8], addr: SocketAddr) -> io::Result { + self.sent.lock().unwrap().push((data.to_vec(), addr)); + Ok(data.len()) + } + + async fn send_to_with_meta( + &self, + data: &[u8], + addr: SocketAddr, + meta: UdpSocketSendMeta, + ) -> io::Result { + self.send_attempts + .lock() + .unwrap() + .push((data.to_vec(), addr, meta)); + if meta.src_ip.is_some() && self.reject_preferred_source.load(Ordering::Relaxed) { + return Err(io::Error::new( + io::ErrorKind::AddrNotAvailable, + "injected preferred source failure", + )); + } + self.send_to(data, addr).await + } + + async fn recv_from(&self, buf: &mut [u8]) -> io::Result<(usize, SocketAddr)> { + let (data, remote_addr) = + self.incoming.lock().unwrap().pop_front().ok_or_else(|| { + io::Error::new(io::ErrorKind::UnexpectedEof, "no incoming datagram") + })?; + let len = data.len().min(buf.len()); + buf[..len].copy_from_slice(&data[..len]); + Ok((len, remote_addr)) + } +} + +fn easytier_stun_request(change_ip: bool, change_port: bool) -> Vec { + let mut request = Message::::new(MessageClass::Request, BINDING, u32_to_tid(7)); + if change_ip || change_port { + request.add_attribute(Attribute::ChangeRequest(ChangeRequest::new( + change_ip, + change_port, + ))); + } + MessageEncoder::new().encode_into_bytes(request).unwrap() +} + +struct FailingSendVirtualUdpSocket { + local_addr: SocketAddr, +} + +#[async_trait] +impl VirtualUdpSocket for FailingSendVirtualUdpSocket { + fn local_addr(&self) -> io::Result { + Ok(self.local_addr) + } + + async fn send_to(&self, _data: &[u8], _addr: SocketAddr) -> io::Result { + Err(io::Error::new( + io::ErrorKind::ConnectionRefused, + "injected send failure", + )) + } + + async fn recv_from(&self, _buf: &mut [u8]) -> io::Result<(usize, SocketAddr)> { + std::future::pending().await + } +} + +#[derive(Debug, Default)] +struct BlockingUdpSessionStunResponder { + started: tokio::sync::Notify, + release: tokio::sync::Notify, +} + +#[async_trait] +impl UdpSessionStunResponder for BlockingUdpSessionStunResponder { + async fn respond_stun( + &self, + _socket: Arc, + _datagram: &[u8], + _remote_addr: SocketAddr, + ) -> io::Result<()> { + self.started.notify_waiters(); + self.release.notified().await; + Ok(()) + } +} + +struct AutoSackVirtualUdpSocket { + local_addr: SocketAddr, + incoming: Mutex, SocketAddr)>>, + sent: Mutex, SocketAddr)>>, + incoming_notify: tokio::sync::Notify, +} + +impl AutoSackVirtualUdpSocket { + fn new(local_addr: SocketAddr) -> Self { + Self { + local_addr, + incoming: Mutex::new(VecDeque::new()), + sent: Mutex::new(Vec::new()), + incoming_notify: tokio::sync::Notify::new(), + } + } + + fn sent(&self) -> Vec<(Vec, SocketAddr)> { + self.sent.lock().unwrap().clone() + } +} + +#[async_trait] +impl VirtualUdpSocket for AutoSackVirtualUdpSocket { + fn local_addr(&self) -> io::Result { + Ok(self.local_addr) + } + + async fn send_to(&self, data: &[u8], addr: SocketAddr) -> io::Result { + self.sent.lock().unwrap().push((data.to_vec(), addr)); + if let Ok(packet) = parse_udp_session_datagram(BytesMut::from(data), false) { + let header = packet.udp_tunnel_header().unwrap(); + if header.msg_type == UdpPacketType::Syn as u8 && packet.udp_payload().len() == 8 { + let conn_id = header.conn_id.get(); + let magic = u64::from_le_bytes(packet.udp_payload()[..8].try_into().unwrap()); + self.incoming + .lock() + .unwrap() + .push_back((new_sack_packet(conn_id, magic).into_bytes().to_vec(), addr)); + self.incoming_notify.notify_one(); + } + } + Ok(data.len()) + } + + async fn recv_from(&self, buf: &mut [u8]) -> io::Result<(usize, SocketAddr)> { + loop { + if let Some((data, remote_addr)) = self.incoming.lock().unwrap().pop_front() { + let len = data.len().min(buf.len()); + buf[..len].copy_from_slice(&data[..len]); + return Ok((len, remote_addr)); + } + self.incoming_notify.notified().await; + } + } +} + +#[async_trait] +impl UdpSessionStunResponder for BlockingUdpSessionStunResponder { + async fn respond_stun( + &self, + _socket: Arc, + _datagram: &[u8], + _remote_addr: SocketAddr, + ) -> io::Result<()> { + self.started.notify_waiters(); + self.release.notified().await; + Ok(()) + } +} + +fn create_test_easy_tier_mux_session( + socket: Arc, + key: UdpSessionKey, + sessions: Arc, +) -> (UdpSession, watch::Sender) +where + S: VirtualUdpSocket, +{ + let (shutdown_tx, shutdown_rx) = watch::channel(false); + ( + create_test_easy_tier_mux_session_with_shutdown(socket, key, sessions, shutdown_rx), + shutdown_tx, + ) +} + +fn create_test_easy_tier_mux_session_with_shutdown( + socket: Arc, + key: UdpSessionKey, + sessions: Arc, + shutdown: watch::Receiver, +) -> UdpSession +where + S: VirtualUdpSocket, +{ + let local_addr = socket.local_addr().unwrap(); + let rings = create_udp_session_rings(); + sessions.insert(key, udp_session_registry_entry(&rings)); + let close = UdpSessionClose::easy_tier(key, rings.close_tx.clone(), sessions); + UdpSession::new( + socket, + local_addr, + key.peer_addr, + UdpSessionKind::EasyTierMux, + UdpSessionCodec::EasyTierData { + conn_id: key.conn_id, + }, + rings, + close, + shutdown, + ) +} + +async fn wait_for_sent(mut sent: F, min_len: usize) -> Vec<(Vec, SocketAddr)> +where + F: FnMut() -> Vec<(Vec, SocketAddr)>, +{ + tokio::time::timeout(Duration::from_secs(1), async move { + loop { + let packets = sent(); + if packets.len() >= min_len { + return packets; + } + tokio::task::yield_now().await; + } + }) + .await + .unwrap() +} + +fn wireguard_transport_packet(payload: &[u8]) -> Vec { + let mut packet = vec![0; 32.max(4 + payload.len())]; + packet[..4].copy_from_slice(&4u32.to_le_bytes()); + packet[4..4 + payload.len()].copy_from_slice(payload); + packet +} + +fn wireguard_packet_with_easy_tier_data_header(payload: &[u8]) -> Vec { + let payload_len = 24.max(payload.len()); + let mut packet = vec![0; UDP_TUNNEL_HEADER_SIZE + payload_len]; + packet[..4].copy_from_slice(&4u32.to_le_bytes()); + packet[4] = UdpPacketType::Data as u8; + packet[6..8].copy_from_slice(&(payload_len as u16).to_le_bytes()); + packet[UDP_TUNNEL_HEADER_SIZE..UDP_TUNNEL_HEADER_SIZE + payload.len()].copy_from_slice(payload); + packet +} + +#[tokio::test] +async fn wireguard_udp_session_sends_to_peer_addr() { + let local_addr = SocketAddr::from(([127, 0, 0, 1], 12000)); + let peer_addr = SocketAddr::from(([127, 0, 0, 1], 12001)); + let socket = Arc::new(AutoSackVirtualUdpSocket::new(local_addr)); + let session = + UdpSession::identity_standalone(socket.clone(), peer_addr, UdpSessionKind::WireGuard) + .unwrap(); + + assert_eq!(session.kind(), UdpSessionKind::WireGuard); + assert_eq!(session.local_addr().unwrap(), local_addr); + assert_eq!(session.peer_addr().unwrap(), peer_addr); + assert_eq!(session.send(b"hello").await.unwrap(), 5); + + let sent = wait_for_sent(|| socket.sent(), 1).await; + assert_eq!(sent, vec![(b"hello".to_vec(), peer_addr)]); +} + +#[tokio::test] +async fn wireguard_udp_session_receives_only_from_peer_addr() { + let local_addr = SocketAddr::from(([127, 0, 0, 1], 12000)); + let peer_addr = SocketAddr::from(([127, 0, 0, 1], 12001)); + let unexpected_addr = SocketAddr::from(([127, 0, 0, 1], 12002)); + let socket = Arc::new(AutoSackVirtualUdpSocket::new(local_addr)); + socket.incoming.lock().unwrap().extend([ + (b"noise".to_vec(), unexpected_addr), + (b"payload".to_vec(), peer_addr), + ]); + socket.incoming_notify.notify_one(); + let session = + UdpSession::identity_standalone(socket, peer_addr, UdpSessionKind::WireGuard).unwrap(); + + let mut buf = [0; 16]; + let len = session.recv(&mut buf).await.unwrap(); + + assert_eq!(&buf[..len], b"payload"); +} + +#[tokio::test] +async fn wireguard_udp_session_send_failure_closes_recv() { + let local_addr = SocketAddr::from(([127, 0, 0, 1], 12000)); + let peer_addr = SocketAddr::from(([127, 0, 0, 1], 12001)); + let socket = Arc::new(FailingSendVirtualUdpSocket { local_addr }); + let session = + UdpSession::identity_standalone(socket, peer_addr, UdpSessionKind::WireGuard).unwrap(); + + let err = tokio::time::timeout(Duration::from_secs(1), session.send(b"payload")) + .await + .unwrap() + .unwrap_err(); + assert_eq!(err.kind(), io::ErrorKind::ConnectionRefused); + + let mut buf = [0; 16]; + let err = tokio::time::timeout(Duration::from_secs(1), session.recv(&mut buf)) + .await + .unwrap() + .unwrap_err(); + assert_eq!(err.kind(), io::ErrorKind::UnexpectedEof); +} + +#[tokio::test] +async fn udp_layer_routes_wireguard_packets_to_registered_wireguard_session() { + let local_addr = SocketAddr::from(([127, 0, 0, 1], 12000)); + let peer_addr = SocketAddr::from(([127, 0, 0, 1], 12001)); + let socket = Arc::new(AutoSackVirtualUdpSocket::new(local_addr)); + let layer = UdpSessionLayer::new(socket.clone()); + let session = layer + .open_classified_session(UdpSessionProtocol::WireGuard, peer_addr) + .unwrap(); + + assert_eq!(session.kind(), UdpSessionKind::WireGuard); + assert_eq!(session.send(b"outbound").await.unwrap(), 8); + let sent = wait_for_sent(|| socket.sent(), 1).await; + assert_eq!(sent, vec![(b"outbound".to_vec(), peer_addr)]); + + let inbound = wireguard_transport_packet(b"inbound"); + socket + .incoming + .lock() + .unwrap() + .push_back((inbound.clone(), peer_addr)); + socket.incoming_notify.notify_one(); + + let mut buf = [0; 64]; + let len = tokio::time::timeout(Duration::from_secs(1), session.recv(&mut buf)) + .await + .unwrap() + .unwrap(); + + assert_eq!(&buf[..len], inbound.as_slice()); +} + +#[tokio::test] +async fn udp_layer_routes_unclaimed_easy_tier_shaped_wireguard_packet_to_wireguard_session() { + let local_addr = SocketAddr::from(([127, 0, 0, 1], 12000)); + let peer_addr = SocketAddr::from(([127, 0, 0, 1], 12001)); + let socket = Arc::new(AutoSackVirtualUdpSocket::new(local_addr)); + let layer = UdpSessionLayer::new(socket.clone()); + let session = layer + .open_classified_session(UdpSessionProtocol::WireGuard, peer_addr) + .unwrap(); + let packet = wireguard_packet_with_easy_tier_data_header(b"wireguard-collision"); + + socket + .incoming + .lock() + .unwrap() + .push_back((packet.clone(), peer_addr)); + socket.incoming_notify.notify_one(); + + let mut buf = [0; 64]; + let len = tokio::time::timeout(Duration::from_secs(1), session.recv(&mut buf)) + .await + .unwrap() + .unwrap(); + + assert_eq!(&buf[..len], packet.as_slice()); + assert_eq!(layer.active_session_count(), 0); +} + +#[tokio::test] +async fn udp_layer_routes_claimed_easy_tier_data_to_mux_before_wireguard_session() { + let local_addr = SocketAddr::from(([127, 0, 0, 1], 12000)); + let peer_addr = SocketAddr::from(([127, 0, 0, 1], 12001)); + let socket = Arc::new(AutoSackVirtualUdpSocket::new(local_addr)); + let layer = UdpSessionLayer::new(socket.clone()); + let wireguard_session = layer + .open_classified_session(UdpSessionProtocol::WireGuard, peer_addr) + .unwrap(); + let key = UdpSessionKey::new(peer_addr, 4); + let (mux_session, _shutdown_tx) = + create_test_easy_tier_mux_session(socket.clone(), key, layer.sessions.clone()); + let packet = wireguard_packet_with_easy_tier_data_header(b"mux-payload"); + let mux_payload = packet[UDP_TUNNEL_HEADER_SIZE..].to_vec(); + + socket + .incoming + .lock() + .unwrap() + .push_back((packet, peer_addr)); + socket.incoming_notify.notify_one(); + + let mut mux_buf = [0; 32]; + let len = tokio::time::timeout(Duration::from_secs(1), mux_session.recv(&mut mux_buf)) + .await + .unwrap() + .unwrap(); + assert_eq!(&mux_buf[..len], mux_payload.as_slice()); + + let mut wireguard_buf = [0; 64]; + assert!( + tokio::time::timeout( + Duration::from_millis(100), + wireguard_session.recv(&mut wireguard_buf) + ) + .await + .is_err() + ); +} + +#[tokio::test] +async fn udp_layer_drops_unknown_datagram_instead_of_creating_session() { + let local_addr = SocketAddr::from(([127, 0, 0, 1], 12000)); + let peer_addr = SocketAddr::from(([127, 0, 0, 1], 12001)); + let socket = Arc::new(AutoSackVirtualUdpSocket::new(local_addr)); + let layer = Arc::new(UdpSessionLayer::new(socket.clone())); + let mut accept_task = tokio::spawn({ + let layer = layer.clone(); + async move { + layer + .accept_classified_session(UdpSessionProtocol::WireGuard) + .await + } + }); + + tokio::time::timeout(Duration::from_secs(1), async { + while !layer + .classified_accepts + .get(&UdpSessionProtocol::WireGuard) + .unwrap() + .accept_enabled + .load(Ordering::Relaxed) + { + tokio::task::yield_now().await; + } + }) + .await + .unwrap(); + + socket + .incoming + .lock() + .unwrap() + .push_back((b"first".to_vec(), peer_addr)); + socket.incoming_notify.notify_one(); + + assert!( + tokio::time::timeout(Duration::from_millis(100), &mut accept_task) + .await + .is_err() + ); + accept_task.abort(); + assert_eq!(layer.active_classified_session_count(), 0); +} + +#[tokio::test] +async fn udp_layer_drops_malformed_quic_like_datagrams_instead_of_creating_session() { + let local_addr = SocketAddr::from(([127, 0, 0, 1], 12000)); + let peer_addr = SocketAddr::from(([127, 0, 0, 1], 12001)); + let socket = Arc::new(AutoSackVirtualUdpSocket::new(local_addr)); + let layer = Arc::new(UdpSessionLayer::new(socket.clone())); + let mut accept_task = tokio::spawn({ + let layer = layer.clone(); + async move { + layer + .accept_classified_session(UdpSessionProtocol::Quic) + .await + } + }); + + tokio::time::timeout(Duration::from_secs(1), async { + while !layer + .classified_accepts + .get(&UdpSessionProtocol::Quic) + .unwrap() + .accept_enabled + .load(Ordering::Relaxed) + { + tokio::task::yield_now().await; + } + }) + .await + .unwrap(); + + socket + .incoming + .lock() + .unwrap() + .push_back((vec![0xc0; 32], peer_addr)); + socket + .incoming + .lock() + .unwrap() + .push_back((vec![0xc0; 1200], peer_addr)); + socket.incoming_notify.notify_one(); + + assert!( + tokio::time::timeout(Duration::from_millis(100), &mut accept_task) + .await + .is_err() + ); + accept_task.abort(); + assert_eq!(layer.active_classified_session_count(), 0); +} + +#[tokio::test] +async fn udp_layer_accepts_unclaimed_easy_tier_shaped_wireguard_packet_when_enabled() { + let local_addr = SocketAddr::from(([127, 0, 0, 1], 12000)); + let peer_addr = SocketAddr::from(([127, 0, 0, 1], 12001)); + let socket = Arc::new(AutoSackVirtualUdpSocket::new(local_addr)); + let layer = Arc::new(UdpSessionLayer::new(socket.clone())); + let accept_task = tokio::spawn({ + let layer = layer.clone(); + async move { + layer + .accept_classified_session(UdpSessionProtocol::WireGuard) + .await + } + }); + let packet = wireguard_packet_with_easy_tier_data_header(b"accepted-wireguard"); + + tokio::time::timeout(Duration::from_secs(1), async { + while !layer + .classified_accepts + .get(&UdpSessionProtocol::WireGuard) + .unwrap() + .accept_enabled + .load(Ordering::Relaxed) + { + tokio::task::yield_now().await; + } + }) + .await + .unwrap(); + + socket + .incoming + .lock() + .unwrap() + .push_back((packet.clone(), peer_addr)); + socket.incoming_notify.notify_one(); + + let accepted = tokio::time::timeout(Duration::from_secs(1), accept_task) + .await + .unwrap() + .unwrap() + .unwrap(); + let mut buf = [0; 192]; + let len = accepted.recv(&mut buf).await.unwrap(); + + assert_eq!(accepted.peer_addr().unwrap(), peer_addr); + assert_eq!(&buf[..len], packet.as_slice()); + assert_eq!(layer.active_session_count(), 0); + assert_eq!(layer.active_classified_session_count(), 1); +} + +#[tokio::test] +async fn udp_layer_pre_enabled_classified_accept_queues_first_packet() { + let local_addr = SocketAddr::from(([127, 0, 0, 1], 12000)); + let peer_addr = SocketAddr::from(([127, 0, 0, 1], 12001)); + let socket = Arc::new(AutoSackVirtualUdpSocket::new(local_addr)); + let layer = UdpSessionLayer::new(socket.clone()); + let packet = wireguard_packet_with_easy_tier_data_header(b"pre-enabled-wireguard"); + + layer + .enable_classified_accept(UdpSessionProtocol::WireGuard) + .unwrap(); + socket + .incoming + .lock() + .unwrap() + .push_back((packet.clone(), peer_addr)); + socket.incoming_notify.notify_one(); + + let accepted = tokio::time::timeout( + Duration::from_secs(1), + layer.accept_classified_session(UdpSessionProtocol::WireGuard), + ) + .await + .unwrap() + .unwrap(); + let mut buf = [0; 192]; + let len = accepted.recv(&mut buf).await.unwrap(); + + assert_eq!(accepted.peer_addr().unwrap(), peer_addr); + assert_eq!(&buf[..len], packet.as_slice()); + assert_eq!(layer.active_session_count(), 0); + assert_eq!(layer.active_classified_session_count(), 1); +} + +#[tokio::test] +async fn udp_layer_routes_quic_like_easytier_packet_to_existing_quic_session() { + let local_addr = SocketAddr::from(([127, 0, 0, 1], 12000)); + let peer_addr = SocketAddr::from(([127, 0, 0, 1], 12001)); + let socket = Arc::new(AutoSackVirtualUdpSocket::new(local_addr)); + let layer = UdpSessionLayer::new(socket.clone()); + let session = layer + .open_classified_session(UdpSessionProtocol::Quic, peer_addr) + .unwrap(); + let packet = new_udp_packet( + |header| { + header.conn_id.set(0x40); + header.msg_type = UdpPacketType::Syn as u8; + header.len.set(8); + }, + b"12345678", + ) + .into_bytes() + .to_vec(); + + socket + .incoming + .lock() + .unwrap() + .push_back((packet.clone(), peer_addr)); + socket.incoming_notify.notify_one(); + + let mut buf = [0; 64]; + let len = tokio::time::timeout(Duration::from_secs(1), session.recv(&mut buf)) + .await + .unwrap() + .unwrap(); + + assert_eq!(&buf[..len], packet.as_slice()); + assert_eq!(layer.active_session_count(), 0); + assert_eq!(layer.active_classified_session_count(), 1); +} + +#[tokio::test] +async fn udp_layer_keeps_easy_tier_syn_out_of_wireguard_session() { + let local_addr = SocketAddr::from(([127, 0, 0, 1], 12000)); + let peer_addr = SocketAddr::from(([127, 0, 0, 1], 12001)); + let socket = Arc::new(AutoSackVirtualUdpSocket::new(local_addr)); + let layer = UdpSessionLayer::new(socket.clone()); + let session = layer + .open_classified_session(UdpSessionProtocol::WireGuard, peer_addr) + .unwrap(); + let syn = new_syn_packet(0x1122_3344, 0x5566_7788).into_bytes(); + + socket + .incoming + .lock() + .unwrap() + .push_back((syn.to_vec(), peer_addr)); + socket.incoming_notify.notify_one(); + + tokio::time::timeout(Duration::from_secs(1), async { + while layer.active_session_count() == 0 { + tokio::task::yield_now().await; + } + }) + .await + .unwrap(); + + let mut buf = [0; 16]; + assert!( + tokio::time::timeout(Duration::from_millis(50), session.recv(&mut buf)) + .await + .is_err() + ); +} + +#[tokio::test] +async fn easy_tier_mux_udp_session_wraps_sent_payloads() { + let local_addr = SocketAddr::from(([127, 0, 0, 1], 12000)); + let peer_addr = SocketAddr::from(([127, 0, 0, 1], 12001)); + let conn_id = 0x1122_3344; + let socket = Arc::new(MockVirtualUdpSocket::new(local_addr, Vec::new())); + let sessions = Arc::new(DashMap::new()); + let (session, _shutdown_tx) = create_test_easy_tier_mux_session( + socket.clone(), + UdpSessionKey::new(peer_addr, conn_id), + sessions, + ); + + assert_eq!(session.kind(), UdpSessionKind::EasyTierMux); + assert_eq!(session.send(b"payload").await.unwrap(), 7); + + let sent = wait_for_sent(|| socket.sent(), 1).await; + assert_eq!(sent.len(), 1); + assert_eq!(sent[0].1, peer_addr); + + let packet = parse_udp_session_datagram(BytesMut::from(sent[0].0.as_slice()), false) + .expect("sent datagram should keep EasyTier UDP packet shape"); + let header = packet.udp_tunnel_header().unwrap(); + assert_eq!(header.conn_id.get(), conn_id); + assert_eq!(header.msg_type, UdpPacketType::Data as u8); + assert_eq!(header.len.get(), 7); + assert_eq!(packet.udp_payload(), b"payload"); +} + +#[tokio::test] +async fn easy_tier_mux_udp_session_rejects_oversized_payload_before_enqueue() { + let local_addr = SocketAddr::from(([127, 0, 0, 1], 12000)); + let peer_addr = SocketAddr::from(([127, 0, 0, 1], 12001)); + let conn_id = 0x1122_3344; + let socket = Arc::new(MockVirtualUdpSocket::new(local_addr, Vec::new())); + let sessions = Arc::new(DashMap::new()); + let (session, _shutdown_tx) = create_test_easy_tier_mux_session( + socket.clone(), + UdpSessionKey::new(peer_addr, conn_id), + sessions, + ); + + let payload = vec![0; u16::MAX as usize + 1]; + let err = session.send(&payload).await.unwrap_err(); + + assert_eq!(err.kind(), io::ErrorKind::InvalidInput); + assert!(socket.sent().is_empty()); +} + +#[tokio::test] +async fn easy_tier_mux_udp_session_send_failure_closes_session() { + let local_addr = SocketAddr::from(([127, 0, 0, 1], 12000)); + let peer_addr = SocketAddr::from(([127, 0, 0, 1], 12001)); + let conn_id = 0x1122_3344; + let key = UdpSessionKey::new(peer_addr, conn_id); + let socket = Arc::new(FailingSendVirtualUdpSocket { local_addr }); + let sessions = Arc::new(DashMap::new()); + let (session, _shutdown_tx) = create_test_easy_tier_mux_session(socket, key, sessions.clone()); + + let err = tokio::time::timeout(Duration::from_secs(1), session.send(b"payload")) + .await + .unwrap() + .unwrap_err(); + + assert_eq!(err.kind(), io::ErrorKind::ConnectionRefused); + tokio::time::timeout(Duration::from_secs(1), async { + while sessions.contains_key(&key) { + tokio::task::yield_now().await; + } + }) + .await + .unwrap(); + + let mut buf = [0; 16]; + let err = tokio::time::timeout(Duration::from_secs(1), session.recv(&mut buf)) + .await + .unwrap() + .unwrap_err(); + assert_eq!(err.kind(), io::ErrorKind::UnexpectedEof); +} + +#[tokio::test] +async fn easy_tier_mux_udp_session_receives_only_peer_data_payloads() { + let local_addr = SocketAddr::from(([127, 0, 0, 1], 12000)); + let peer_addr = SocketAddr::from(([127, 0, 0, 1], 12001)); + let unexpected_addr = SocketAddr::from(([127, 0, 0, 1], 12002)); + let conn_id = 0x1122_3344; + let socket = Arc::new(MockVirtualUdpSocket::new(local_addr, Vec::new())); + let sessions = Arc::new(DashMap::new()); + let (session, _shutdown_tx) = create_test_easy_tier_mux_session( + socket, + UdpSessionKey::new(peer_addr, conn_id), + sessions.clone(), + ); + + dispatch_data_packet( + &sessions, + unexpected_addr, + conn_id, + &new_data_packet(conn_id, b"wrong-peer").unwrap(), + Default::default(), + ); + dispatch_data_packet( + &sessions, + peer_addr, + conn_id + 1, + &new_data_packet(conn_id + 1, b"wrong-conn").unwrap(), + Default::default(), + ); + dispatch_data_packet( + &sessions, + peer_addr, + conn_id, + &new_data_packet(conn_id, b"payload").unwrap(), + Default::default(), + ); + + let mut buf = [0; 16]; + let len = tokio::time::timeout(Duration::from_secs(1), session.recv(&mut buf)) + .await + .unwrap() + .unwrap(); + + assert_eq!(&buf[..len], b"payload"); +} + +#[tokio::test] +async fn udp_session_layer_connects_with_shared_recv_loop() { + let local_addr = SocketAddr::from(([127, 0, 0, 1], 12000)); + let peer_addr = SocketAddr::from(([127, 0, 0, 1], 12001)); + let socket = Arc::new(AutoSackVirtualUdpSocket::new(local_addr)); + let layer = UdpSessionLayer::new(socket.clone()); + + let session = tokio::time::timeout(Duration::from_secs(1), layer.connect(peer_addr)) + .await + .unwrap() + .unwrap(); + + assert_eq!(layer.local_addr().unwrap(), local_addr); + assert_eq!(session.kind(), UdpSessionKind::EasyTierMux); + assert_eq!(session.peer_addr().unwrap(), peer_addr); + + let sent = socket.sent(); + assert!(!sent.is_empty()); + let packet = parse_udp_session_datagram(BytesMut::from(sent[0].0.as_slice()), false) + .expect("first sent datagram should be syn"); + assert_eq!( + packet.udp_tunnel_header().unwrap().msg_type, + UdpPacketType::Syn as u8 + ); +} + +#[tokio::test] +async fn cancelled_udp_session_connect_cleans_registered_state() { + let local_addr = SocketAddr::from(([127, 0, 0, 1], 12000)); + let peer_addr = SocketAddr::from(([127, 0, 0, 1], 12001)); + let socket = Arc::new(MockVirtualUdpSocket::new(local_addr, Vec::new())); + let layer = UdpSessionLayer::new(socket); + + let result = tokio::time::timeout(Duration::from_millis(50), layer.connect(peer_addr)).await; + + assert!(result.is_err()); + assert!(layer.pending_connects.is_empty()); + assert!(layer.sessions.is_empty()); +} + +#[tokio::test] +async fn dropping_udp_session_layer_closes_session_recv() { + let local_addr = SocketAddr::from(([127, 0, 0, 1], 12000)); + let peer_addr = SocketAddr::from(([127, 0, 0, 1], 12001)); + let socket = Arc::new(AutoSackVirtualUdpSocket::new(local_addr)); + let layer = UdpSessionLayer::new(socket.clone()); + let session = layer.connect(peer_addr).await.unwrap(); + drop(layer); + + let mut buf = [0; 16]; + let err = tokio::time::timeout(Duration::from_secs(1), session.recv(&mut buf)) + .await + .unwrap() + .unwrap_err(); + + assert_eq!(err.kind(), io::ErrorKind::UnexpectedEof); +} + +#[tokio::test] +async fn udp_session_recv_loop_error_closes_registered_sessions() { + let local_addr = SocketAddr::from(([127, 0, 0, 1], 12000)); + let peer_addr = SocketAddr::from(([127, 0, 0, 1], 12001)); + let conn_id = 0x1122_3344; + let socket = Arc::new(MockVirtualUdpSocket::new(local_addr, Vec::new())); + let sessions = Arc::new(DashMap::new()); + let (session_shutdown_tx, session_shutdown_rx) = watch::channel(false); + let session = create_test_easy_tier_mux_session_with_shutdown( + socket.clone(), + UdpSessionKey::new(peer_addr, conn_id), + sessions.clone(), + session_shutdown_rx, + ); + let pending_connects = Arc::new(DashMap::new()); + let classified_sessions = Arc::new(DashMap::new()); + let classified_accepts = create_classified_udp_session_accepts(); + let (mux_accepted_tx, _mux_accepted_rx) = mpsc::channel(UDP_SESSION_QUEUE_CAPACITY); + let (control_tx, _control_rx) = mpsc::channel(UDP_SESSION_QUEUE_CAPACITY); + + udp_session_layer_recv_task( + socket, + sessions.clone(), + classified_sessions.clone(), + classified_accepts, + pending_connects.clone(), + mux_accepted_tx, + control_tx, + Arc::new(NoopUdpSessionStunResponder), + session_shutdown_tx, + ) + .await; + + assert!(sessions.is_empty()); + assert!(classified_sessions.is_empty()); + assert!(pending_connects.is_empty()); + + let mut buf = [0; 16]; + let err = tokio::time::timeout(Duration::from_secs(1), session.recv(&mut buf)) + .await + .unwrap() + .unwrap_err(); + assert_eq!(err.kind(), io::ErrorKind::UnexpectedEof); + + let err = tokio::time::timeout(Duration::from_secs(1), async { + loop { + if let Err(err) = session.send(b"payload").await { + return err; + } + tokio::task::yield_now().await; + } + }) + .await + .unwrap(); + assert_eq!(err.kind(), io::ErrorKind::UnexpectedEof); +} + +#[tokio::test] +async fn udp_session_layer_accepts_syn_and_sends_sack() { + let local_addr = SocketAddr::from(([127, 0, 0, 1], 12000)); + let peer_addr = SocketAddr::from(([127, 0, 0, 1], 12001)); + let conn_id = 0x1122_3344; + let magic = 0x0102_0304_0506_0708; + let socket = Arc::new(MockVirtualUdpSocket::new(local_addr, Vec::new())); + let sessions = Arc::new(DashMap::new()); + let (mux_accepted_tx, mut mux_accepted_rx) = mpsc::channel(UDP_SESSION_QUEUE_CAPACITY); + let (_shutdown_tx, shutdown_rx) = watch::channel(false); + + handle_new_easy_tier_mux_connect( + socket.clone(), + sessions.clone(), + mux_accepted_tx, + peer_addr, + conn_id, + &new_syn_packet(conn_id, magic), + shutdown_rx, + ); + + let accepted = tokio::time::timeout(Duration::from_secs(1), mux_accepted_rx.recv()) + .await + .unwrap() + .unwrap(); + assert_eq!(accepted.kind(), UdpSessionKind::EasyTierMux); + assert_eq!(accepted.peer_addr().unwrap(), peer_addr); + assert!(sessions.contains_key(&UdpSessionKey::new(peer_addr, conn_id))); + + let sent = wait_for_sent(|| socket.sent(), 1).await; + assert_eq!(sent.len(), 1); + assert_eq!(sent[0].1, peer_addr); + let packet = parse_udp_session_datagram(BytesMut::from(sent[0].0.as_slice()), false) + .expect("sent datagram should be sack"); + let header = packet.udp_tunnel_header().unwrap(); + assert_eq!(header.conn_id.get(), conn_id); + assert_eq!(header.msg_type, UdpPacketType::Sack as u8); + assert_eq!(packet.udp_payload(), magic.to_le_bytes()); +} + +#[tokio::test] +async fn duplicate_syn_sack_send_failure_closes_existing_session() { + let local_addr = SocketAddr::from(([127, 0, 0, 1], 12000)); + let peer_addr = SocketAddr::from(([127, 0, 0, 1], 12001)); + let conn_id = 0x1122_3344; + let magic = 0x0102_0304_0506_0708; + let key = UdpSessionKey::new(peer_addr, conn_id); + let socket = Arc::new(FailingSendVirtualUdpSocket { local_addr }); + let sessions = Arc::new(DashMap::new()); + let (session, _shutdown_tx) = + create_test_easy_tier_mux_session(socket.clone(), key, sessions.clone()); + let (mux_accepted_tx, _mux_accepted_rx) = mpsc::channel(UDP_SESSION_QUEUE_CAPACITY); + let (_session_shutdown_tx, session_shutdown_rx) = watch::channel(false); + + handle_new_easy_tier_mux_connect( + socket, + sessions.clone(), + mux_accepted_tx, + peer_addr, + conn_id, + &new_syn_packet(conn_id, magic), + session_shutdown_rx, + ); + + let mut buf = [0; 16]; + let err = tokio::time::timeout(Duration::from_secs(1), session.recv(&mut buf)) + .await + .unwrap() + .unwrap_err(); + assert_eq!(err.kind(), io::ErrorKind::UnexpectedEof); + assert!(!sessions.contains_key(&key)); +} + +#[tokio::test] +async fn full_accept_queue_does_not_block_udp_session_recv_loop() { + let local_addr = SocketAddr::from(([127, 0, 0, 1], 12000)); + let peer_addr = SocketAddr::from(([127, 0, 0, 1], 12001)); + let conn_id = 0x1122_3344; + let magic = 0x0102_0304_0506_0708; + let socket = Arc::new(MockVirtualUdpSocket::new(local_addr, Vec::new())); + let sessions = Arc::new(DashMap::new()); + let full_sessions = Arc::new(DashMap::new()); + let (mux_accepted_tx, mut mux_accepted_rx) = mpsc::channel(1); + let (queued_session, _queued_shutdown_tx) = create_test_easy_tier_mux_session( + socket.clone(), + UdpSessionKey::new(SocketAddr::from(([127, 0, 0, 1], 12002)), 7), + full_sessions, + ); + mux_accepted_tx.try_send(queued_session).unwrap(); + let (_shutdown_tx, shutdown_rx) = watch::channel(false); + + handle_new_easy_tier_mux_connect( + socket.clone(), + sessions.clone(), + mux_accepted_tx, + peer_addr, + conn_id, + &new_syn_packet(conn_id, magic), + shutdown_rx, + ); + + assert!(mux_accepted_rx.try_recv().is_ok()); + assert!(!sessions.contains_key(&UdpSessionKey::new(peer_addr, conn_id))); + assert!(socket.sent().is_empty()); +} + +#[tokio::test] +async fn udp_session_layer_routes_stun_and_hole_punch_control_packets() { + let local_addr = SocketAddr::from(([127, 0, 0, 1], 12000)); + let stun_remote_addr = SocketAddr::from(([127, 0, 0, 1], 12001)); + let change_stun_remote_addr = SocketAddr::from(([127, 0, 0, 1], 12004)); + let rejected_remote_addr = SocketAddr::from(([192, 0, 2, 1], 12001)); + let v4_remote_addr = SocketAddr::from(([127, 0, 0, 1], 12002)); + let v6_remote_addr = "[::1]:12003".parse::().unwrap(); + let dst_v4 = SocketAddrV4::new(Ipv4Addr::new(192, 0, 2, 10), 1234); + let dst_v6 = "[2001:db8::1]:2345".parse::().unwrap(); + let preferred_src = PreferredIpv6Source { + ip: "2001:db8::2".parse().unwrap(), + ifindex: 42, + }; + let stun = easytier_stun_request(false, false); + let change_stun = easytier_stun_request(true, false); + let socket = Arc::new(MockVirtualUdpSocket::new( + local_addr, + vec![ + (stun.clone(), stun_remote_addr), + (change_stun.clone(), change_stun_remote_addr), + ( + new_v4_hole_punch_packet(&dst_v4).into_bytes().to_vec(), + rejected_remote_addr, + ), + ( + new_v4_hole_punch_packet(&dst_v4).into_bytes().to_vec(), + v4_remote_addr, + ), + ( + new_v6_hole_punch_packet(&dst_v6, Some(preferred_src)) + .into_bytes() + .to_vec(), + v6_remote_addr, + ), + ], + )); + socket + .reject_preferred_source + .store(true, Ordering::Relaxed); + let stun_responder = Arc::new(MockVirtualUdpSocketFactory::new(13000)); + let layer = UdpSessionLayer::new_with_stun_responder(socket.clone(), stun_responder.clone()); + + let mut events = Vec::new(); + for _ in 0..4 { + events.push( + tokio::time::timeout(Duration::from_secs(1), layer.recv_control()) + .await + .unwrap() + .unwrap(), + ); + } + assert!(events.contains(&UdpSessionLayerControl::Stun { + remote_addr: stun_remote_addr, + datagram: BytesMut::from(stun.as_slice()), + })); + assert!(events.contains(&UdpSessionLayerControl::Stun { + remote_addr: change_stun_remote_addr, + datagram: BytesMut::from(change_stun.as_slice()), + })); + assert!(events.contains(&UdpSessionLayerControl::V4HolePunch { + remote_addr: v4_remote_addr, + dst_addr: dst_v4, + })); + assert!(events.contains(&UdpSessionLayerControl::V6HolePunch { + remote_addr: v6_remote_addr, + dst_addr: dst_v6, + preferred_src: Some(preferred_src), + })); + tokio::time::timeout(Duration::from_secs(1), async { + loop { + let responder_sockets = stun_responder.sockets(); + let send_attempts = socket.send_attempts(); + let hole_punch_attempts = send_attempts + .iter() + .filter(|attempt| { + matches!(attempt.1, SocketAddr::V4(addr) if addr == dst_v4) + || matches!(attempt.1, SocketAddr::V6(addr) if addr == dst_v6) + }) + .count(); + if responder_sockets + .first() + .is_some_and(|socket| !socket.sent().is_empty()) + && hole_punch_attempts == 3 + { + return; + } + tokio::task::yield_now().await; + } + }) + .await + .unwrap(); + assert_eq!( + stun_responder.bind_options(), + vec![ + UdpBindOptions::hole_punch_control().with_local_addr(Some(SocketAddr::V4( + SocketAddrV4::new(Ipv4Addr::UNSPECIFIED, 0) + ))) + ] + ); + let responder_sockets = stun_responder.sockets(); + assert_eq!(responder_sockets.len(), 1); + assert_eq!( + responder_sockets[0].sent()[0].1, + change_stun_remote_addr, + "ChangeRequest STUN responses should be sent through a fresh socket" + ); + let attempts = socket + .send_attempts() + .into_iter() + .filter(|attempt| { + matches!(attempt.1, SocketAddr::V4(addr) if addr == dst_v4) + || matches!(attempt.1, SocketAddr::V6(addr) if addr == dst_v6) + }) + .collect::>(); + assert_eq!(attempts.len(), 3); + assert!(attempts.iter().all(|attempt| { + hole_punch_packet_tid(&attempt.0, UDP_SESSION_HOLE_PUNCH_PACKET_BODY_LEN) == Some(1) + })); + assert!(attempts.iter().any(|attempt| { + attempt.1 == SocketAddr::V4(dst_v4) && attempt.2 == UdpSocketSendMeta::default() + })); + let preferred_meta = UdpSocketSendMeta { + src_ip: Some(preferred_src.ip.into()), + src_ifindex: Some(preferred_src.ifindex), + }; + let preferred_index = attempts + .iter() + .position(|attempt| attempt.1 == SocketAddr::V6(dst_v6) && attempt.2 == preferred_meta) + .unwrap(); + let fallback_index = attempts + .iter() + .position(|attempt| { + attempt.1 == SocketAddr::V6(dst_v6) && attempt.2 == UdpSocketSendMeta::default() + }) + .unwrap(); + assert!(preferred_index < fallback_index); + assert_eq!(attempts[preferred_index].0, attempts[fallback_index].0); + assert!( + socket + .sent() + .iter() + .any(|(_, destination)| *destination == stun_remote_addr), + "normal STUN responses should use the listener socket" + ); +} + +#[tokio::test] +async fn udp_session_recv_loop_does_not_wait_for_stun_responder() { + let local_addr = SocketAddr::from(([127, 0, 0, 1], 12000)); + let stun_remote_addr = SocketAddr::from(([127, 0, 0, 1], 12001)); + let peer_addr = SocketAddr::from(([127, 0, 0, 1], 12002)); + let conn_id = 0x1122_3344; + let mut stun = vec![0; UDP_TUNNEL_HEADER_SIZE]; + stun[4..8].copy_from_slice(&[0x21, 0x12, 0xA4, 0x42]); + let socket = Arc::new(AutoSackVirtualUdpSocket::new(local_addr)); + socket.incoming.lock().unwrap().extend([ + (stun, stun_remote_addr), + ( + new_data_packet(conn_id, b"payload") + .unwrap() + .into_bytes() + .to_vec(), + peer_addr, + ), + ]); + socket.incoming_notify.notify_one(); + let sessions = Arc::new(DashMap::new()); + let (session, _shutdown_tx) = create_test_easy_tier_mux_session( + socket.clone(), + UdpSessionKey::new(peer_addr, conn_id), + sessions.clone(), + ); + let pending_connects = Arc::new(DashMap::new()); + let classified_sessions = Arc::new(DashMap::new()); + let classified_accepts = create_classified_udp_session_accepts(); + let (mux_accepted_tx, _mux_accepted_rx) = mpsc::channel(UDP_SESSION_QUEUE_CAPACITY); + let (control_tx, _control_rx) = mpsc::channel(UDP_SESSION_QUEUE_CAPACITY); + let stun_responder = Arc::new(BlockingUdpSessionStunResponder::default()); + let (session_shutdown_tx, _) = watch::channel(false); + let recv_task = tokio::spawn(udp_session_layer_recv_task( + socket, + sessions, + classified_sessions, + classified_accepts, + pending_connects, + mux_accepted_tx, + control_tx, + stun_responder.clone(), + session_shutdown_tx, + )); + + tokio::time::timeout(Duration::from_secs(1), stun_responder.started.notified()) + .await + .unwrap(); + + let mut buf = [0; 16]; + let len = tokio::time::timeout(Duration::from_secs(1), session.recv(&mut buf)) + .await + .unwrap() + .unwrap(); + assert_eq!(&buf[..len], b"payload"); + + stun_responder.release.notify_waiters(); + recv_task.abort(); + let _ = recv_task.await; +} + +#[tokio::test] +async fn local_hole_punch_control_is_dispatched_to_control_queue() { + let remote_addr = SocketAddr::from(([127, 0, 0, 1], 12001)); + let dst_addr = SocketAddrV4::new(Ipv4Addr::new(192, 0, 2, 10), 1234); + let (control_tx, mut control_rx) = mpsc::channel(1); + let socket = Arc::new(MockVirtualUdpSocket::new(remote_addr, vec![])); + + dispatch_v4_hole_punch_control( + socket, + Arc::new(Semaphore::new(UDP_SESSION_QUEUE_CAPACITY)), + &control_tx, + remote_addr, + &new_v4_hole_punch_packet(&dst_addr), + ); + + assert_eq!( + control_rx.recv().await.unwrap(), + UdpSessionLayerControl::V4HolePunch { + remote_addr, + dst_addr, + } + ); +} + +#[tokio::test] +async fn sack_from_actual_remote_rekeys_pending_session_before_data_dispatch() { + let expected_addr = SocketAddr::from(([127, 0, 0, 1], 12001)); + let actual_addr = SocketAddr::from(([127, 0, 0, 1], 12002)); + let conn_id = 0x1122_3344; + let magic = 0x0102_0304_0506_0708; + let sessions = DashMap::new(); + let pending_connects = DashMap::new(); + let expected_key = UdpSessionKey::new(expected_addr, conn_id); + let actual_key = UdpSessionKey::new(actual_addr, conn_id); + let session_key = Arc::new(StdMutex::new(None)); + let rings = create_udp_session_rings(); + let entry = udp_session_registry_entry(&rings); + let mut incoming_rx = rings.session_recv_rx; + let (control_tx, _control_rx) = mpsc::channel(1); + let (sack_tx, mut sack_rx) = watch::channel(None); + control_tx + .try_send(UdpConnectControl::HolePunch { + recv_addr: expected_addr, + }) + .unwrap(); + pending_connects.insert( + conn_id, + PendingUdpSessionConnect { + expected_addr, + magic, + session_key: session_key.clone(), + entry, + control: control_tx, + sack: sack_tx, + }, + ); + + dispatch_data_packet( + &sessions, + expected_addr, + conn_id, + &new_data_packet(conn_id, b"pre-sack").unwrap(), + Default::default(), + ); + dispatch_sack_packet( + &sessions, + &pending_connects, + actual_addr, + conn_id, + &new_sack_packet(conn_id, magic), + ); + dispatch_data_packet( + &sessions, + actual_addr, + conn_id, + &new_data_packet(conn_id, b"payload").unwrap(), + Default::default(), + ); + + assert!(sessions.contains_key(&actual_key)); + assert!(!sessions.contains_key(&expected_key)); + assert!(pending_connects.is_empty()); + assert_eq!(*session_key.lock().unwrap(), Some(actual_key)); + sack_rx.changed().await.unwrap(); + assert_eq!(*sack_rx.borrow_and_update(), Some(actual_addr)); + + let payload = futures::StreamExt::next(&mut incoming_rx) + .await + .unwrap() + .unwrap(); + assert_eq!(payload.payload, BytesMut::from(&b"payload"[..])); +} + +#[tokio::test] +async fn replayed_sack_cannot_rekey_pending_session_after_first_success() { + let expected_addr = SocketAddr::from(([127, 0, 0, 1], 12001)); + let first_addr = SocketAddr::from(([127, 0, 0, 1], 12002)); + let replay_addr = SocketAddr::from(([127, 0, 0, 1], 12003)); + let conn_id = 0x1122_3344; + let magic = 0x0102_0304_0506_0708; + let sessions = DashMap::new(); + let pending_connects = DashMap::new(); + let expected_key = UdpSessionKey::new(expected_addr, conn_id); + let first_key = UdpSessionKey::new(first_addr, conn_id); + let replay_key = UdpSessionKey::new(replay_addr, conn_id); + let session_key = Arc::new(StdMutex::new(None)); + let rings = create_udp_session_rings(); + let entry = udp_session_registry_entry(&rings); + let (control_tx, _control_rx) = mpsc::channel(UDP_SESSION_QUEUE_CAPACITY); + let (sack_tx, mut sack_rx) = watch::channel(None); + pending_connects.insert( + conn_id, + PendingUdpSessionConnect { + expected_addr, + magic, + session_key: session_key.clone(), + entry, + control: control_tx, + sack: sack_tx, + }, + ); + + dispatch_sack_packet( + &sessions, + &pending_connects, + first_addr, + conn_id, + &new_sack_packet(conn_id, magic), + ); + dispatch_sack_packet( + &sessions, + &pending_connects, + replay_addr, + conn_id, + &new_sack_packet(conn_id, magic), + ); + + assert!(sessions.contains_key(&first_key)); + assert!(!sessions.contains_key(&expected_key)); + assert!(!sessions.contains_key(&replay_key)); + assert_eq!(*session_key.lock().unwrap(), Some(first_key)); + sack_rx.changed().await.unwrap(); + assert_eq!(*sack_rx.borrow_and_update(), Some(first_addr)); +} + +#[tokio::test] +async fn stale_sack_after_pending_removal_does_not_register_session() { + let expected_addr = SocketAddr::from(([127, 0, 0, 1], 12001)); + let actual_addr = SocketAddr::from(([127, 0, 0, 1], 12002)); + let conn_id = 0x1122_3344; + let magic = 0x0102_0304_0506_0708; + let sessions = DashMap::new(); + let pending_connects = DashMap::new(); + let rings = create_udp_session_rings(); + let entry = udp_session_registry_entry(&rings); + let (control_tx, _control_rx) = mpsc::channel(UDP_SESSION_QUEUE_CAPACITY); + let (sack_tx, sack_rx) = watch::channel(None); + pending_connects.insert( + conn_id, + PendingUdpSessionConnect { + expected_addr, + magic, + session_key: Arc::new(StdMutex::new(None)), + entry, + control: control_tx, + sack: sack_tx, + }, + ); + pending_connects.remove(&conn_id); + + dispatch_sack_packet( + &sessions, + &pending_connects, + actual_addr, + conn_id, + &new_sack_packet(conn_id, magic), + ); + + assert!(sessions.is_empty()); + assert_eq!(*sack_rx.borrow(), None); +} + +#[tokio::test] +async fn sack_after_connect_receiver_drop_removes_registered_session() { + let expected_addr = SocketAddr::from(([127, 0, 0, 1], 12001)); + let actual_addr = SocketAddr::from(([127, 0, 0, 1], 12002)); + let conn_id = 0x1122_3344; + let magic = 0x0102_0304_0506_0708; + let sessions = DashMap::new(); + let pending_connects = DashMap::new(); + let rings = create_udp_session_rings(); + let entry = udp_session_registry_entry(&rings); + let (control_tx, _control_rx) = mpsc::channel(UDP_SESSION_QUEUE_CAPACITY); + let (sack_tx, sack_rx) = watch::channel(None); + drop(sack_rx); + pending_connects.insert( + conn_id, + PendingUdpSessionConnect { + expected_addr, + magic, + session_key: Arc::new(StdMutex::new(None)), + entry, + control: control_tx, + sack: sack_tx, + }, + ); + + dispatch_sack_packet( + &sessions, + &pending_connects, + actual_addr, + conn_id, + &new_sack_packet(conn_id, magic), + ); + + assert!(sessions.is_empty()); +} + +type MockIncomingDatagrams = VecDeque, SocketAddr)>>; + +struct MockVirtualUdpSocketFactory { + next_port: AtomicU16, + bind_options: Mutex>, + sockets: Mutex>>, + incoming: Mutex, +} + +impl MockVirtualUdpSocketFactory { + fn new(next_port: u16) -> Self { + Self { + next_port: AtomicU16::new(next_port), + bind_options: Mutex::new(Vec::new()), + sockets: Mutex::new(Vec::new()), + incoming: Mutex::new(VecDeque::new()), + } + } + + fn with_socket_incoming(next_port: u16, incoming: Vec<(Vec, SocketAddr)>) -> Self { + let factory = Self::new(next_port); + factory.incoming.lock().unwrap().push_back(incoming); + factory + } + + fn bind_options(&self) -> Vec { + self.bind_options.lock().unwrap().clone() + } + + fn sockets(&self) -> Vec> { + self.sockets.lock().unwrap().clone() + } +} + +#[async_trait] +impl VirtualUdpSocketFactory for MockVirtualUdpSocketFactory { + type Socket = MockVirtualUdpSocket; + + async fn bind_udp(&self, options: UdpBindOptions) -> anyhow::Result> { + self.bind_options.lock().unwrap().push(options.clone()); + let local_addr = options.local_addr.unwrap_or_else(|| { + SocketAddr::from(( + [127, 0, 0, 1], + self.next_port.fetch_add(1, Ordering::Relaxed), + )) + }); + let incoming = self + .incoming + .lock() + .unwrap() + .pop_front() + .unwrap_or_default(); + let socket = Arc::new(MockVirtualUdpSocket::new(local_addr, incoming)); + self.sockets.lock().unwrap().push(socket.clone()); + Ok(socket) + } +} + +#[tokio::test] +async fn udp_session_dialer_binds_socket_and_returns_wireguard_session() { + let factory = Arc::new(MockVirtualUdpSocketFactory::new(13000)); + let mut dialer = UdpSessionDialer::new(factory.clone()); + let remote_addr = SocketAddr::from(([192, 0, 2, 10], 11010)); + let bind_addr = SocketAddr::from(([127, 0, 0, 1], 14000)); + let request = UdpSessionConnectRequest::wireguard(remote_addr) + .with_bind(UdpBindOptions::port_bound_listener(bind_addr)); + let expected_bind = request.bind.clone(); + + let session = dialer.connect(request).await.unwrap(); + + assert_eq!(factory.bind_options(), vec![expected_bind]); + assert_eq!(session.kind(), UdpSessionKind::WireGuard); + assert_eq!(session.local_addr().unwrap(), bind_addr); + assert_eq!(session.peer_addr().unwrap(), remote_addr); +} + +#[tokio::test] +async fn udp_session_socket_listener_builds_port_bound_bind_options() { + let factory = Arc::new(MockVirtualUdpSocketFactory::new(13000)); + let local_addr = SocketAddr::from(([0, 0, 0, 0], 11010)); + let mut listener = UdpSessionSocketListener::new( + "udp://0.0.0.0:0".parse().unwrap(), + local_addr, + factory.clone(), + ); + + listener.listen().await.unwrap(); + + assert_eq!( + factory.bind_options(), + vec![UdpBindOptions::port_bound_listener(local_addr).with_only_v6(true)] + ); + assert_eq!(listener.local_url().port(), Some(11010)); + assert_eq!(listener.connection_counter().get(), Some(0)); + assert!(Arc::ptr_eq( + &listener.bound_socket().unwrap(), + &factory.sockets()[0] + )); +} + +#[tokio::test] +async fn udp_session_socket_listener_accepts_easy_tier_mux_session() { + let local_addr = SocketAddr::from(([127, 0, 0, 1], 11010)); + let peer_addr = SocketAddr::from(([127, 0, 0, 1], 12010)); + let factory = Arc::new(MockVirtualUdpSocketFactory::with_socket_incoming( + 13000, + vec![( + new_syn_packet(0x1122_3344, 0x5566_7788) + .into_bytes() + .to_vec(), + peer_addr, + )], + )); + let mut listener = + UdpSessionSocketListener::new("udp://127.0.0.1:0".parse().unwrap(), local_addr, factory); + + listener.listen().await.unwrap(); + let session = tokio::time::timeout(Duration::from_secs(1), listener.accept_session()) + .await + .unwrap() + .unwrap(); + + assert_eq!(session.kind(), UdpSessionKind::EasyTierMux); + assert_eq!(session.local_addr().unwrap(), local_addr); + assert_eq!(session.peer_addr().unwrap(), peer_addr); +} + +#[tokio::test] +async fn udp_session_dialer_uses_factory_as_stun_responder() { + let remote_addr = SocketAddr::from(([192, 0, 2, 10], 11010)); + let stun_remote_addr = SocketAddr::from(([127, 0, 0, 1], 12001)); + let request = UdpSessionConnectRequest::wireguard(remote_addr); + let expected_session_bind = request.bind.clone(); + let factory = Arc::new(MockVirtualUdpSocketFactory::with_socket_incoming( + 13000, + vec![(easytier_stun_request(true, false), stun_remote_addr)], + )); + let mut dialer = UdpSessionDialer::new(factory.clone()); + + let _session = dialer.connect(request).await.unwrap(); + + tokio::time::timeout(Duration::from_secs(1), async { + loop { + let sockets = factory.sockets(); + if sockets.len() == 2 && !sockets[1].sent().is_empty() { + return; + } + tokio::task::yield_now().await; + } + }) + .await + .unwrap(); + assert_eq!( + factory.bind_options(), + vec![ + expected_session_bind, + UdpBindOptions::hole_punch_control().with_local_addr(Some(SocketAddr::V4( + SocketAddrV4::new(Ipv4Addr::UNSPECIFIED, 0) + ))) + ] + ); + let sockets = factory.sockets(); + assert_eq!(sockets[1].sent()[0].1, stun_remote_addr); +} + +#[tokio::test] +async fn v4_hole_punch_control_sender_uses_factory_socket() { + let factory = MockVirtualUdpSocketFactory::new(13000); + let dst_addr = SocketAddrV4::new(Ipv4Addr::new(192, 0, 2, 1), 11010); + let context = SocketContext::default() + .with_socket_mark(Some(0)) + .with_netns(Some(NetNamespace::new("instance-a"))); + + send_v4_hole_punch_control_packet(&factory, context.clone(), 22020, dst_addr) + .await + .unwrap(); + + assert_eq!( + factory.bind_options(), + vec![ + UdpBindOptions::hole_punch_control() + .with_context(context.with_ip_version(IpVersion::V4)) + .with_local_addr(Some(SocketAddr::V4(SocketAddrV4::new( + Ipv4Addr::LOCALHOST, + 0 + )))) + ] + ); + let sockets = factory.sockets(); + assert_eq!(sockets.len(), 1); + assert_eq!( + sockets[0].sent(), + vec![( + new_v4_hole_punch_packet(&dst_addr).into_bytes().to_vec(), + SocketAddr::V4(SocketAddrV4::new(Ipv4Addr::LOCALHOST, 22020)) + )] + ); +} + +#[tokio::test] +async fn v6_hole_punch_control_sender_uses_factory_socket() { + let factory = MockVirtualUdpSocketFactory::new(13000); + let dst_addr = "[2001:db8::1]:11010".parse::().unwrap(); + let preferred_src = PreferredIpv6Source { + ip: "2001:db8::2".parse().unwrap(), + ifindex: 42, + }; + let context = SocketContext::default() + .with_socket_mark(Some(0)) + .with_netns(Some(NetNamespace::new("instance-a"))); + + send_v6_hole_punch_control_packet( + &factory, + context.clone(), + 22020, + dst_addr, + Some(preferred_src), + ) + .await + .unwrap(); + + assert_eq!( + factory.bind_options(), + vec![ + UdpBindOptions::hole_punch_control() + .with_context(context.with_ip_version(IpVersion::V6)) + .with_local_addr(Some(SocketAddr::V6(SocketAddrV6::new( + Ipv6Addr::LOCALHOST, + 0, + 0, + 0 + )))) + ] + ); + let sockets = factory.sockets(); + assert_eq!(sockets.len(), 1); + assert_eq!( + sockets[0].sent(), + vec![( + new_v6_hole_punch_packet(&dst_addr, Some(preferred_src)) + .into_bytes() + .to_vec(), + SocketAddr::V6(SocketAddrV6::new(Ipv6Addr::LOCALHOST, 22020, 0, 0)) + )] + ); +} + +#[test] +fn builds_syn_and_sack_packets_without_changing_wire_shape() { + let conn_id = 0x1234_5678; + let magic = 0x0102_0304_0506_0708; + + for (packet, msg_type) in [ + (new_syn_packet(conn_id, magic), UdpPacketType::Syn as u8), + (new_sack_packet(conn_id, magic), UdpPacketType::Sack as u8), + ] { + let header = packet.udp_tunnel_header().unwrap(); + assert_eq!(header.conn_id.get(), conn_id); + assert_eq!(header.msg_type, msg_type); + assert_eq!(header.len.get(), 8); + assert_eq!(packet.udp_payload(), magic.to_le_bytes()); + } +} + +#[test] +fn v6_hole_punch_packet_preserves_preferred_source() { + let dst_addr = "[2001:db8::1]:10001".parse::().unwrap(); + let preferred_src = PreferredIpv6Source { + ip: "2001:db8::2".parse().unwrap(), + ifindex: 42, + }; + + let packet = new_v6_hole_punch_packet(&dst_addr, Some(preferred_src)); + let (parsed_dst_addr, parsed_preferred_src) = + extract_v6_hole_punch_packet(packet.udp_payload()).unwrap(); + + assert_eq!(parsed_dst_addr, dst_addr); + assert_eq!(parsed_preferred_src, Some(preferred_src)); +} + +#[test] +fn parses_udp_session_datagram_and_rejects_bad_payload_len() { + let packet = new_syn_packet(7, 42).into_bytes(); + let parsed = parse_udp_session_datagram(packet.clone().into(), false).unwrap(); + assert_eq!(parsed.udp_tunnel_header().unwrap().conn_id.get(), 7); + + let mut bad_packet = packet.to_vec(); + bad_packet.pop(); + + assert!(matches!( + parse_udp_session_datagram(BytesMut::from(bad_packet.as_slice()), false), + Err(UdpSessionPacketError::PayloadLenMismatch { .. }) + )); +} + +#[test] +fn inspects_easytier_udp_datagram_without_owning_buffer() { + let packet = new_syn_packet(7, 42).into_bytes(); + let info = inspect_easytier_udp_datagram(&packet).unwrap().unwrap(); + + assert_eq!(info.kind, EasyTierUdpPacketKind::Syn); + assert_eq!(info.conn_id, 7); + + let unknown_packet = new_udp_packet( + |header| { + header.conn_id.set(9); + header.msg_type = 0xff; + header.len.set(0); + }, + &[], + ) + .into_bytes(); + assert_eq!( + inspect_easytier_udp_datagram(&unknown_packet).unwrap(), + None + ); + + let mut bad_packet = packet.to_vec(); + bad_packet.pop(); + + assert!(matches!( + inspect_easytier_udp_datagram(&bad_packet), + Err(EasyTierUdpDatagramInspectError::PayloadLenMismatch { .. }) + )); +} + +#[test] +fn stun_classifier_requires_cookie_and_stun_bits() { + let mut stun = [0; UDP_TUNNEL_HEADER_SIZE]; + stun[4..8].copy_from_slice(&[0x21, 0x12, 0xA4, 0x42]); + + assert!(is_stun_packet(&stun)); + + stun[0] = 0xC0; + assert!(!is_stun_packet(&stun)); + assert!(!is_stun_packet(&stun[..UDP_TUNNEL_HEADER_SIZE - 1])); +} diff --git a/easytier-core/src/socket/udp/virtual_socket.rs b/easytier-core/src/socket/udp/virtual_socket.rs new file mode 100644 index 00000000..9ace385c --- /dev/null +++ b/easytier-core/src/socket/udp/virtual_socket.rs @@ -0,0 +1,265 @@ +use std::{ + io, + net::{IpAddr, Ipv4Addr, Ipv6Addr, SocketAddr, SocketAddrV4, SocketAddrV6}, + sync::Arc, +}; + +use async_trait::async_trait; +use serde::{Deserialize, Serialize}; + +use crate::socket::{IpVersion, SocketContext}; + +use super::packet::{new_v4_hole_punch_packet, new_v6_hole_punch_packet}; + +#[derive(Debug, Default, Clone, Copy, PartialEq, Eq)] +pub struct UdpSocketRecvMeta { + pub dst_ip: Option, +} + +#[derive(Debug, Default, Clone, Copy, PartialEq, Eq)] +pub struct UdpSocketSendMeta { + pub src_ip: Option, + pub src_ifindex: Option, +} + +#[async_trait] +pub trait VirtualUdpSocket: Send + Sync + 'static { + fn local_addr(&self) -> std::io::Result; + + fn socket_context(&self) -> SocketContext { + SocketContext::default() + } + + async fn send_to(&self, data: &[u8], addr: SocketAddr) -> std::io::Result; + + async fn recv_from(&self, buf: &mut [u8]) -> std::io::Result<(usize, SocketAddr)>; + + async fn send_to_with_meta( + &self, + data: &[u8], + addr: SocketAddr, + meta: UdpSocketSendMeta, + ) -> std::io::Result { + let _ = meta; + self.send_to(data, addr).await + } + + async fn recv_from_with_meta( + &self, + buf: &mut [u8], + ) -> std::io::Result<(usize, SocketAddr, UdpSocketRecvMeta)> { + let (len, addr) = self.recv_from(buf).await?; + Ok((len, addr, UdpSocketRecvMeta::default())) + } +} + +#[async_trait] +pub trait UdpSessionStunResponder: Send + Sync + 'static +where + S: VirtualUdpSocket, +{ + async fn respond_stun( + &self, + _socket: Arc, + _datagram: &[u8], + _remote_addr: SocketAddr, + ) -> io::Result<()> { + Ok(()) + } +} + +#[derive(Debug, Default)] +pub struct NoopUdpSessionStunResponder; + +#[async_trait] +impl UdpSessionStunResponder for NoopUdpSessionStunResponder where S: VirtualUdpSocket {} + +pub async fn send_v4_hole_punch_control_packet( + factory: &F, + context: SocketContext, + listener_port: u16, + dst_addr: SocketAddrV4, +) -> anyhow::Result<()> +where + F: VirtualUdpSocketFactory, +{ + let socket = factory + .bind_udp( + UdpBindOptions::hole_punch_control() + .with_context(context.with_ip_version(IpVersion::V4)) + .with_local_addr(Some(SocketAddr::V4(SocketAddrV4::new( + Ipv4Addr::LOCALHOST, + 0, + )))), + ) + .await?; + let packet = new_v4_hole_punch_packet(&dst_addr).into_bytes(); + let listener_addr = SocketAddr::V4(SocketAddrV4::new(Ipv4Addr::LOCALHOST, listener_port)); + socket.send_to(&packet, listener_addr).await?; + Ok(()) +} + +pub async fn send_v6_hole_punch_control_packet( + factory: &F, + context: SocketContext, + listener_port: u16, + dst_addr: SocketAddrV6, + preferred_src: Option, +) -> anyhow::Result<()> +where + F: VirtualUdpSocketFactory, +{ + let socket = factory + .bind_udp( + UdpBindOptions::hole_punch_control() + .with_context(context.with_ip_version(IpVersion::V6)) + .with_local_addr(Some(SocketAddr::V6(SocketAddrV6::new( + Ipv6Addr::LOCALHOST, + 0, + 0, + 0, + )))), + ) + .await?; + let packet = new_v6_hole_punch_packet(&dst_addr, preferred_src).into_bytes(); + let listener_addr = SocketAddr::V6(SocketAddrV6::new(Ipv6Addr::LOCALHOST, listener_port, 0, 0)); + socket.send_to(&packet, listener_addr).await?; + Ok(()) +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] +pub enum UdpSocketPurpose { + HolePunchControl, + HolePunchCandidate, + DirectConnect, + PortBoundListener, + ProxyNat, + StunProbe, + Socks5, + PortForward, + PortLease, +} + +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +pub struct UdpBindOptions { + #[serde(default)] + pub context: SocketContext, + pub local_addr: Option, + pub bind_device: Option, + pub reuse_addr: bool, + pub reuse_port: bool, + pub only_v6: bool, + pub purpose: UdpSocketPurpose, +} + +impl UdpBindOptions { + fn for_purpose(purpose: UdpSocketPurpose) -> Self { + Self { + context: SocketContext::default(), + local_addr: None, + bind_device: None, + reuse_addr: false, + reuse_port: false, + only_v6: false, + purpose, + } + } + + pub fn hole_punch_control() -> Self { + Self::for_purpose(UdpSocketPurpose::HolePunchControl) + } + + pub fn hole_punch_candidate() -> Self { + Self::for_purpose(UdpSocketPurpose::HolePunchCandidate) + } + + pub fn direct_connect() -> Self { + Self::for_purpose(UdpSocketPurpose::DirectConnect) + } + + pub fn port_bound_listener(local_addr: SocketAddr) -> Self { + Self { + local_addr: Some(local_addr), + ..Self::for_purpose(UdpSocketPurpose::PortBoundListener) + } + } + + pub fn proxy_nat() -> Self { + Self::for_purpose(UdpSocketPurpose::ProxyNat) + } + + pub fn stun_probe() -> Self { + Self::for_purpose(UdpSocketPurpose::StunProbe) + } + + pub fn socks5() -> Self { + Self::for_purpose(UdpSocketPurpose::Socks5) + } + + pub fn port_forward(local_addr: SocketAddr) -> Self { + Self::for_purpose(UdpSocketPurpose::PortForward).with_local_addr(Some(local_addr)) + } + + pub fn port_lease(local_addr: SocketAddr) -> Self { + Self::for_purpose(UdpSocketPurpose::PortLease).with_local_addr(Some(local_addr)) + } + + pub fn with_local_addr(mut self, local_addr: Option) -> Self { + self.local_addr = local_addr; + self + } + + pub fn with_socket_mark(mut self, socket_mark: Option) -> Self { + self.context.socket_mark = socket_mark; + self + } + + pub fn with_context(mut self, context: SocketContext) -> Self { + self.context = context; + self + } + + pub fn with_ip_version(mut self, ip_version: IpVersion) -> Self { + self.context.ip_version = ip_version; + self + } + + pub fn with_bind_device(mut self, bind_device: Option) -> Self { + self.bind_device = bind_device; + self + } + + pub fn with_reuse_addr(mut self, reuse_addr: bool) -> Self { + self.reuse_addr = reuse_addr; + self + } + + pub fn with_reuse_port(mut self, reuse_port: bool) -> Self { + self.reuse_port = reuse_port; + self + } + + pub fn with_only_v6(mut self, only_v6: bool) -> Self { + self.only_v6 = only_v6; + self + } +} + +impl Default for UdpBindOptions { + fn default() -> Self { + Self::hole_punch_control() + } +} + +#[async_trait] +pub trait VirtualUdpSocketFactory: Send + Sync + 'static { + type Socket: VirtualUdpSocket; + + async fn bind_udp(&self, options: UdpBindOptions) -> anyhow::Result>; +} + +#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize, Deserialize)] +pub struct PreferredIpv6Source { + pub ip: Ipv6Addr, + pub ifindex: u32, +} diff --git a/easytier/src/peers/encrypt/aes_gcm.rs b/easytier-core/src/tunnel/encrypt/aes_gcm.rs similarity index 96% rename from easytier/src/peers/encrypt/aes_gcm.rs rename to easytier-core/src/tunnel/encrypt/aes_gcm.rs index c1018d9f..a90548fd 100644 --- a/easytier/src/peers/encrypt/aes_gcm.rs +++ b/easytier-core/src/tunnel/encrypt/aes_gcm.rs @@ -2,7 +2,7 @@ use aes_gcm::{AeadCore, AeadInPlace, Aes128Gcm, Aes256Gcm, Key, KeyInit}; use rand::rngs::OsRng; use zerocopy::{AsBytes, FromBytes}; -use crate::tunnel::packet_def::{StandardAeadTail, ZCPacket}; +use crate::packet::{StandardAeadTail, ZCPacket}; use super::{Encryptor, Error}; @@ -137,8 +137,8 @@ impl Encryptor for AesGcmCipher { #[cfg(test)] mod tests { use crate::{ - peers::encrypt::{Encryptor, aes_gcm::AesGcmCipher}, - tunnel::packet_def::{StandardAeadTail, ZCPacket}, + packet::{StandardAeadTail, ZCPacket}, + tunnel::encrypt::{Encryptor, aes_gcm::AesGcmCipher}, }; use zerocopy::FromBytes; diff --git a/easytier-core/src/tunnel/encrypt/chacha20.rs b/easytier-core/src/tunnel/encrypt/chacha20.rs new file mode 100644 index 00000000..3fbaa51d --- /dev/null +++ b/easytier-core/src/tunnel/encrypt/chacha20.rs @@ -0,0 +1,150 @@ +use chacha20poly1305::{AeadCore, AeadInPlace, ChaCha20Poly1305, Key, KeyInit}; +use rand::rngs::OsRng; +use zerocopy::{AsBytes, FromBytes}; + +use crate::packet::{StandardAeadTail, ZCPacket}; + +use super::{Encryptor, Error}; + +#[derive(Clone)] +pub struct ChaCha20Cipher { + cipher: Box, +} + +impl ChaCha20Cipher { + pub fn new(key: [u8; 32]) -> Self { + let key: &Key = &key.into(); + Self { + cipher: Box::new(ChaCha20Poly1305::new(key)), + } + } +} + +impl Encryptor for ChaCha20Cipher { + fn decrypt(&self, zc_packet: &mut ZCPacket) -> Result<(), Error> { + let pm_header = zc_packet.peer_manager_header().unwrap(); + if !pm_header.is_encrypted() { + return Ok(()); + } + + let payload_len = zc_packet.payload().len(); + if payload_len < StandardAeadTail::SIZE { + return Err(Error::PacketTooShort(zc_packet.payload().len())); + } + + let text_len = payload_len - StandardAeadTail::SIZE; + + let tail = StandardAeadTail::ref_from_suffix(zc_packet.payload()) + .unwrap() + .clone(); + + let nonce = tail.nonce.into(); + let tag = tail.tag.into(); + + self.cipher + .decrypt_in_place_detached(&nonce, &[], &mut zc_packet.mut_payload()[..text_len], &tag) + .map_err(|_| Error::DecryptionFailed)?; + + let pm_header = zc_packet.mut_peer_manager_header().unwrap(); + pm_header.set_encrypted(false); + let old_len = zc_packet.buf_len(); + zc_packet + .mut_inner() + .truncate(old_len - StandardAeadTail::SIZE); + Ok(()) + } + + fn encrypt(&self, zc_packet: &mut ZCPacket) -> Result<(), Error> { + self.encrypt_with_nonce(zc_packet, None) + } + + fn encrypt_with_nonce( + &self, + zc_packet: &mut ZCPacket, + nonce: Option<&[u8]>, + ) -> Result<(), Error> { + let pm_header = zc_packet.peer_manager_header().unwrap(); + if pm_header.is_encrypted() { + tracing::warn!(?zc_packet, "packet is already encrypted"); + return Ok(()); + } + + let nonce = nonce + .map(|n| { + <[u8; StandardAeadTail::NONCE_SIZE]>::try_from(n) + .map(Into::into) + .map_err(|_| Error::EncryptionFailed) + }) + .transpose()? + .unwrap_or_else(|| ChaCha20Poly1305::generate_nonce(&mut OsRng)); + + let tag = self + .cipher + .encrypt_in_place_detached(&nonce, &[], zc_packet.mut_payload()) + .map_err(|_| Error::EncryptionFailed)?; + + let tail = StandardAeadTail { + tag: tag.into(), + nonce: nonce.into(), + }; + + let pm_header = zc_packet.mut_peer_manager_header().unwrap(); + pm_header.set_encrypted(true); + zc_packet.mut_inner().extend_from_slice(tail.as_bytes()); + Ok(()) + } +} + +#[cfg(test)] +mod tests { + use crate::{ + packet::{StandardAeadTail, ZCPacket}, + tunnel::encrypt::{Encryptor, chacha20::ChaCha20Cipher}, + }; + use zerocopy::FromBytes; + + #[test] + fn test_chacha20_cipher() { + let key = [7u8; 32]; + let cipher = ChaCha20Cipher::new(key); + let text = b"Hello, ChaCha20"; + let mut packet = ZCPacket::new_with_payload(text); + packet.fill_peer_manager_hdr(0, 0, 0); + + cipher.encrypt(&mut packet).unwrap(); + assert_eq!(packet.payload().len(), text.len() + StandardAeadTail::SIZE); + assert!(packet.peer_manager_header().unwrap().is_encrypted()); + + cipher.decrypt(&mut packet).unwrap(); + assert_eq!(packet.payload(), text); + assert!(!packet.peer_manager_header().unwrap().is_encrypted()); + } + + #[test] + fn test_chacha20_cipher_with_nonce() { + let key = [7u8; 32]; + let cipher = ChaCha20Cipher::new(key); + let text = b"Hello"; + let nonce = [3u8; 12]; + + let mut packet1 = ZCPacket::new_with_payload(text); + packet1.fill_peer_manager_hdr(0, 0, 0); + cipher + .encrypt_with_nonce(&mut packet1, Some(&nonce)) + .unwrap(); + + let mut packet2 = ZCPacket::new_with_payload(text); + packet2.fill_peer_manager_hdr(0, 0, 0); + cipher + .encrypt_with_nonce(&mut packet2, Some(&nonce)) + .unwrap(); + + assert_eq!(packet1.payload(), packet2.payload()); + + let tail = StandardAeadTail::ref_from_suffix(packet1.payload()).unwrap(); + assert_eq!(tail.nonce, nonce); + + cipher.decrypt(&mut packet1).unwrap(); + assert_eq!(packet1.payload(), text); + } +} diff --git a/easytier-core/src/tunnel/encrypt/mod.rs b/easytier-core/src/tunnel/encrypt/mod.rs new file mode 100644 index 00000000..e6e1348f --- /dev/null +++ b/easytier-core/src/tunnel/encrypt/mod.rs @@ -0,0 +1,272 @@ +use crate::{config::EncryptionAlgorithm, packet::ZCPacket}; +use std::{collections::hash_map::DefaultHasher, hash::Hasher, sync::Arc}; + +#[cfg(feature = "aes-gcm")] +pub mod aes_gcm; +#[cfg(feature = "chacha20")] +pub mod chacha20; + +pub mod xor; + +// The disabled backends keep the same error Interface as the AEAD backends. +#[allow(dead_code)] +#[derive(thiserror::Error, Debug)] +pub enum Error { + #[error("packet is too short. len: {0}")] + PacketTooShort(usize), + #[error("decryption failed")] + DecryptionFailed, + #[error("encryption failed")] + EncryptionFailed, + #[error("invalid encryption algorithm: {0}")] + InvalidAlgorithm(String), + #[error("encryption algorithm is unavailable in this build: {0}")] + AlgorithmUnavailable(String), +} + +pub trait Encryptor: Send + Sync + 'static { + fn decrypt(&self, zc_packet: &mut ZCPacket) -> Result<(), Error>; + fn encrypt(&self, zc_packet: &mut ZCPacket) -> Result<(), Error>; + fn encrypt_with_nonce( + &self, + zc_packet: &mut ZCPacket, + _nonce: Option<&[u8]>, + ) -> Result<(), Error> { + self.encrypt(zc_packet) + } +} + +pub struct NullCipher; + +struct UnsupportedCipher { + algorithm: String, + unavailable: bool, +} + +pub fn derive_key_128(secret: &str) -> [u8; 16] { + let mut key = [0u8; 16]; + let mut hasher = DefaultHasher::new(); + hasher.write(secret.as_bytes()); + key[0..8].copy_from_slice(&hasher.finish().to_be_bytes()); + hasher.write(&key[0..8]); + key[8..16].copy_from_slice(&hasher.finish().to_be_bytes()); + hasher.write(&key); + key +} + +pub fn derive_key_256(secret: &str) -> [u8; 32] { + let mut key = [0u8; 32]; + let mut hasher = DefaultHasher::new(); + hasher.write(secret.as_bytes()); + hasher.write(b"easytier-256bit-key"); + for i in 0..4 { + let chunk_start = i * 8; + let chunk_end = chunk_start + 8; + hasher.write(&key[0..chunk_start]); + hasher.write(&[i as u8]); + key[chunk_start..chunk_end].copy_from_slice(&hasher.finish().to_be_bytes()); + } + key +} + +impl Encryptor for NullCipher { + fn decrypt(&self, zc_packet: &mut ZCPacket) -> Result<(), Error> { + let pm_header = zc_packet.peer_manager_header().unwrap(); + if pm_header.is_encrypted() { + Err(Error::DecryptionFailed) + } else { + Ok(()) + } + } + + fn encrypt(&self, _zc_packet: &mut ZCPacket) -> Result<(), Error> { + Ok(()) + } +} + +impl UnsupportedCipher { + fn error(&self) -> Error { + if self.unavailable { + Error::AlgorithmUnavailable(self.algorithm.clone()) + } else { + Error::InvalidAlgorithm(self.algorithm.clone()) + } + } +} + +impl Encryptor for UnsupportedCipher { + fn decrypt(&self, _zc_packet: &mut ZCPacket) -> Result<(), Error> { + Err(self.error()) + } + + fn encrypt(&self, _zc_packet: &mut ZCPacket) -> Result<(), Error> { + Err(self.error()) + } +} + +fn invalid_encryptor(algorithm: &str) -> Arc { + Arc::new(UnsupportedCipher { + algorithm: algorithm.to_owned(), + unavailable: false, + }) +} + +#[allow(dead_code)] // Selected disabled backends call this in reduced profiles. +fn unavailable_encryptor(algorithm: &str) -> Arc { + Arc::new(UnsupportedCipher { + algorithm: algorithm.to_owned(), + unavailable: true, + }) +} + +fn algorithm_is_available(algorithm: EncryptionAlgorithm) -> bool { + match algorithm { + EncryptionAlgorithm::Xor => true, + EncryptionAlgorithm::AesGcm | EncryptionAlgorithm::Aes256Gcm => cfg!(feature = "aes-gcm"), + EncryptionAlgorithm::ChaCha20 => cfg!(feature = "chacha20"), + } +} + +fn create_aes_128(key: [u8; 16]) -> Arc { + #[cfg(feature = "aes-gcm")] + { + Arc::new(aes_gcm::AesGcmCipher::new_128(key)) + } + #[cfg(not(feature = "aes-gcm"))] + { + let _ = key; + unavailable_encryptor("aes-gcm") + } +} + +fn create_aes_256(key: [u8; 32]) -> Arc { + #[cfg(feature = "aes-gcm")] + { + Arc::new(aes_gcm::AesGcmCipher::new_256(key)) + } + #[cfg(not(feature = "aes-gcm"))] + { + let _ = key; + unavailable_encryptor("aes-256-gcm") + } +} + +fn create_chacha20(key: [u8; 32]) -> Arc { + #[cfg(feature = "chacha20")] + { + Arc::new(chacha20::ChaCha20Cipher::new(key)) + } + #[cfg(not(feature = "chacha20"))] + { + let _ = key; + unavailable_encryptor("chacha20") + } +} + +pub(crate) fn validate_algorithm(algorithm: &str) -> Result<(), Error> { + let parsed = algorithm + .parse::() + .map_err(|()| Error::InvalidAlgorithm(algorithm.to_owned()))?; + if algorithm_is_available(parsed) { + Ok(()) + } else { + Err(Error::AlgorithmUnavailable(parsed.to_string())) + } +} + +pub(super) fn effective_algorithm_uses_xor(algorithm: &str) -> bool { + algorithm.parse() == Ok(EncryptionAlgorithm::Xor) +} + +/// Create an encryptor based on the algorithm name. +/// +/// Callers that accept user configuration validate it during construction. +/// Protocol paths remain infallible here and receive an encryptor that returns +/// an explicit error if a peer names an invalid or unavailable algorithm. +pub fn create_encryptor( + algorithm: &str, + key_128: [u8; 16], + key_256: [u8; 32], +) -> Arc { + let Ok(algorithm) = algorithm.parse::() else { + return invalid_encryptor(algorithm); + }; + + match algorithm { + EncryptionAlgorithm::Xor => Arc::new(xor::XorCipher::new(&key_128)), + EncryptionAlgorithm::AesGcm => create_aes_128(key_128), + EncryptionAlgorithm::Aes256Gcm => create_aes_256(key_256), + EncryptionAlgorithm::ChaCha20 => create_chacha20(key_256), + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn network_secret_key_derivation_is_stable() { + assert_eq!( + derive_key_128("secret"), + [ + 86, 90, 25, 219, 78, 240, 193, 33, 168, 172, 88, 14, 218, 248, 78, 166, + ] + ); + assert_eq!( + derive_key_256("secret"), + [ + 199, 205, 248, 94, 194, 101, 97, 138, 79, 69, 167, 248, 140, 5, 165, 163, 192, 139, + 166, 217, 166, 152, 28, 230, 146, 109, 150, 196, 66, 242, 231, 140, + ] + ); + } + + #[test] + fn effective_algorithm_only_reports_explicit_xor() { + assert!(effective_algorithm_uses_xor("xor")); + assert!(!effective_algorithm_uses_xor("")); + assert!(!effective_algorithm_uses_xor("unsupported")); + assert!(!effective_algorithm_uses_xor("aes-gcm")); + } + + #[cfg(not(feature = "aes-gcm"))] + #[test] + fn unavailable_aes_is_known_but_rejected() { + assert_eq!( + validate_algorithm("aes-gcm").unwrap_err().to_string(), + "encryption algorithm is unavailable in this build: aes-gcm" + ); + } + + #[cfg(feature = "aes-gcm")] + #[test] + fn compiled_aes_is_available() { + validate_algorithm("aes-gcm").unwrap(); + validate_algorithm("aes-256-gcm").unwrap(); + } + + #[cfg(not(feature = "chacha20"))] + #[test] + fn unavailable_chacha20_is_known_but_rejected() { + assert_eq!( + validate_algorithm("chacha20-poly1305") + .unwrap_err() + .to_string(), + "encryption algorithm is unavailable in this build: chacha20" + ); + } + + #[cfg(feature = "chacha20")] + #[test] + fn compiled_chacha20_is_available() { + validate_algorithm("chacha20").unwrap(); + } + + #[test] + fn invalid_algorithm_is_rejected() { + assert_eq!( + validate_algorithm("rot13").unwrap_err().to_string(), + "invalid encryption algorithm: rot13" + ); + } +} diff --git a/easytier/src/peers/encrypt/xor.rs b/easytier-core/src/tunnel/encrypt/xor.rs similarity index 89% rename from easytier/src/peers/encrypt/xor.rs rename to easytier-core/src/tunnel/encrypt/xor.rs index a6d3165c..ad8b1108 100644 --- a/easytier/src/peers/encrypt/xor.rs +++ b/easytier-core/src/tunnel/encrypt/xor.rs @@ -1,10 +1,7 @@ -use crate::tunnel::packet_def::ZCPacket; +use crate::packet::ZCPacket; use super::{Encryptor, Error}; -// XOR 加密不需要额外的尾部数据,因为它是对称的 -pub const XOR_ENCRYPTION_RESERVED: usize = 0; - #[derive(Clone)] pub struct XorCipher { pub(crate) key: Vec, @@ -61,8 +58,8 @@ impl Encryptor for XorCipher { #[cfg(test)] mod tests { use crate::{ - peers::encrypt::{Encryptor, xor::XorCipher}, - tunnel::packet_def::ZCPacket, + packet::ZCPacket, + tunnel::encrypt::{Encryptor, xor::XorCipher}, }; #[test] diff --git a/easytier/src/tunnel/filter.rs b/easytier-core/src/tunnel/filter.rs similarity index 78% rename from easytier/src/tunnel/filter.rs rename to easytier-core/src/tunnel/filter.rs index 787eeab9..4cb00ea7 100644 --- a/easytier/src/tunnel/filter.rs +++ b/easytier-core/src/tunnel/filter.rs @@ -6,11 +6,13 @@ use std::{ use auto_impl::auto_impl; use futures::{Sink, SinkExt, Stream, StreamExt}; -use crate::proto::common::TunnelInfo; - -use self::stats::Throughput; - -use super::*; +use crate::{ + packet::ZCPacket, + proto::common::TunnelInfo, + tunnel::{ + SinkError, SinkItem, StreamItem, Tunnel, ZCPacketSink, ZCPacketStream, stats::Throughput, + }, +}; #[auto_impl(Arc, Box)] pub trait TunnelFilter: Send + Sync { @@ -41,14 +43,17 @@ where B: TunnelFilter, { type FilterOutput = (OA, OB); + fn before_send(&self, data: SinkItem) -> Option { let data = self.a.before_send(data)?; self.b.before_send(data) } + fn after_received(&self, data: StreamItem) -> Option { let data = self.b.after_received(data)?; self.a.after_received(data) } + fn filter_output(&self) -> Self::FilterOutput { (self.a.filter_output(), self.b.filter_output()) } @@ -65,8 +70,10 @@ impl TunnelFilterChain { } pub struct EmptyFilter; + impl TunnelFilter for EmptyFilter { type FilterOutput = (); + fn filter_output(&self) {} } @@ -199,7 +206,12 @@ where self.inner.info() } - fn split(&self) -> (Pin>, Pin>) { + fn split( + &self, + ) -> ( + std::pin::Pin>, + std::pin::Pin>, + ) { let (stream, sink) = self.inner.split(); let filter = self.filter.clone(); ( @@ -299,75 +311,3 @@ impl StatsRecorderTunnelFilter { self.throughput.clone() } } - -#[cfg(test)] -pub mod tests { - use std::sync::atomic::{AtomicU32, Ordering}; - - use filter::ring::create_ring_tunnel_pair; - - use super::*; - - pub struct DropSendTunnelFilter { - start: AtomicU32, - end: AtomicU32, - cur: AtomicU32, - } - - impl TunnelFilter for DropSendTunnelFilter { - type FilterOutput = (); - - fn before_send(&self, data: SinkItem) -> Option { - self.cur.fetch_add(1, Ordering::SeqCst); - if self.cur.load(Ordering::SeqCst) >= self.start.load(Ordering::SeqCst) - && self.cur.load(std::sync::atomic::Ordering::SeqCst) - < self.end.load(Ordering::SeqCst) - { - tracing::trace!("drop packet: {:?}", data); - return None; - } - Some(data) - } - - fn filter_output(&self) {} - } - - impl DropSendTunnelFilter { - pub fn new(start: u32, end: u32) -> Self { - Self { - start: AtomicU32::new(start), - end: AtomicU32::new(end), - cur: AtomicU32::new(0), - } - } - } - - #[tokio::test] - async fn test_nested_filter() { - let filter = Arc::new( - PacketRecorderTunnelFilter::new() - .to_chain() - .chain(PacketRecorderTunnelFilter::new()) - .chain(PacketRecorderTunnelFilter::new()) - .chain(PacketRecorderTunnelFilter::new()), - ); - let (s, _b) = create_ring_tunnel_pair(); - let tunnel = TunnelWithFilter::new(s, filter.clone()); - - let (_r, mut s) = tunnel.split(); - s.send(ZCPacket::new_with_payload("ab".as_bytes())) - .await - .unwrap(); - - let out = filter.filter_output(); - - let a = out.0.0.0.1; - let b = out.0.0.1; - let c = out.0.1; - let _d = out.1; - - assert_eq!(1, a.0.len()); - assert_eq!(1, b.0.len()); - assert_eq!(1, c.0.len()); - } -} diff --git a/easytier-core/src/tunnel/framed.rs b/easytier-core/src/tunnel/framed.rs new file mode 100644 index 00000000..76e5995a --- /dev/null +++ b/easytier-core/src/tunnel/framed.rs @@ -0,0 +1,375 @@ +use std::{ + any::Any, + collections::VecDeque, + io::IoSlice, + pin::Pin, + task::{Poll, ready}, +}; + +use bytes::{Buf, BufMut, Bytes, BytesMut}; +use futures::{Sink, Stream}; +use pin_project_lite::pin_project; +use tokio::io::{AsyncRead, AsyncWrite, ReadBuf}; +use tokio_util::io::poll_write_buf; +use zerocopy::FromBytes as _; + +use crate::{ + packet::{ + PEER_MANAGER_HEADER_SIZE, TCP_TUNNEL_HEADER_SIZE, TCPTunnelHeader, ZCPacket, ZCPacketType, + }, + tunnel::{SinkError, SinkItem, StreamItem, TunnelError}, +}; + +pub const TCP_MTU_BYTES: usize = 2000; + +pub fn reserve_buf(buf: &mut BytesMut, min_size: usize, max_size: usize) { + if buf.capacity() < min_size { + buf.reserve(max_size); + } +} + +pin_project! { + pub struct FramedReader { + #[pin] + reader: R, + buf: BytesMut, + max_packet_size: usize, + _associate_data: Option>, + error: Option, + } +} + +impl FramedReader { + pub fn new(reader: R, max_packet_size: usize) -> Self { + Self::new_with_associate_data(reader, max_packet_size, None) + } + + pub fn new_with_associate_data( + reader: R, + max_packet_size: usize, + associate_data: Option>, + ) -> Self { + Self { + reader, + buf: BytesMut::with_capacity(max_packet_size), + max_packet_size, + _associate_data: associate_data, + error: None, + } + } + + pub fn extract_one_packet( + buf: &mut BytesMut, + max_packet_size: usize, + ) -> Option> { + if buf.len() < TCP_TUNNEL_HEADER_SIZE { + return None; + } + + let header = TCPTunnelHeader::ref_from_prefix(&buf[..]).unwrap(); + let body_len = header.len.get() as usize; + if body_len > max_packet_size { + return Some(Err(TunnelError::InvalidPacket("body too long".to_owned()))); + } + + if body_len < PEER_MANAGER_HEADER_SIZE { + return Some(Err(TunnelError::InvalidPacket("body too short".to_owned()))); + } + + if buf.len() < TCP_TUNNEL_HEADER_SIZE + body_len { + return None; + } + + let packet_buf = buf.split_to(TCP_TUNNEL_HEADER_SIZE + body_len); + Some(Ok(ZCPacket::new_from_buf(packet_buf, ZCPacketType::TCP))) + } +} + +impl Stream for FramedReader +where + R: AsyncRead + Send + 'static + Unpin, +{ + type Item = StreamItem; + + fn poll_next( + self: Pin<&mut Self>, + cx: &mut std::task::Context<'_>, + ) -> Poll> { + let mut this = self.project(); + + loop { + if let Some(error) = this.error.as_ref() { + tracing::warn!("poll_next on a failed FramedReader, {:?}", error); + return Poll::Ready(None); + } + + if let Some(packet) = Self::extract_one_packet(this.buf, *this.max_packet_size) { + if let Err(TunnelError::InvalidPacket(msg)) = packet.as_ref() { + this.error.replace(TunnelError::InvalidPacket(msg.clone())); + } + return Poll::Ready(Some(packet)); + } + + reserve_buf(this.buf, *this.max_packet_size, *this.max_packet_size * 2); + + let cap = this.buf.capacity() - this.buf.len(); + let buf = this.buf.chunk_mut().as_mut_ptr(); + let buf = unsafe { std::slice::from_raw_parts_mut(buf, cap) }; + let mut buf = ReadBuf::new(buf); + + let ret = ready!(this.reader.as_mut().poll_read(cx, &mut buf)); + let len = buf.filled().len(); + unsafe { this.buf.advance_mut(len) }; + + match ret { + Ok(_) if len == 0 => return Poll::Ready(None), + Ok(_) => {} + Err(error) => return Poll::Ready(Some(Err(TunnelError::IOError(error)))), + } + } + } +} + +pub trait ZCPacketToBytes { + fn zcpacket_into_bytes(&self, zc_packet: ZCPacket) -> Result; +} + +pub struct TcpZCPacketToBytes; + +impl ZCPacketToBytes for TcpZCPacketToBytes { + fn zcpacket_into_bytes(&self, item: ZCPacket) -> Result { + let mut item = item.convert_type(ZCPacketType::TCP); + + let tcp_len = PEER_MANAGER_HEADER_SIZE + item.payload_len(); + let Some(header) = item.mut_tcp_tunnel_header() else { + return Err(TunnelError::InvalidPacket("packet too short".to_owned())); + }; + header.len.set(tcp_len.try_into().unwrap()); + + Ok(item.into_bytes()) + } +} + +struct SendBufs { + bufs: VecDeque, +} + +impl SendBufs { + fn new() -> Self { + Self { + bufs: VecDeque::new(), + } + } + + fn len(&self) -> usize { + self.bufs.len() + } + + fn push(&mut self, buf: Bytes) { + debug_assert!(buf.has_remaining()); + self.bufs.push_back(buf); + } +} + +impl Buf for SendBufs { + fn remaining(&self) -> usize { + self.bufs.iter().map(Buf::remaining).sum() + } + + fn chunk(&self) -> &[u8] { + self.bufs.front().map(Buf::chunk).unwrap_or_default() + } + + fn advance(&mut self, mut cnt: usize) { + while cnt > 0 { + let Some(front) = self.bufs.front_mut() else { + return; + }; + let rem = front.remaining(); + if rem > cnt { + front.advance(cnt); + return; + } + front.advance(rem); + cnt -= rem; + self.bufs.pop_front(); + } + } + + fn chunks_vectored<'a>(&'a self, dst: &mut [IoSlice<'a>]) -> usize { + if dst.is_empty() { + return 0; + } + + let mut count = 0; + for buf in &self.bufs { + count += buf.chunks_vectored(&mut dst[count..]); + if count == dst.len() { + break; + } + } + count + } + + fn copy_to_bytes(&mut self, len: usize) -> Bytes { + match self.bufs.front_mut() { + Some(front) if front.remaining() == len => { + let bytes = front.copy_to_bytes(len); + self.bufs.pop_front(); + bytes + } + Some(front) if front.remaining() > len => front.copy_to_bytes(len), + _ => { + assert!(len <= self.remaining(), "len greater than remaining"); + let mut bytes = BytesMut::with_capacity(len); + bytes.put(self.take(len)); + bytes.freeze() + } + } + } +} + +pin_project! { + pub struct FramedWriter { + #[pin] + writer: W, + sending_bufs: SendBufs, + _associate_data: Option>, + converter: C, + } +} + +impl FramedWriter { + fn max_buffer_count(&self) -> usize { + 64 + } +} + +impl FramedWriter { + pub fn new(writer: W) -> Self { + Self::new_with_associate_data(writer, None) + } + + pub fn new_with_associate_data( + writer: W, + associate_data: Option>, + ) -> Self { + Self { + writer, + sending_bufs: SendBufs::new(), + _associate_data: associate_data, + converter: TcpZCPacketToBytes, + } + } +} + +impl FramedWriter { + pub fn new_with_converter(writer: W, converter: C) -> Self { + Self::new_with_converter_and_associate_data(writer, converter, None) + } + + pub fn new_with_converter_and_associate_data( + writer: W, + converter: C, + associate_data: Option>, + ) -> Self { + Self { + writer, + sending_bufs: SendBufs::new(), + _associate_data: associate_data, + converter, + } + } +} + +impl Sink for FramedWriter +where + W: AsyncWrite + Send + 'static, + C: ZCPacketToBytes + Send + 'static, +{ + type Error = SinkError; + + fn poll_ready( + mut self: Pin<&mut Self>, + cx: &mut std::task::Context<'_>, + ) -> Poll> { + let max_buffer_count = self.max_buffer_count(); + if self.sending_bufs.len() >= max_buffer_count { + self.as_mut().poll_flush(cx) + } else { + tracing::trace!(bufs_cnt = self.sending_bufs.len(), "ready to send"); + Poll::Ready(Ok(())) + } + } + + fn start_send(self: Pin<&mut Self>, item: SinkItem) -> Result<(), Self::Error> { + let this = self.project(); + this.sending_bufs + .push(this.converter.zcpacket_into_bytes(item)?); + Ok(()) + } + + fn poll_flush( + self: Pin<&mut Self>, + cx: &mut std::task::Context<'_>, + ) -> Poll> { + let mut this = self.project(); + let mut remaining = this.sending_bufs.remaining(); + while remaining != 0 { + let n = ready!(poll_write_buf(this.writer.as_mut(), cx, this.sending_bufs))?; + if n == 0 { + return Poll::Ready(Err(TunnelError::IOError(std::io::Error::new( + std::io::ErrorKind::WriteZero, + "failed to write frame to transport", + )))); + } + remaining -= n; + } + + tracing::trace!(?remaining, "flushed"); + ready!(this.writer.poll_flush(cx))?; + Poll::Ready(Ok(())) + } + + fn poll_close( + mut self: Pin<&mut Self>, + cx: &mut std::task::Context<'_>, + ) -> Poll> { + ready!(self.as_mut().poll_flush(cx))?; + ready!(self.project().writer.poll_shutdown(cx))?; + Poll::Ready(Ok(())) + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn framed_reader_rejects_short_peer_manager_body() { + let mut buf = BytesMut::new(); + buf.put_u32_le((PEER_MANAGER_HEADER_SIZE - 1) as u32); + buf.resize(TCP_TUNNEL_HEADER_SIZE + PEER_MANAGER_HEADER_SIZE - 1, 0); + + let ret = FramedReader::::extract_one_packet(&mut buf, 2000); + + assert!(matches!( + ret, + Some(Err(TunnelError::InvalidPacket(msg))) if msg == "body too short" + )); + } + + #[test] + fn send_bufs_exposes_all_queued_buffers_for_vectored_write() { + let mut bufs = SendBufs::new(); + bufs.push(Bytes::from_static(b"abc")); + bufs.push(Bytes::from_static(b"defg")); + + let mut slices = [IoSlice::new(&[]); 4]; + let count = bufs.chunks_vectored(&mut slices); + + assert_eq!(count, 2); + assert_eq!(&*slices[0], b"abc"); + assert_eq!(&*slices[1], b"defg"); + } +} diff --git a/easytier-core/src/tunnel/mod.rs b/easytier-core/src/tunnel/mod.rs new file mode 100644 index 00000000..7ba9878d --- /dev/null +++ b/easytier-core/src/tunnel/mod.rs @@ -0,0 +1,98 @@ +//! Shared tunnel contract and per-transport tunnel implementations. +//! +//! This root file declares the shared tunnel contract: [`TunnelError`], the +//! [`Tunnel`] trait, and the [`ZCPacket`] stream/sink aliases that tunnels +//! carry. The sibling files in this directory hold the per-transport tunnel +//! implementations (TCP, UDP, ring, mpsc), the filters and wrappers layered +//! on top of them, and the packet encryption primitives (`encrypt`, +//! `secure_datagram`) shared by tunnels and peer sessions. + +use std::{fmt::Debug, pin::Pin}; + +use futures::{Sink, Stream}; + +use crate::{foundation::time::error::Elapsed, packet::ZCPacket, proto::common::TunnelInfo}; + +pub use crate::socket::IpVersion; + +pub(crate) mod encrypt; +pub mod filter; +pub mod framed; +pub mod mpsc; +pub mod ring; +pub(crate) mod secure_datagram; +pub mod stats; +pub mod tcp; +pub mod udp; +pub mod web_security; +pub mod wrapper; + +/// Reports whether the configured name resolves to the insecure XOR cipher. +pub fn effective_encryption_uses_xor(algorithm: &str) -> bool { + encrypt::effective_algorithm_uses_xor(algorithm) +} + +#[derive(Debug, thiserror::Error)] +pub enum TunnelError { + #[error("io error: {0}")] + IOError(#[from] std::io::Error), + #[error("invalid packet. msg: {0}")] + InvalidPacket(String), + #[error("exceed max packet size. max: {0}, input: {1}")] + ExceedMaxPacketSize(usize, usize), + #[error("invalid protocol: {0}")] + InvalidProtocol(String), + #[error("invalid addr: {0}")] + InvalidAddr(String), + #[error("internal error {0}")] + InternalError(String), + #[error("conn id not match, expect: {0}, actual: {1}")] + ConnIdNotMatch(u32, u32), + #[error("buffer full")] + BufferFull, + #[error("timeout")] + Timeout(#[from] Elapsed), + #[error("anyhow error: {0}")] + Anyhow(#[from] anyhow::Error), + #[error("shutdown")] + Shutdown, + #[error("no dns record found")] + NoDnsRecordFound(IpVersion), + #[error("{0}")] + ProtocolError(String), + #[error("tunnel error: {0}")] + TunError(String), +} + +impl From for crate::proto::rpc_types::error::Error { + fn from(value: TunnelError) -> Self { + Self::TunnelError(value.to_string()) + } +} + +pub type StreamT = ZCPacket; +pub type StreamItem = Result; +pub type SinkItem = ZCPacket; +pub type SinkError = TunnelError; + +pub trait ZCPacketStream: Stream + Send {} +impl ZCPacketStream for T where T: Stream + Send {} + +pub trait ZCPacketSink: Sink + Send {} +impl ZCPacketSink for T where T: Sink + Send {} + +pub type SplitTunnel = (Pin>, Pin>); + +#[auto_impl::auto_impl(Box, Arc)] +pub trait Tunnel: Send { + fn split(&self) -> SplitTunnel; + fn info(&self) -> Option; +} + +impl std::fmt::Debug for dyn Tunnel { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("Tunnel") + .field("info", &self.info()) + .finish() + } +} diff --git a/easytier-core/src/tunnel/mpsc.rs b/easytier-core/src/tunnel/mpsc.rs new file mode 100644 index 00000000..18d8ab3d --- /dev/null +++ b/easytier-core/src/tunnel/mpsc.rs @@ -0,0 +1,139 @@ +use std::{pin::Pin, time::Duration}; + +use futures::SinkExt; +use tokio::{ + sync::mpsc::{Receiver, Sender, channel, error::TrySendError}, + task::JoinHandle, +}; + +use crate::{ + foundation::time::timeout, + packet::ZCPacket, + proto::common::TunnelInfo, + tunnel::{Tunnel, TunnelError, ZCPacketSink, ZCPacketStream}, +}; + +#[derive(Clone)] +pub struct MpscTunnelSender(Sender); + +impl MpscTunnelSender { + pub async fn send(&self, item: ZCPacket) -> Result<(), TunnelError> { + self.0.send(item).await.map_err(|_| TunnelError::Shutdown) + } + + pub fn try_send(&self, item: ZCPacket) -> Result<(), TunnelError> { + self.0.try_send(item).map_err(|e| match e { + TrySendError::Full(_) => TunnelError::BufferFull, + TrySendError::Closed(_) => TunnelError::Shutdown, + }) + } +} + +pub struct MpscTunnel { + tx: Option>, + tunnel: T, + stream: Option>>, + task: JoinHandle<()>, +} + +impl MpscTunnel { + pub fn new(tunnel: T, send_timeout: Option) -> Self { + let (tx, mut rx) = channel(32); + let (stream, mut sink) = tunnel.split(); + + let task = tokio::spawn(async move { + loop { + if let Err(e) = Self::forward_one_round(&mut rx, &mut sink, send_timeout).await { + tracing::error!(?e, "forward error"); + break; + } + } + rx.close(); + let close_ret = timeout(Duration::from_secs(5), sink.close()).await; + tracing::warn!(?close_ret, "mpsc close sink"); + }); + + Self { + tx: Some(tx), + tunnel, + stream: Some(stream), + task, + } + } + + async fn forward_one_round( + rx: &mut Receiver, + sink: &mut Pin>, + send_timeout_ms: Option, + ) -> Result<(), TunnelError> { + let item = rx.recv().await.ok_or(TunnelError::Shutdown)?; + if let Some(timeout_ms) = send_timeout_ms { + Self::forward_one_round_with_timeout(rx, sink, item, timeout_ms).await + } else { + Self::forward_one_round_no_timeout(rx, sink, item).await + } + } + + async fn forward_one_round_no_timeout( + rx: &mut Receiver, + sink: &mut Pin>, + initial_item: ZCPacket, + ) -> Result<(), TunnelError> { + sink.feed(initial_item).await?; + + while let Ok(item) = rx.try_recv() { + if let Err(e) = sink.feed(item).await { + tracing::error!(?e, "feed error"); + return Err(e); + } + } + + sink.flush().await + } + + async fn forward_one_round_with_timeout( + rx: &mut Receiver, + sink: &mut Pin>, + initial_item: ZCPacket, + timeout_ms: Duration, + ) -> Result<(), TunnelError> { + match timeout(timeout_ms, async move { + Self::forward_one_round_no_timeout(rx, sink, initial_item).await + }) + .await + { + Ok(Ok(_)) => Ok(()), + Ok(Err(e)) => { + tracing::error!(?e, "forward error"); + Err(e) + } + Err(e) => { + tracing::error!(?e, "forward timeout"); + Err(e.into()) + } + } + } + + pub fn get_stream(&mut self) -> Pin> { + self.stream.take().unwrap() + } + + pub fn get_sink(&self) -> MpscTunnelSender { + MpscTunnelSender(self.tx.as_ref().unwrap().clone()) + } + + pub fn close(&mut self) { + self.tx.take(); + self.task.abort(); + } + + pub fn tunnel_info(&self) -> Option { + self.tunnel.info() + } +} + +impl Drop for MpscTunnel { + fn drop(&mut self) { + self.task.abort(); + } +} diff --git a/easytier-core/src/tunnel/ring.rs b/easytier-core/src/tunnel/ring.rs new file mode 100644 index 00000000..51086777 --- /dev/null +++ b/easytier-core/src/tunnel/ring.rs @@ -0,0 +1,591 @@ +use std::{ + collections::HashMap, + fmt::Debug, + io, + pin::Pin, + sync::{Arc, Mutex}, + task::{Context, Poll, ready}, +}; + +use async_trait::async_trait; +use futures::{Sink, Stream}; +use tokio::io::{AsyncRead, AsyncWrite, ReadBuf}; +use tokio::sync::mpsc::{UnboundedReceiver, UnboundedSender, unbounded_channel}; + +use crate::{ + proto::common::TunnelInfo, + socket::{ + SocketListener, + ring::{ + RING_SOCKET_CAPACITY, RingSocket, RingSocketId, RingSocketReceiver, + RingSocketSendError, RingSocketSender, + }, + }, + tunnel::{SinkError, SinkItem, StreamItem, Tunnel, TunnelError, ZCPacketSink, ZCPacketStream}, +}; + +pub const RING_TUNNEL_CAP: usize = RING_SOCKET_CAPACITY; + +type RingItem = SinkItem; +pub type RingTunnelSocket = RingSocket; +pub type RingSinkSendError = RingSocketSendError; + +pub struct RingByteStream { + receiver: RingSocketReceiver, + sender: RingSocketSender, + buffered: Option<(RingItem, usize)>, +} + +impl RingByteStream { + pub fn new(socket: Arc) -> io::Result { + let (receiver, sender) = socket + .try_split() + .map_err(|error| io::Error::other(error.to_string()))?; + Ok(Self { + receiver, + sender, + buffered: None, + }) + } +} + +impl AsyncRead for RingByteStream { + fn poll_read( + mut self: Pin<&mut Self>, + cx: &mut Context<'_>, + buffer: &mut ReadBuf<'_>, + ) -> Poll> { + if buffer.remaining() == 0 { + return Poll::Ready(Ok(())); + } + loop { + if let Some((packet, offset)) = &mut self.buffered { + if *offset == packet.payload().len() { + self.buffered = None; + continue; + } + let remaining = &packet.payload()[*offset..]; + let copy_len = remaining.len().min(buffer.remaining()); + buffer.put_slice(&remaining[..copy_len]); + *offset += copy_len; + return Poll::Ready(Ok(())); + } + + match ready!(Pin::new(&mut self.receiver).poll_next(cx)) { + Some(Ok(packet)) => self.buffered = Some((packet, 0)), + Some(Err(error)) => return Poll::Ready(Err(io::Error::other(error.to_string()))), + None => return Poll::Ready(Ok(())), + } + } + } +} + +impl AsyncWrite for RingByteStream { + fn poll_write( + mut self: Pin<&mut Self>, + cx: &mut Context<'_>, + buffer: &[u8], + ) -> Poll> { + if buffer.is_empty() { + return Poll::Ready(Ok(0)); + } + ready!(Pin::new(&mut self.sender).poll_ready(cx)) + .map_err(|error| io::Error::other(error.to_string()))?; + Pin::new(&mut self.sender) + .start_send(crate::packet::ZCPacket::new_with_payload(buffer)) + .map_err(|error| io::Error::other(error.to_string()))?; + Poll::Ready(Ok(buffer.len())) + } + + fn poll_flush(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + Pin::new(&mut self.sender) + .poll_flush(cx) + .map_err(|error| io::Error::other(error.to_string())) + } + + fn poll_shutdown(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + Pin::new(&mut self.sender) + .poll_close(cx) + .map_err(|error| io::Error::other(error.to_string())) + } +} + +#[derive(Debug, thiserror::Error)] +pub enum RingTunnelRegistryError { + #[error("ring listener already registered: {0}")] + AlreadyRegistered(RingSocketId), + #[error("ring listener not found: {0}")] + NotFound(RingSocketId), + #[error("ring listener closed: {0}")] + Closed(RingSocketId), + #[error("ring listener {listener_id} received socket for {socket_id}")] + SocketIdMismatch { + listener_id: RingSocketId, + socket_id: RingSocketId, + }, +} + +struct PendingRingConnection { + client: Arc, + server: Arc, +} + +type ConnectionMap = HashMap>>; + +#[derive(Default)] +pub(crate) struct RingTunnelRegistry { + connections: Mutex, +} + +impl RingTunnelRegistry { + pub(crate) fn bind( + self: &Arc, + local_id: RingSocketId, + ) -> Result { + let (conn_sender, conn_receiver) = unbounded_channel(); + let mut connections = self.connections.lock().unwrap(); + if connections.contains_key(&local_id) { + return Err(RingTunnelRegistryError::AlreadyRegistered(local_id)); + } + + connections.insert(local_id, conn_sender.clone()); + Ok(RingTunnelSocketListener { + registry: self.clone(), + local_id, + conn_sender, + conn_receiver, + }) + } + + pub(crate) fn connect( + &self, + remote_id: RingSocketId, + ) -> Result { + let conn_sender = self + .connections + .lock() + .unwrap() + .get(&remote_id) + .cloned() + .ok_or(RingTunnelRegistryError::NotFound(remote_id))?; + let (client, server) = + RingSocket::pair_with_ids(uuid::Uuid::new_v4(), remote_id, RING_TUNNEL_CAP); + let conn = Arc::new(PendingRingConnection { + client: client.clone(), + server, + }); + + conn_sender + .send(conn) + .map_err(|_| RingTunnelRegistryError::Closed(remote_id))?; + + Ok(DialedRingSocket { + local_id: client.id(), + socket: client, + remote_id, + }) + } +} + +pub struct AcceptedRingSocket { + pub socket: Arc, + pub local_id: RingSocketId, + pub remote_id: RingSocketId, +} + +impl AcceptedRingSocket { + pub fn into_tunnel(self) -> Box { + Box::new(RingTunnel::new( + self.socket, + Some(ring_tunnel_info(self.local_id, self.remote_id)), + )) + } +} + +pub struct DialedRingSocket { + pub socket: Arc, + pub local_id: RingSocketId, + pub remote_id: RingSocketId, +} + +impl DialedRingSocket { + pub fn into_tunnel(self) -> Box { + Box::new(RingTunnel::new( + self.socket, + Some(ring_tunnel_info(self.local_id, self.remote_id)), + )) + } +} + +pub struct RingTunnelSocketListener { + registry: Arc, + local_id: RingSocketId, + conn_sender: UnboundedSender>, + conn_receiver: UnboundedReceiver>, +} + +impl Debug for RingTunnelSocketListener { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("RingTunnelSocketListener") + .field("local_id", &self.local_id) + .finish_non_exhaustive() + } +} + +impl RingTunnelSocketListener { + pub fn local_id(&self) -> RingSocketId { + self.local_id + } + + pub async fn accept(&mut self) -> Result { + let conn = self + .conn_receiver + .recv() + .await + .ok_or(RingTunnelRegistryError::Closed(self.local_id))?; + + let socket_id = conn.server.id(); + if socket_id != self.local_id { + return Err(RingTunnelRegistryError::SocketIdMismatch { + listener_id: self.local_id, + socket_id, + }); + } + + Ok(AcceptedRingSocket { + socket: conn.server.clone(), + local_id: self.local_id, + remote_id: conn.client.id(), + }) + } +} + +#[async_trait] +impl SocketListener for RingTunnelSocketListener { + type Accepted = Box; + + async fn listen(&mut self) -> anyhow::Result<()> { + Ok(()) + } + + async fn accept(&mut self) -> anyhow::Result { + Ok(RingTunnelSocketListener::accept(self).await?.into_tunnel()) + } + + fn local_url(&self) -> url::Url { + ring_url(self.local_id) + } +} + +impl Drop for RingTunnelSocketListener { + fn drop(&mut self) { + let mut connections = self.registry.connections.lock().unwrap(); + if connections + .get(&self.local_id) + .is_some_and(|sender| sender.same_channel(&self.conn_sender)) + { + connections.remove(&self.local_id); + } + } +} + +pub struct RingStream { + id: RingSocketId, + inner: RingSocketReceiver, +} + +impl RingStream { + pub fn new(inner: RingSocketReceiver, id: RingSocketId) -> Self { + Self { id, inner } + } +} + +impl Stream for RingStream { + type Item = StreamItem; + + fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + let ret = std::task::ready!(Pin::new(&mut self.get_mut().inner).poll_next(cx)); + Poll::Ready(ret.map(|item| item.map_err(|_| TunnelError::Shutdown))) + } +} + +impl Debug for RingStream { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("RingStream") + .field("id", &self.id) + .finish_non_exhaustive() + } +} + +pub struct RingSink { + id: RingSocketId, + inner: RingSocketSender, +} + +impl RingSink { + pub fn new(inner: RingSocketSender, id: RingSocketId) -> Self { + Self { id, inner } + } + + pub fn try_send(&mut self, item: RingItem) -> Result<(), RingSinkSendError> { + self.inner.try_send(item) + } + + pub fn force_send(&mut self, item: RingItem) -> Result<(), RingSinkSendError> { + self.inner.force_send(item) + } +} + +fn map_ring_send_error(error: RingSinkSendError) -> TunnelError { + match error { + RingSocketSendError::Closed(_) => TunnelError::Shutdown, + RingSocketSendError::Full(_) => TunnelError::BufferFull, + } +} + +impl Sink for RingSink { + type Error = SinkError; + + fn poll_ready(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + Pin::new(&mut self.get_mut().inner) + .poll_ready(cx) + .map_err(|_| TunnelError::Shutdown) + } + + fn start_send(self: Pin<&mut Self>, item: SinkItem) -> Result<(), Self::Error> { + self.get_mut() + .inner + .force_send(item) + .map_err(map_ring_send_error) + } + + fn poll_flush(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + Pin::new(&mut self.get_mut().inner) + .poll_flush(cx) + .map_err(|_| TunnelError::Shutdown) + } + + fn poll_close(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + Pin::new(&mut self.get_mut().inner) + .poll_close(cx) + .map_err(|_| TunnelError::Shutdown) + } +} + +impl Debug for RingSink { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("RingSink") + .field("id", &self.id) + .finish_non_exhaustive() + } +} + +pub fn create_ring_socket_pair(capacity: usize) -> (Arc, Arc) { + RingSocket::pair(capacity) +} + +pub fn split_ring_socket(socket: Arc) -> (RingStream, RingSink) { + let id = socket.id(); + let (recv, send) = socket.split(); + (RingStream::new(recv, id), RingSink::new(send, id)) +} + +pub struct RingTunnel { + socket: Arc, + info: Option, +} + +impl RingTunnel { + pub fn new(socket: Arc, info: Option) -> Self { + Self { socket, info } + } + + pub fn socket(&self) -> Arc { + self.socket.clone() + } +} + +impl Debug for RingTunnel { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("RingTunnel") + .field("socket", &self.socket) + .field("info", &self.info) + .finish() + } +} + +impl Tunnel for RingTunnel { + fn split(&self) -> (Pin>, Pin>) { + let (stream, sink) = split_ring_socket(self.socket.clone()); + (Box::pin(stream), Box::pin(sink)) + } + + fn info(&self) -> Option { + self.info.clone() + } +} + +fn ring_tunnel_info(local_id: RingSocketId, remote_id: RingSocketId) -> TunnelInfo { + TunnelInfo { + tunnel_type: "ring".to_owned(), + local_addr: Some(ring_url(local_id).into()), + remote_addr: Some(ring_url(remote_id).into()), + resolved_remote_addr: Some(ring_url(remote_id).into()), + } +} + +fn ring_url(id: RingSocketId) -> url::Url { + format!("ring://{id}") + .parse() + .expect("ring socket id should form a valid URL") +} + +pub fn create_ring_tunnel_pair() -> (Box, Box) { + let (first, second) = create_ring_socket_pair(RING_TUNNEL_CAP); + ( + Box::new(RingTunnel::new(first, None)), + Box::new(RingTunnel::new(second, None)), + ) +} + +#[cfg(test)] +mod tests { + use futures::{SinkExt, StreamExt}; + use tokio::{ + io::{AsyncReadExt, AsyncWriteExt}, + time::{Duration, timeout}, + }; + + use crate::packet::ZCPacket; + + use super::*; + + #[tokio::test] + async fn explicit_registries_isolate_ring_listener_namespaces() { + let listener_id = uuid::Uuid::new_v4(); + let first = Arc::new(RingTunnelRegistry::default()); + let second = Arc::new(RingTunnelRegistry::default()); + let mut listener = first.bind(listener_id).unwrap(); + + assert!(matches!( + second.connect(listener_id), + Err(RingTunnelRegistryError::NotFound(id)) if id == listener_id + )); + + let dialed = first.connect(listener_id).unwrap(); + let accepted = listener.accept().await.unwrap(); + assert_eq!(dialed.remote_id, listener_id); + assert_eq!(accepted.local_id, listener_id); + assert_eq!(accepted.remote_id, dialed.local_id); + + drop(listener); + assert!(matches!( + first.connect(listener_id), + Err(RingTunnelRegistryError::NotFound(id)) if id == listener_id + )); + } + + #[tokio::test] + async fn registry_endpoints_upgrade_to_packet_native_tunnels() { + let listener_id = uuid::Uuid::new_v4(); + let registry = Arc::new(RingTunnelRegistry::default()); + let mut listener = registry.bind(listener_id).unwrap(); + let client = registry.connect(listener_id).unwrap().into_tunnel(); + let server = listener.accept().await.unwrap().into_tunnel(); + let (_client_stream, mut client_sink) = client.split(); + let (mut server_stream, _server_sink) = server.split(); + + client_sink + .send(ZCPacket::new_with_payload(b"packet-native")) + .await + .unwrap(); + let packet = server_stream.next().await.unwrap().unwrap(); + assert_eq!(packet.payload(), b"packet-native"); + assert_eq!(client.info().unwrap().tunnel_type, "ring"); + } + + #[tokio::test] + async fn ring_tunnel_pair_transfers_packets() { + let (left, right) = create_ring_tunnel_pair(); + let (_left_stream, mut left_sink) = left.split(); + let (mut right_stream, _right_sink) = right.split(); + + let packet = ZCPacket::new_with_payload(&[1, 2, 3]); + left_sink.send(packet.clone()).await.unwrap(); + + let received = right_stream.next().await.unwrap().unwrap(); + assert_eq!(received.payload(), packet.payload()); + } + + #[tokio::test] + async fn ring_tunnel_observes_remote_close() { + let (left, right) = create_ring_tunnel_pair(); + drop(left); + + let (mut right_stream, _right_sink) = right.split(); + assert!(right_stream.next().await.is_none()); + } + + #[tokio::test] + async fn ring_stream_can_be_aborted_while_waiting() { + let (_left, right) = create_ring_tunnel_pair(); + let (mut right_stream, _right_sink) = right.split(); + let task = tokio::spawn(async move { right_stream.next().await }); + + tokio::task::yield_now().await; + task.abort(); + assert!(task.await.unwrap_err().is_cancelled()); + } + + #[tokio::test] + async fn ring_stream_wait_can_time_out() { + let (_left, right) = create_ring_tunnel_pair(); + let (mut right_stream, _right_sink) = right.split(); + + assert!( + timeout(Duration::from_millis(10), right_stream.next()) + .await + .is_err() + ); + } + + #[tokio::test] + async fn byte_stream_preserves_bytes_across_partial_reads() { + let (left, right) = RingTunnelSocket::pair(8); + let mut left = RingByteStream::new(left).unwrap(); + let mut right = RingByteStream::new(right).unwrap(); + + timeout(Duration::from_millis(50), right.read(&mut [])) + .await + .expect("zero-capacity read should not wait") + .unwrap(); + left.write_all(b"abcdef").await.unwrap(); + + let mut prefix = [0; 2]; + right.read_exact(&mut prefix).await.unwrap(); + let mut suffix = [0; 4]; + right.read_exact(&mut suffix).await.unwrap(); + assert_eq!(&prefix, b"ab"); + assert_eq!(&suffix, b"cdef"); + } + + #[tokio::test] + async fn byte_stream_skips_empty_packets_without_reporting_eof() { + let (left, right) = RingTunnelSocket::pair(8); + let (_left_receiver, mut left_sender) = left.split(); + let mut right = RingByteStream::new(right).unwrap(); + + left_sender + .send(crate::packet::ZCPacket::new_with_payload(b"")) + .await + .unwrap(); + left_sender + .send(crate::packet::ZCPacket::new_with_payload(b"data")) + .await + .unwrap(); + + let mut received = [0; 4]; + right.read_exact(&mut received).await.unwrap(); + assert_eq!(&received, b"data"); + } +} diff --git a/easytier/src/peers/secure_datagram.rs b/easytier-core/src/tunnel/secure_datagram.rs similarity index 97% rename from easytier/src/peers/secure_datagram.rs rename to easytier-core/src/tunnel/secure_datagram.rs index f4722e44..0d57e3b7 100644 --- a/easytier/src/peers/secure_datagram.rs +++ b/easytier-core/src/tunnel/secure_datagram.rs @@ -14,8 +14,8 @@ use sha2::Sha256; use zerocopy::FromBytes; use crate::{ - peers::encrypt::{Encryptor, create_encryptor}, - tunnel::packet_def::{StandardAeadTail, ZCPacket}, + packet::{StandardAeadTail, ZCPacket}, + tunnel::encrypt::{Encryptor, create_encryptor}, }; type HmacSha256 = Hmac; @@ -687,20 +687,6 @@ impl SecureDatagramSession { accepted } - fn check_replay( - &self, - epoch: u32, - seq: u64, - dir: SecureDatagramDirection, - now_ms: u64, - ) -> bool { - if self.precheck_replay(epoch, seq, dir, now_ms) { - return self.commit_replay(epoch, seq, dir, now_ms); - } - - false - } - pub fn encrypt_payload( &self, dir: SecureDatagramDirection, @@ -757,17 +743,6 @@ impl SecureDatagramSession { Ok(()) } - - #[cfg(test)] - fn check_replay_for_test( - &self, - epoch: u32, - seq: u64, - dir: SecureDatagramDirection, - now_ms: u64, - ) -> bool { - self.check_replay(epoch, seq, dir, now_ms) - } } fn now_ms() -> u64 { @@ -780,9 +755,27 @@ fn now_ms() -> u64 { #[cfg(test)] mod tests { use super::*; - use crate::tunnel::packet_def::PacketType; + #[cfg(all(feature = "aes-gcm", feature = "chacha20"))] + use crate::packet::PacketType; + + impl SecureDatagramSession { + fn check_replay_for_test( + &self, + epoch: u32, + seq: u64, + dir: SecureDatagramDirection, + now_ms: u64, + ) -> bool { + if self.precheck_replay(epoch, seq, dir, now_ms) { + return self.commit_replay(epoch, seq, dir, now_ms); + } + + false + } + } #[test] + #[cfg(all(feature = "aes-gcm", feature = "chacha20"))] fn secure_datagram_supports_asymmetric_algorithms() { let root_key = SecureDatagramSession::new_root_key(); let generation = 1u32; @@ -844,7 +837,10 @@ mod tests { } #[test] + #[cfg(feature = "aes-gcm")] fn failed_decrypt_does_not_poison_replay_window() { + use crate::packet::PacketType; + let root_key = SecureDatagramSession::new_root_key(); let sender = SecureDatagramSession::new( root_key, diff --git a/easytier/src/tunnel/stats.rs b/easytier-core/src/tunnel/stats.rs similarity index 98% rename from easytier/src/tunnel/stats.rs rename to easytier-core/src/tunnel/stats.rs index f3b707a5..4149acdc 100644 --- a/easytier/src/tunnel/stats.rs +++ b/easytier-core/src/tunnel/stats.rs @@ -80,7 +80,6 @@ impl Clone for Throughput { } } -// add sync::Send and sync::Sync traits to Throughput unsafe impl Send for Throughput {} unsafe impl Sync for Throughput {} diff --git a/easytier-core/src/tunnel/tcp.rs b/easytier-core/src/tunnel/tcp.rs new file mode 100644 index 00000000..d5a9aaca --- /dev/null +++ b/easytier-core/src/tunnel/tcp.rs @@ -0,0 +1,198 @@ +use std::sync::Mutex as StdMutex; + +use crate::{ + proto::common::TunnelInfo, + socket::tcp::VirtualTcpSocket, + tunnel::framed::{FramedReader, FramedWriter, TCP_MTU_BYTES}, + tunnel::{SplitTunnel, Tunnel, TunnelError}, +}; + +pub struct TcpTunnel { + info: Option, + socket: StdMutex>, + max_packet_size: usize, +} + +impl TcpTunnel { + fn new(socket: S, tunnel_info: TunnelInfo, max_packet_size: usize) -> Self { + Self { + info: Some(tunnel_info), + socket: StdMutex::new(Some(socket)), + max_packet_size, + } + } +} + +impl Tunnel for TcpTunnel +where + S: VirtualTcpSocket, +{ + fn split(&self) -> SplitTunnel { + let socket = self + .socket + .lock() + .unwrap() + .take() + .expect("TcpTunnel can only be split once"); + let (reader, writer) = tokio::io::split(socket); + ( + Box::pin(FramedReader::new(reader, self.max_packet_size)), + Box::pin(FramedWriter::new(writer)), + ) + } + + fn info(&self) -> Option { + self.info.clone() + } +} + +pub struct TcpTunnelUpgrader { + tunnel_info: TunnelInfo, + max_packet_size: usize, +} + +impl TcpTunnelUpgrader { + pub fn new(tunnel_info: TunnelInfo) -> Self { + Self { + tunnel_info, + max_packet_size: TCP_MTU_BYTES, + } + } + + pub(crate) fn with_max_packet_size(mut self, max_packet_size: usize) -> Self { + self.max_packet_size = max_packet_size; + self + } + + pub fn upgrade(self, socket: S) -> Result, TunnelError> + where + S: VirtualTcpSocket, + { + Ok(Box::new(TcpTunnel::new( + socket, + self.tunnel_info, + self.max_packet_size, + ))) + } +} + +#[cfg(test)] +mod tests { + use std::{ + io, + net::SocketAddr, + pin::Pin, + task::{Context, Poll}, + }; + + use crate::packet::{PEER_MANAGER_HEADER_SIZE, ZCPacket, ZCPacketType}; + use futures::{SinkExt, StreamExt}; + use tokio::io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt, DuplexStream, ReadBuf}; + + use super::*; + + struct MockTcpSocket { + stream: DuplexStream, + local_addr: SocketAddr, + peer_addr: SocketAddr, + } + + impl MockTcpSocket { + fn new(stream: DuplexStream, local_addr: SocketAddr, peer_addr: SocketAddr) -> Self { + Self { + stream, + local_addr, + peer_addr, + } + } + } + + impl AsyncRead for MockTcpSocket { + fn poll_read( + mut self: Pin<&mut Self>, + cx: &mut Context<'_>, + buf: &mut ReadBuf<'_>, + ) -> Poll> { + Pin::new(&mut self.stream).poll_read(cx, buf) + } + } + + impl AsyncWrite for MockTcpSocket { + fn poll_write( + mut self: Pin<&mut Self>, + cx: &mut Context<'_>, + buf: &[u8], + ) -> Poll> { + Pin::new(&mut self.stream).poll_write(cx, buf) + } + + fn poll_flush(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + Pin::new(&mut self.stream).poll_flush(cx) + } + + fn poll_shutdown(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + Pin::new(&mut self.stream).poll_shutdown(cx) + } + } + + impl VirtualTcpSocket for MockTcpSocket { + fn local_addr(&self) -> io::Result { + Ok(self.local_addr) + } + + fn peer_addr(&self) -> io::Result { + Ok(self.peer_addr) + } + } + + fn set_tcp_tunnel_len(packet: &mut ZCPacket) { + let tcp_len = PEER_MANAGER_HEADER_SIZE + packet.payload_len(); + packet + .mut_tcp_tunnel_header() + .unwrap() + .len + .set(tcp_len.try_into().unwrap()); + } + + #[tokio::test] + async fn tcp_tunnel_upgrader_preserves_metadata_and_framing() { + let (socket_stream, mut peer_stream) = tokio::io::duplex(65536); + let socket = MockTcpSocket::new( + socket_stream, + "127.0.0.1:1000".parse().unwrap(), + "127.0.0.1:2000".parse().unwrap(), + ); + let info = TunnelInfo { + tunnel_type: "tcp".to_owned(), + local_addr: None, + remote_addr: None, + resolved_remote_addr: None, + }; + let tunnel = TcpTunnelUpgrader::new(info.clone()) + .upgrade(socket) + .unwrap(); + assert_eq!(tunnel.info(), Some(info)); + + let (mut stream, mut sink) = tunnel.split(); + let outbound = ZCPacket::new_with_payload(b"outbound"); + let mut expected = outbound.clone().convert_type(ZCPacketType::TCP); + set_tcp_tunnel_len(&mut expected); + let expected_raw = expected.into_bytes(); + let read_peer = tokio::spawn(async move { + let mut raw = vec![0; expected_raw.len()]; + peer_stream.read_exact(&mut raw).await.unwrap(); + assert_eq!(raw, expected_raw); + + let mut inbound = + ZCPacket::new_with_payload(b"inbound").convert_type(ZCPacketType::TCP); + set_tcp_tunnel_len(&mut inbound); + peer_stream.write_all(&inbound.into_bytes()).await.unwrap(); + }); + sink.send(outbound).await.unwrap(); + sink.flush().await.unwrap(); + + let packet = stream.next().await.unwrap().unwrap(); + assert_eq!(packet.payload(), b"inbound"); + read_peer.await.unwrap(); + } +} diff --git a/easytier-core/src/tunnel/udp.rs b/easytier-core/src/tunnel/udp.rs new file mode 100644 index 00000000..2cc98049 --- /dev/null +++ b/easytier-core/src/tunnel/udp.rs @@ -0,0 +1,422 @@ +use std::{ + pin::Pin, + sync::{Arc, Mutex as StdMutex}, + task::{Context, Poll}, +}; + +use bytes::BytesMut; +use futures::{Sink, Stream}; +use tokio::sync::{oneshot, watch}; + +use crate::{ + packet::{UDP_TUNNEL_HEADER_SIZE, UdpPacketType, ZCPacket, ZCPacketType}, + proto::common::TunnelInfo, + socket::{ + ring::{RingSocketError, RingSocketReceiver, RingSocketSendError, RingSocketSender}, + udp::{ + UdpSession, UdpSessionCleanup, UdpSessionCodec, UdpSessionDatagram, UdpSessionOutbound, + UdpSessionTunnelParts, + }, + }, + tunnel::{SinkError, SinkItem, SplitTunnel, StreamItem, Tunnel, TunnelError}, +}; + +fn zcpacket_from_udp_session_payload(payload: &[u8]) -> Result { + let payload_len = u16::try_from(payload.len()) + .map_err(|_| TunnelError::ExceedMaxPacketSize(u16::MAX as usize, payload.len()))?; + let mut buf = BytesMut::new(); + buf.resize(UDP_TUNNEL_HEADER_SIZE + payload.len(), 0); + buf[UDP_TUNNEL_HEADER_SIZE..].copy_from_slice(payload); + + let mut packet = ZCPacket::new_from_buf(buf, ZCPacketType::UDP); + let header = packet.mut_udp_tunnel_header().unwrap(); + header.msg_type = UdpPacketType::Data as u8; + header.len.set(payload_len); + Ok(packet) +} + +fn ring_socket_error_to_tunnel(error: RingSocketError) -> TunnelError { + match error { + RingSocketError::Closed => TunnelError::Shutdown, + RingSocketError::Full => TunnelError::BufferFull, + RingSocketError::AlreadySplit => { + TunnelError::InternalError("udp session ring already split".to_owned()) + } + } +} + +fn ring_send_error_to_tunnel(error: RingSocketSendError) -> TunnelError { + match error { + RingSocketSendError::Closed(_) => TunnelError::Shutdown, + RingSocketSendError::Full(_) => TunnelError::BufferFull, + } +} + +struct UdpTunnelSessionGuard { + cleanup: StdMutex>, + _layer_guard: Option>, +} + +impl UdpTunnelSessionGuard { + fn close_session(&self) { + drop(self.cleanup.lock().unwrap().take()); + } +} + +struct UdpTunnelStream { + session_recv_rx: RingSocketReceiver, + session_guard: Arc, +} + +impl Stream for UdpTunnelStream { + type Item = StreamItem; + + fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + let ret = std::task::ready!(Pin::new(&mut self.get_mut().session_recv_rx).poll_next(cx)); + Poll::Ready(ret.map(|payload| { + payload + .map_err(ring_socket_error_to_tunnel) + .and_then(|datagram| zcpacket_from_udp_session_payload(&datagram.payload)) + })) + } +} + +impl Drop for UdpTunnelStream { + fn drop(&mut self) { + self.session_guard.close_session(); + } +} + +struct UdpTunnelSink { + codec: UdpSessionCodec, + session_send_tx: RingSocketSender, + closed: watch::Receiver, + session_guard: Arc, +} + +impl Sink for UdpTunnelSink { + type Error = SinkError; + + fn poll_ready(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + let this = self.get_mut(); + if *this.closed.borrow() { + return Poll::Ready(Err(TunnelError::Shutdown)); + } + Pin::new(&mut this.session_send_tx) + .poll_ready(cx) + .map_err(ring_socket_error_to_tunnel) + } + + fn start_send(self: Pin<&mut Self>, item: SinkItem) -> Result<(), Self::Error> { + let this = self.get_mut(); + if *this.closed.borrow() { + return Err(TunnelError::Shutdown); + } + + let packet = item.convert_type(ZCPacketType::UDP); + let payload = BytesMut::from(packet.udp_payload()); + this.codec + .validate_payload(&payload) + .map_err(TunnelError::IOError)?; + let (completion, _sent) = oneshot::channel(); + let outbound = UdpSessionOutbound { + payload, + completion, + }; + this.session_send_tx + .force_send(outbound) + .map_err(ring_send_error_to_tunnel) + } + + fn poll_flush(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + Pin::new(&mut self.get_mut().session_send_tx) + .poll_flush(cx) + .map_err(ring_socket_error_to_tunnel) + } + + fn poll_close(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + let this = self.get_mut(); + let result = Pin::new(&mut this.session_send_tx) + .poll_close(cx) + .map_err(ring_socket_error_to_tunnel); + if result.is_ready() { + this.session_guard.close_session(); + } + result + } +} + +impl Drop for UdpTunnelSink { + fn drop(&mut self) { + self.session_guard.close_session(); + } +} + +struct UdpTunnelParts { + codec: UdpSessionCodec, + session_recv_rx: RingSocketReceiver, + session_send_tx: RingSocketSender, + closed: watch::Receiver, + cleanup: UdpSessionCleanup, + keep_alive: Option>, +} + +pub struct UdpTunnel { + info: Option, + parts: StdMutex>, +} + +impl UdpTunnel { + fn new( + tunnel_info: TunnelInfo, + session_parts: UdpSessionTunnelParts, + keep_alive: Option>, + ) -> Self { + let UdpSessionTunnelParts { + local_addr, + peer_addr, + kind, + codec, + session_recv_rx, + session_send_tx, + closed, + cleanup, + } = session_parts; + tracing::debug!( + ?local_addr, + ?peer_addr, + ?kind, + "udp build tunnel from session" + ); + Self { + info: Some(tunnel_info), + parts: StdMutex::new(Some(UdpTunnelParts { + codec, + session_recv_rx, + session_send_tx, + closed, + cleanup, + keep_alive, + })), + } + } +} + +impl Tunnel for UdpTunnel { + fn split(&self) -> SplitTunnel { + let parts = self + .parts + .lock() + .unwrap() + .take() + .expect("UdpTunnel can only be split once"); + let session_guard = Arc::new(UdpTunnelSessionGuard { + cleanup: StdMutex::new(Some(parts.cleanup)), + _layer_guard: parts.keep_alive, + }); + ( + Box::pin(UdpTunnelStream { + session_recv_rx: parts.session_recv_rx, + session_guard: session_guard.clone(), + }), + Box::pin(UdpTunnelSink { + codec: parts.codec, + session_send_tx: parts.session_send_tx, + closed: parts.closed, + session_guard, + }), + ) + } + + fn info(&self) -> Option { + self.info.clone() + } +} + +pub struct UdpTunnelUpgrader { + tunnel_info: TunnelInfo, + keep_alive: Option>, +} + +impl UdpTunnelUpgrader { + pub fn new(tunnel_info: TunnelInfo) -> Self { + Self { + tunnel_info, + keep_alive: None, + } + } + + pub fn with_keep_alive(tunnel_info: TunnelInfo, keep_alive: T) -> Self + where + T: Send + Sync + 'static, + { + Self { + tunnel_info, + keep_alive: Some(Box::new(keep_alive)), + } + } + + pub fn upgrade(self, session: UdpSession) -> Result, TunnelError> { + let Self { + tunnel_info, + keep_alive, + } = self; + Ok(Box::new(UdpTunnel::new( + tunnel_info, + session.into_tunnel_parts(), + keep_alive, + ))) + } +} + +#[cfg(test)] +mod tests { + use std::{ + io, + net::SocketAddr, + sync::{Arc, Mutex as StdMutex}, + time::Duration, + }; + + use async_trait::async_trait; + use futures::{SinkExt, StreamExt}; + use tokio::{ + sync::mpsc::{Receiver, channel}, + time::{sleep, timeout}, + }; + + use crate::socket::udp::{UdpSessionKind, VirtualUdpSocket}; + + use super::*; + + struct MockVirtualUdpSocket { + local_addr: SocketAddr, + inbound: tokio::sync::Mutex, SocketAddr)>>, + sent: StdMutex, SocketAddr)>>, + } + + impl MockVirtualUdpSocket { + fn new(local_addr: SocketAddr, inbound: Receiver<(Vec, SocketAddr)>) -> Self { + Self { + local_addr, + inbound: tokio::sync::Mutex::new(inbound), + sent: StdMutex::new(Vec::new()), + } + } + + fn sent(&self) -> Vec<(Vec, SocketAddr)> { + self.sent.lock().unwrap().clone() + } + } + + #[async_trait] + impl VirtualUdpSocket for MockVirtualUdpSocket { + fn local_addr(&self) -> io::Result { + Ok(self.local_addr) + } + + async fn send_to(&self, data: &[u8], addr: SocketAddr) -> io::Result { + self.sent.lock().unwrap().push((data.to_vec(), addr)); + Ok(data.len()) + } + + async fn recv_from(&self, buf: &mut [u8]) -> io::Result<(usize, SocketAddr)> { + let mut inbound = self.inbound.lock().await; + let Some((data, addr)) = inbound.recv().await else { + return Err(io::Error::new( + io::ErrorKind::UnexpectedEof, + "mock socket closed", + )); + }; + let len = data.len().min(buf.len()); + buf[..len].copy_from_slice(&data[..len]); + Ok((len, addr)) + } + } + + async fn wait_until(mut condition: F) + where + F: FnMut() -> bool, + { + timeout(Duration::from_secs(1), async { + loop { + if condition() { + return; + } + sleep(Duration::from_millis(10)).await; + } + }) + .await + .unwrap(); + } + + #[tokio::test] + async fn udp_tunnel_upgrader_consumes_udp_session_rings() { + let local_addr = "127.0.0.1:1".parse().unwrap(); + let peer_addr = "127.0.0.1:2".parse().unwrap(); + let (network_sender, network_recv) = channel(8); + let socket = Arc::new(MockVirtualUdpSocket::new(local_addr, network_recv)); + let session = + UdpSession::identity_standalone(socket.clone(), peer_addr, UdpSessionKind::Quic) + .unwrap(); + let tunnel = UdpTunnelUpgrader::new(TunnelInfo { + tunnel_type: "udp".to_owned(), + local_addr: None, + remote_addr: None, + resolved_remote_addr: None, + }) + .upgrade(session) + .unwrap(); + let (mut stream, mut sink) = tunnel.split(); + + let outbound_packet = ZCPacket::new_with_payload(b"outbound"); + let expected_session_payload = outbound_packet + .clone() + .convert_type(ZCPacketType::UDP) + .udp_payload() + .to_vec(); + sink.send(outbound_packet).await.unwrap(); + wait_until(|| !socket.sent().is_empty()).await; + assert_eq!(socket.sent(), vec![(expected_session_payload, peer_addr)]); + + network_sender + .send((b"inbound".to_vec(), peer_addr)) + .await + .unwrap(); + let packet = timeout(Duration::from_secs(1), stream.next()) + .await + .unwrap() + .unwrap() + .unwrap(); + + assert_eq!(packet.udp_payload(), b"inbound"); + } + + #[tokio::test] + async fn udp_tunnel_sink_close_closes_session() { + let local_addr = "127.0.0.1:1".parse().unwrap(); + let peer_addr = "127.0.0.1:2".parse().unwrap(); + let (_network_sender, network_recv) = channel(8); + let socket = Arc::new(MockVirtualUdpSocket::new(local_addr, network_recv)); + let session = + UdpSession::identity_standalone(socket, peer_addr, UdpSessionKind::Quic).unwrap(); + let tunnel = UdpTunnelUpgrader::new(TunnelInfo { + tunnel_type: "udp".to_owned(), + local_addr: None, + remote_addr: None, + resolved_remote_addr: None, + }) + .upgrade(session) + .unwrap(); + let (mut stream, mut sink) = tunnel.split(); + + sink.close().await.unwrap(); + let item = timeout(Duration::from_secs(1), stream.next()) + .await + .unwrap(); + assert!( + item.is_none() || matches!(item, Some(Err(TunnelError::Shutdown))), + "session stream must close after sink close" + ); + } +} diff --git a/easytier/src/web_client/security.rs b/easytier-core/src/tunnel/web_security.rs similarity index 94% rename from easytier/src/web_client/security.rs rename to easytier-core/src/tunnel/web_security.rs index 2240a9ab..b09d4790 100644 --- a/easytier/src/web_client/security.rs +++ b/easytier-core/src/tunnel/web_security.rs @@ -1,17 +1,16 @@ use std::sync::{Arc, Mutex}; -use std::time::Duration; use futures::{SinkExt, StreamExt}; use snow::{Builder, params::NoiseParams}; use crate::{ - common::config::EncryptionAlgorithm, - peers::secure_datagram::{SecureDatagramDirection, SecureDatagramSession}, + foundation::time::{Duration, timeout}, + packet::{PacketType, ZCPacket, ZCPacketType}, proto::common::TunnelInfo, tunnel::{ SplitTunnel, StreamItem, Tunnel, TunnelError, ZCPacketSink, ZCPacketStream, filter::{TunnelFilter, TunnelWithFilter}, - packet_def::{PacketType, ZCPacket, ZCPacketType}, + secure_datagram::{SecureDatagramDirection, SecureDatagramSession}, }, }; @@ -159,9 +158,7 @@ fn decode_noise_payload(payload: &[u8]) -> Option<&[u8]> { } pub fn web_secure_tunnel_supported() -> bool { - WEB_SECURE_CIPHER_ALGORITHM - .parse::() - .is_ok() + cfg!(feature = "aes-gcm") } fn web_secure_cipher_algorithm() -> Result<&'static str, TunnelError> { @@ -223,8 +220,7 @@ pub async fn upgrade_client_tunnel( ))) .await?; - let msg2_packet = match tokio::time::timeout(WEB_SECURE_HANDSHAKE_TIMEOUT, stream.next()).await - { + let msg2_packet = match timeout(WEB_SECURE_HANDSHAKE_TIMEOUT, stream.next()).await { Ok(Some(Ok(packet))) => packet, Ok(Some(Err(error))) => return Err(error), Ok(None) => return Err(TunnelError::Shutdown), @@ -260,7 +256,7 @@ pub async fn accept_or_upgrade_server_tunnel( let mut stream = stream; let mut sink = sink; - let first_packet = match tokio::time::timeout(WEB_SECURE_ACCEPT_TIMEOUT, stream.next()).await { + let first_packet = match timeout(WEB_SECURE_ACCEPT_TIMEOUT, stream.next()).await { Ok(Some(Ok(packet))) => packet, Ok(Some(Err(error))) => return Err(error), Ok(None) => return Err(TunnelError::Shutdown), @@ -319,10 +315,12 @@ pub async fn accept_or_upgrade_server_tunnel( #[cfg(test)] mod tests { - use super::*; - use crate::tunnel::ring::create_ring_tunnel_pair; use bytes::BytesMut; + use crate::{foundation::time::sleep, tunnel::ring::create_ring_tunnel_pair}; + + use super::*; + #[test] fn web_secure_cipher_algorithm_matches_support_flag() { let result = web_secure_cipher_algorithm(); @@ -350,6 +348,10 @@ mod tests { #[tokio::test] async fn upgrade_client_tunnel_times_out_when_server_never_replies() { + if !web_secure_tunnel_supported() { + return; + } + let (server_tunnel, client_tunnel) = create_ring_tunnel_pair(); let _server_tunnel = server_tunnel; @@ -381,12 +383,16 @@ mod tests { #[tokio::test] async fn accept_secure_tunnel_after_short_client_delay() { + if !web_secure_tunnel_supported() { + return; + } + let (server_tunnel, client_tunnel) = create_ring_tunnel_pair(); let server_task = tokio::spawn(async move { accept_or_upgrade_server_tunnel(server_tunnel).await }); - tokio::time::sleep(Duration::from_millis(1500)).await; + sleep(Duration::from_millis(1500)).await; let client_task = tokio::spawn(async move { upgrade_client_tunnel(client_tunnel).await }); diff --git a/easytier-core/src/tunnel/wrapper.rs b/easytier-core/src/tunnel/wrapper.rs new file mode 100644 index 00000000..03af61ea --- /dev/null +++ b/easytier-core/src/tunnel/wrapper.rs @@ -0,0 +1,52 @@ +use std::{ + any::Any, + pin::Pin, + sync::{Arc, Mutex}, +}; + +use crate::proto::common::TunnelInfo; + +use super::{Tunnel, ZCPacketSink, ZCPacketStream}; + +pub struct TunnelWrapper { + reader: Arc>>, + writer: Arc>>, + info: Option, + _associate_data: Option>, +} + +impl TunnelWrapper { + pub fn new(reader: R, writer: W, info: Option) -> Self { + Self::new_with_associate_data(reader, writer, info, None) + } + + pub fn new_with_associate_data( + reader: R, + writer: W, + info: Option, + associate_data: Option>, + ) -> Self { + Self { + reader: Arc::new(Mutex::new(Some(reader))), + writer: Arc::new(Mutex::new(Some(writer))), + info, + _associate_data: associate_data, + } + } +} + +impl Tunnel for TunnelWrapper +where + R: ZCPacketStream + Send + 'static, + W: ZCPacketSink + Send + 'static, +{ + fn split(&self) -> (Pin>, Pin>) { + let reader = self.reader.lock().unwrap().take().unwrap(); + let writer = self.writer.lock().unwrap().take().unwrap(); + (Box::pin(reader), Box::pin(writer)) + } + + fn info(&self) -> Option { + self.info.clone() + } +} diff --git a/easytier-core/src/wasi/abi.rs b/easytier-core/src/wasi/abi.rs new file mode 100644 index 00000000..4a1031a0 --- /dev/null +++ b/easytier-core/src/wasi/abi.rs @@ -0,0 +1,85 @@ +//! Stable ABI contract between the EasyTier WASI guest and its runtime. +//! +//! A runtime must provide every function in the `easytier_host` import module +//! and call the guest lifecycle exports listed in [`GUEST_EXPORTS`]. The raw +//! import declarations live in the target-only `imports` Module. +//! Rust visibility does not define this cross-language contract; the imported +//! and exported WebAssembly symbol names and signatures do. +//! +//! Every `u32` pointer is an offset in wasm32 guest linear memory, never a +//! native host pointer. The runtime must copy input bytes before an import +//! returns and may write result bytes only during the matching `take_*` call. +//! An operation ID belongs to core until a terminal `take_*` call consumes it +//! or the runtime receives the `cancel_operation` import. + +/// WebAssembly import module a WASI runtime must implement. +pub const HOST_IMPORT_MODULE: &str = "easytier_host"; + +/// Version of the JSON document accepted by `easytier_instance_create`. +pub const CORE_INSTANCE_CONFIG_VERSION: u32 = 14; + +/// Version of the public data-plane guest export contract. +pub const DATA_PLANE_ABI_VERSION: u32 = 2; + +/// The guest exposes an instance-scoped data-plane operation broker. +pub const DATA_PLANE_CAPABILITY: u64 = 1 << 0; +/// The guest data plane supports TCP streams and listeners. +pub const DATA_PLANE_TCP_CAPABILITY: u64 = 1 << 1; +/// The guest data plane supports UDP sockets. +pub const DATA_PLANE_UDP_CAPABILITY: u64 = 1 << 2; + +/// Guest exports a WASI runtime calls to manage a core instance. +/// +/// Buffer allocation precedes config, packet, and error-copy calls. Instance +/// creation and lifecycle use the returned instance handle. A runtime drives +/// all asynchronous guest work through `easytier_instance_drive` and host +/// completion notifications. +pub const GUEST_EXPORTS: &[&str] = &[ + // Guest-memory buffers. + "easytier_buffer_alloc", + "easytier_buffer_free", + // Instance lifecycle and external runtime driving. + "easytier_instance_create", + "easytier_instance_start", + "easytier_instance_stop", + "easytier_instance_drive", + "easytier_instance_notify_completions", + "easytier_instance_state", + "easytier_instance_next_deadline_millis", + // Raw IP packet ingress and error retrieval. + "easytier_instance_send_packet", + "easytier_instance_drop", + "easytier_instance_error_len", + "easytier_instance_error_copy", +]; + +/// Guest exports present when the core is built with the smoltcp data plane. +#[cfg(feature = "proxy-smoltcp-stack")] +pub const DATA_PLANE_GUEST_EXPORTS: &[&str] = &[ + // ABI discovery. + "easytier_data_plane_abi_version", + "easytier_data_plane_capabilities", + // Data-plane operation submission. + "easytier_data_plane_tcp_connect_submit", + "easytier_data_plane_tcp_bind_submit", + "easytier_data_plane_tcp_accept_submit", + "easytier_data_plane_tcp_read_submit", + "easytier_data_plane_tcp_write_submit", + "easytier_data_plane_udp_bind_submit", + "easytier_data_plane_udp_receive_submit", + "easytier_data_plane_udp_send_submit", + // Completion, result, and resource lifecycle. + "easytier_data_plane_completion_drain", + "easytier_data_plane_result_size", + "easytier_data_plane_tcp_connect_result_take", + "easytier_data_plane_tcp_bind_result_take", + "easytier_data_plane_tcp_accept_result_take", + "easytier_data_plane_tcp_read_result_take", + "easytier_data_plane_tcp_write_result_take", + "easytier_data_plane_udp_bind_result_take", + "easytier_data_plane_udp_receive_result_take", + "easytier_data_plane_udp_send_result_take", + "easytier_data_plane_operation_cancel", + "easytier_data_plane_operation_free", + "easytier_data_plane_resource_close", +]; diff --git a/easytier-core/src/wasi/adapter/dns.rs b/easytier-core/src/wasi/adapter/dns.rs new file mode 100644 index 00000000..cc6478b0 --- /dev/null +++ b/easytier-core/src/wasi/adapter/dns.rs @@ -0,0 +1,121 @@ +use std::{io, net::IpAddr, task::Poll}; + +use crate::{ + host::{ + dns::{DnsQuery, DnsSrvRecord, HostDnsIo}, + socket::HostOperationId, + }, + wasi::{ + imports::{ + HOST_PENDING, cancel_operation, start_dns_resolve, start_dns_srv, start_dns_txt, + take_dns_resolve, take_dns_srv, take_dns_txt, + }, + wire::{ + common::host_error, + dns::{decode_addresses, decode_srv, decode_txt, encode_query}, + }, + }, +}; + +const MAX_DNS_RESULT_LEN: usize = 1024 * 1024; + +#[derive(Default)] +pub struct WasiHostDnsIo; + +impl HostDnsIo for WasiHostDnsIo { + fn submit_resolve(&self, operation: HostOperationId, query: &DnsQuery) -> io::Result<()> { + submit_query("start_dns_resolve", operation, query, start_dns_resolve) + } + + fn take_resolve(&self, operation: HostOperationId) -> Poll>> { + take_result( + "take_dns_resolve", + operation, + take_dns_resolve, + decode_addresses, + ) + } + + fn submit_txt(&self, operation: HostOperationId, query: &DnsQuery) -> io::Result<()> { + submit_query("start_dns_txt", operation, query, start_dns_txt) + } + + fn take_txt(&self, operation: HostOperationId) -> Poll> { + take_result("take_dns_txt", operation, take_dns_txt, decode_txt) + } + + fn submit_srv(&self, operation: HostOperationId, query: &DnsQuery) -> io::Result<()> { + submit_query("start_dns_srv", operation, query, start_dns_srv) + } + + fn take_srv(&self, operation: HostOperationId) -> Poll>> { + take_result("take_dns_srv", operation, take_dns_srv, decode_srv) + } + + fn cancel_operation(&self, operation: HostOperationId) -> io::Result<()> { + host_status("cancel_operation", unsafe { cancel_operation(operation.0) }) + } +} + +fn submit_query( + name: &'static str, + operation: HostOperationId, + query: &DnsQuery, + submit: unsafe extern "C" fn(u64, u32, u32) -> i32, +) -> io::Result<()> { + let encoded = encode_query(query)?; + let length = u32::try_from(encoded.len()) + .map_err(|_| io::Error::new(io::ErrorKind::InvalidInput, "DNS query is too long"))?; + host_status(name, unsafe { + submit(operation.0, encoded.as_ptr() as u32, length) + }) +} + +fn take_result( + name: &'static str, + operation: HostOperationId, + take: unsafe extern "C" fn(u64, u32, u32) -> i32, + decode: fn(&[u8]) -> io::Result, +) -> Poll> { + let required = unsafe { take(operation.0, 0, 0) }; + if required == HOST_PENDING { + return Poll::Pending; + } + if required <= 0 { + return Poll::Ready(Err(host_error(name, required))); + } + let required = usize::try_from(required).expect("positive i32 fits usize"); + if required > MAX_DNS_RESULT_LEN { + cancel_probed_result(operation); + return Poll::Ready(Err(io::Error::new( + io::ErrorKind::InvalidData, + format!("host {name} result exceeds {MAX_DNS_RESULT_LEN} bytes"), + ))); + } + let mut encoded = vec![0_u8; required]; + let capacity = u32::try_from(required).expect("DNS result limit fits u32"); + let copied = unsafe { take(operation.0, encoded.as_mut_ptr() as u32, capacity) }; + if copied != i32::try_from(required).expect("positive DNS result length fits i32") { + cancel_probed_result(operation); + return Poll::Ready(Err(io::Error::new( + io::ErrorKind::InvalidData, + format!("host {name} changed result length from {required} to {copied}"), + ))); + } + Poll::Ready(decode(&encoded)) +} + +fn cancel_probed_result(operation: HostOperationId) { + // The size probe does not consume host state. A second take may have + // consumed it before returning a malformed status, so this best-effort + // ownership cleanup must not hide the original protocol error. + let _ = unsafe { cancel_operation(operation.0) }; +} + +fn host_status(name: &'static str, result: i32) -> io::Result<()> { + if result == 0 { + Ok(()) + } else { + Err(host_error(name, result)) + } +} diff --git a/easytier-core/src/wasi/adapter/environment.rs b/easytier-core/src/wasi/adapter/environment.rs new file mode 100644 index 00000000..33b97273 --- /dev/null +++ b/easytier-core/src/wasi/adapter/environment.rs @@ -0,0 +1,76 @@ +//! WASI imports for connector environment operations. + +use std::{io, net::SocketAddr, task::Poll}; + +use crate::socket::SocketContext; +use crate::{ + host::{environment::HostConnectorEnvironmentIo, socket::HostOperationId}, + wasi::{ + imports::{ + HOST_PENDING, cancel_operation, start_local_addr_for_remote, take_local_addr_for_remote, + }, + wire::{ + common::{host_error, status}, + options::encode_socket_context, + socket::{SOCKET_ADDRESS_LEN, decode_socket_address, encode_socket_address}, + }, + }, +}; + +#[derive(Default)] +pub struct WasiHostConnectorEnvironmentIo; + +impl HostConnectorEnvironmentIo for WasiHostConnectorEnvironmentIo { + fn submit_local_addr_for_remote( + &self, + operation: HostOperationId, + remote_addr: SocketAddr, + context: &SocketContext, + ) -> io::Result<()> { + let encoded = encode_socket_address(remote_addr); + let encoded_context = encode_socket_context(context)?; + status("start_local_addr_for_remote", unsafe { + start_local_addr_for_remote( + operation.0, + encoded.as_ptr() as u32, + SOCKET_ADDRESS_LEN as u32, + encoded_context.as_ptr() as u32, + encoded_context.len() as u32, + ) + }) + } + + fn take_local_addr_for_remote( + &self, + operation: HostOperationId, + ) -> Poll> { + take_address( + "take_local_addr_for_remote", + operation, + take_local_addr_for_remote, + ) + } + + fn cancel_operation(&self, operation: HostOperationId) -> io::Result<()> { + status("cancel_operation", unsafe { cancel_operation(operation.0) }) + } +} + +fn take_address( + name: &'static str, + operation: HostOperationId, + take: unsafe extern "C" fn(u64, u32, u32) -> i32, +) -> Poll> { + let mut encoded = [0_u8; SOCKET_ADDRESS_LEN]; + match unsafe { + take( + operation.0, + encoded.as_mut_ptr() as u32, + SOCKET_ADDRESS_LEN as u32, + ) + } { + HOST_PENDING => Poll::Pending, + 0 => Poll::Ready(decode_socket_address(&encoded)), + value => Poll::Ready(Err(host_error(name, value))), + } +} diff --git a/easytier-core/src/wasi/adapter/mod.rs b/easytier-core/src/wasi/adapter/mod.rs new file mode 100644 index 00000000..70008e19 --- /dev/null +++ b/easytier-core/src/wasi/adapter/mod.rs @@ -0,0 +1,6 @@ +//! Concrete implementations of portable Host capability seams for WASI. + +pub mod dns; +pub mod environment; +pub mod packet; +pub mod socket; diff --git a/easytier-core/src/wasi/adapter/packet.rs b/easytier-core/src/wasi/adapter/packet.rs new file mode 100644 index 00000000..0e5c56ee --- /dev/null +++ b/easytier-core/src/wasi/adapter/packet.rs @@ -0,0 +1,59 @@ +use std::{io, task::Poll}; + +use crate::{ + host::{ + packet::{HostPacketIo, HostPacketSinkHandle}, + socket::HostOperationId, + }, + wasi::{ + imports::{ + HOST_PENDING, HOST_WOULD_BLOCK, cancel_operation, start_packet_write_ready, + take_packet_write_ready, try_packet_write, + }, + wire::common::{host_error, status}, + }, +}; + +const MAX_HOST_PACKET_LEN: usize = 1024 * 1024; + +#[derive(Default)] +pub struct WasiHostPacketIo; + +impl HostPacketIo for WasiHostPacketIo { + fn try_write_packet(&self, handle: HostPacketSinkHandle, packet: &[u8]) -> io::Result<()> { + if packet.len() > MAX_HOST_PACKET_LEN { + return Err(io::Error::new( + io::ErrorKind::InvalidInput, + format!("host packet exceeds {MAX_HOST_PACKET_LEN} bytes"), + )); + } + let length = u32::try_from(packet.len()).expect("host packet limit fits u32"); + match unsafe { try_packet_write(handle.0, packet.as_ptr() as u32, length) } { + 0 => Ok(()), + HOST_WOULD_BLOCK => Err(io::ErrorKind::WouldBlock.into()), + value => Err(host_error("try_packet_write", value)), + } + } + + fn submit_write_ready( + &self, + handle: HostPacketSinkHandle, + operation: HostOperationId, + ) -> io::Result<()> { + status("start_packet_write_ready", unsafe { + start_packet_write_ready(handle.0, operation.0) + }) + } + + fn take_write_ready(&self, operation: HostOperationId) -> Poll> { + match unsafe { take_packet_write_ready(operation.0) } { + HOST_PENDING => Poll::Pending, + 0 => Poll::Ready(Ok(())), + value => Poll::Ready(Err(host_error("take_packet_write_ready", value))), + } + } + + fn cancel_operation(&self, operation: HostOperationId) -> io::Result<()> { + status("cancel_operation", unsafe { cancel_operation(operation.0) }) + } +} diff --git a/easytier-core/src/wasi/adapter/socket/backend.rs b/easytier-core/src/wasi/adapter/socket/backend.rs new file mode 100644 index 00000000..46914887 --- /dev/null +++ b/easytier-core/src/wasi/adapter/socket/backend.rs @@ -0,0 +1,266 @@ +use std::{io, task::Poll}; + +use crate::host::socket::{ + HostOperationId, HostSocketHandle, HostSocketIo, HostTcpIo, + factory::{HostSocketFactoryIo, HostTcpConnectResult, HostUdpBindResult}, + listener::{HostTcpBindResult, HostTcpListenerIo}, + udp::{HostUdpDatagram, HostUdpIo}, +}; +use crate::socket::{ + tcp::{TcpConnectOptions, TcpListenOptions}, + udp::{UdpBindOptions, UdpSocketSendMeta}, +}; + +use crate::wasi::{ + imports::{ + HOST_PENDING, cancel_operation, close, start_tcp_accept, start_tcp_bind, start_tcp_connect, + start_udp_bind, take_tcp_accept, take_tcp_bind, take_tcp_connect, take_udp_bind, + }, + wire::{ + common::{host_error, status, tcp_connect_error}, + options::{ + BOUND_SOCKET_RESULT_LEN, TCP_SOCKET_RESULT_LEN, decode_tcp_bind_result, + decode_tcp_socket_result, decode_udp_bind_result, encode_tcp_connect_options, + encode_tcp_listen_options, encode_udp_bind_options, + }, + }, +}; + +use super::{WasiHostTcpIo, udp::WasiHostUdpIo}; + +#[derive(Default)] +pub struct WasiHostSocketBackend { + tcp: WasiHostTcpIo, + udp: WasiHostUdpIo, +} + +impl WasiHostSocketBackend { + fn decode_transferred(&self, encoded: &[u8], decoded: io::Result) -> io::Result { + match decoded { + Ok(result) => Ok(result), + Err(decode_error) => { + let handle = HostSocketHandle(u64::from_be_bytes(encoded[..8].try_into().unwrap())); + match self.close(handle) { + Ok(()) => Err(decode_error), + Err(close_error) => Err(io::Error::new( + decode_error.kind(), + format!( + "{decode_error}; additionally failed to close malformed host result: {close_error}" + ), + )), + } + } + } + } +} + +impl HostSocketIo for WasiHostSocketBackend { + fn cancel_operation(&self, operation: HostOperationId) -> io::Result<()> { + self.tcp.forget_operation(operation); + self.udp.forget_operation(operation); + status("cancel_operation", unsafe { cancel_operation(operation.0) }) + } + + fn close(&self, handle: HostSocketHandle) -> io::Result<()> { + status("close", unsafe { close(handle.0) }) + } +} + +impl HostTcpIo for WasiHostSocketBackend { + fn submit_read( + &self, + handle: HostSocketHandle, + operation: HostOperationId, + capacity: usize, + ) -> io::Result<()> { + self.tcp.submit_read(handle, operation, capacity) + } + + fn take_read(&self, operation: HostOperationId) -> Poll>> { + self.tcp.take_read(operation) + } + + fn submit_write( + &self, + handle: HostSocketHandle, + operation: HostOperationId, + source: &[u8], + ) -> io::Result<()> { + self.tcp.submit_write(handle, operation, source) + } + + fn take_write(&self, operation: HostOperationId) -> Poll> { + self.tcp.take_write(operation) + } +} + +impl HostUdpIo for WasiHostSocketBackend { + fn submit_recv( + &self, + handle: HostSocketHandle, + operation: HostOperationId, + capacity: usize, + ) -> io::Result<()> { + self.udp.submit_recv(handle, operation, capacity) + } + + fn take_recv(&self, operation: HostOperationId) -> Poll> { + self.udp.take_recv(operation) + } + + fn try_send( + &self, + handle: HostSocketHandle, + source: &[u8], + peer_addr: std::net::SocketAddr, + meta: UdpSocketSendMeta, + ) -> io::Result<()> { + self.udp.try_send(handle, source, peer_addr, meta) + } + + fn submit_send_ready( + &self, + handle: HostSocketHandle, + operation: HostOperationId, + ) -> io::Result<()> { + self.udp.submit_send_ready(handle, operation) + } + + fn take_send_ready(&self, operation: HostOperationId) -> Poll> { + self.udp.take_send_ready(operation) + } +} + +impl HostSocketFactoryIo for WasiHostSocketBackend { + fn submit_tcp_connect( + &self, + operation: HostOperationId, + options: &TcpConnectOptions, + ) -> io::Result<()> { + let encoded = encode_tcp_connect_options(options)?; + status("start_tcp_connect", unsafe { + start_tcp_connect( + operation.0, + encoded.as_ptr() as u32, + encoded_len("TCP connect options", &encoded)?, + ) + }) + } + + fn take_tcp_connect( + &self, + operation: HostOperationId, + ) -> Poll> { + let mut encoded = [0_u8; TCP_SOCKET_RESULT_LEN]; + match unsafe { + take_tcp_connect( + operation.0, + encoded.as_mut_ptr() as u32, + TCP_SOCKET_RESULT_LEN as u32, + ) + } { + HOST_PENDING => Poll::Pending, + 0 => Poll::Ready(self.decode_transferred(&encoded, decode_tcp_socket_result(&encoded))), + value => Poll::Ready(Err(tcp_connect_error(value))), + } + } + + fn submit_udp_bind( + &self, + operation: HostOperationId, + options: &UdpBindOptions, + ) -> io::Result<()> { + let encoded = encode_udp_bind_options(options)?; + status("start_udp_bind", unsafe { + start_udp_bind( + operation.0, + encoded.as_ptr() as u32, + encoded_len("UDP bind options", &encoded)?, + ) + }) + } + + fn take_udp_bind(&self, operation: HostOperationId) -> Poll> { + let mut encoded = [0_u8; BOUND_SOCKET_RESULT_LEN]; + match unsafe { + take_udp_bind( + operation.0, + encoded.as_mut_ptr() as u32, + BOUND_SOCKET_RESULT_LEN as u32, + ) + } { + HOST_PENDING => Poll::Pending, + 0 => Poll::Ready(self.decode_transferred(&encoded, decode_udp_bind_result(&encoded))), + value => Poll::Ready(Err(host_error("take_udp_bind", value))), + } + } +} + +impl HostTcpListenerIo for WasiHostSocketBackend { + fn submit_tcp_bind( + &self, + operation: HostOperationId, + options: &TcpListenOptions, + ) -> io::Result<()> { + let encoded = encode_tcp_listen_options(options)?; + status("start_tcp_bind", unsafe { + start_tcp_bind( + operation.0, + encoded.as_ptr() as u32, + encoded_len("TCP listen options", &encoded)?, + ) + }) + } + + fn take_tcp_bind(&self, operation: HostOperationId) -> Poll> { + let mut encoded = [0_u8; BOUND_SOCKET_RESULT_LEN]; + match unsafe { + take_tcp_bind( + operation.0, + encoded.as_mut_ptr() as u32, + BOUND_SOCKET_RESULT_LEN as u32, + ) + } { + HOST_PENDING => Poll::Pending, + 0 => Poll::Ready(self.decode_transferred(&encoded, decode_tcp_bind_result(&encoded))), + value => Poll::Ready(Err(host_error("take_tcp_bind", value))), + } + } + + fn submit_tcp_accept( + &self, + handle: HostSocketHandle, + operation: HostOperationId, + ) -> io::Result<()> { + status("start_tcp_accept", unsafe { + start_tcp_accept(handle.0, operation.0) + }) + } + + fn take_tcp_accept( + &self, + operation: HostOperationId, + ) -> Poll> { + let mut encoded = [0_u8; TCP_SOCKET_RESULT_LEN]; + match unsafe { + take_tcp_accept( + operation.0, + encoded.as_mut_ptr() as u32, + TCP_SOCKET_RESULT_LEN as u32, + ) + } { + HOST_PENDING => Poll::Pending, + 0 => Poll::Ready(self.decode_transferred(&encoded, decode_tcp_socket_result(&encoded))), + value => Poll::Ready(Err(host_error("take_tcp_accept", value))), + } + } +} + +fn encoded_len(description: &str, encoded: &[u8]) -> io::Result { + u32::try_from(encoded.len()).map_err(|_| { + io::Error::new( + io::ErrorKind::InvalidInput, + format!("{description} exceed WASI guest memory"), + ) + }) +} diff --git a/easytier-core/src/wasi/adapter/socket/mod.rs b/easytier-core/src/wasi/adapter/socket/mod.rs new file mode 100644 index 00000000..8f9d5231 --- /dev/null +++ b/easytier-core/src/wasi/adapter/socket/mod.rs @@ -0,0 +1,124 @@ +//! WASI TCP adapter behind the portable host socket seams. +//! +//! [`WasiHostTcpIo`] implements the TCP half of the host socket I/O traits on +//! top of the `easytier_host` import module, and +//! [`backend::WasiHostSocketBackend`] composes the TCP, UDP, factory, and +//! listener imports into one host socket backend. Shared status mapping and +//! wire codecs live in [`crate::wasi::wire`]. + +pub mod backend; +pub mod udp; + +use std::{collections::HashMap, io, sync::Mutex, task::Poll}; + +use crate::{ + host::socket::{HostOperationId, HostSocketHandle, HostSocketIo, HostTcpIo}, + wasi::{ + imports::{ + HOST_PENDING, cancel_operation, close, start_read, start_write, take_read, take_write, + }, + wire::common::{host_error, status}, + }, +}; + +#[derive(Debug, Default)] +pub struct WasiHostTcpIo { + read_buffers: Mutex>>, +} + +impl WasiHostTcpIo { + pub(super) fn forget_operation(&self, operation: HostOperationId) { + self.read_buffers + .lock() + .expect("WASI read buffer registry poisoned") + .remove(&operation); + } +} + +impl HostSocketIo for WasiHostTcpIo { + fn cancel_operation(&self, operation: HostOperationId) -> io::Result<()> { + self.forget_operation(operation); + status("cancel_operation", unsafe { cancel_operation(operation.0) }) + } + + fn close(&self, handle: HostSocketHandle) -> io::Result<()> { + status("close", unsafe { close(handle.0) }) + } +} + +impl HostTcpIo for WasiHostTcpIo { + fn submit_read( + &self, + handle: HostSocketHandle, + operation: HostOperationId, + capacity: usize, + ) -> io::Result<()> { + let capacity_u32 = u32::try_from(capacity) + .map_err(|_| io::Error::new(io::ErrorKind::InvalidInput, "read buffer is too large"))?; + status("start_read", unsafe { + start_read(handle.0, operation.0, capacity_u32) + })?; + self.read_buffers + .lock() + .expect("WASI read buffer registry poisoned") + .insert(operation, vec![0; capacity]); + Ok(()) + } + + fn take_read(&self, operation: HostOperationId) -> Poll>> { + let mut buffers = self + .read_buffers + .lock() + .expect("WASI read buffer registry poisoned"); + let Some(buffer) = buffers.get_mut(&operation) else { + return Poll::Ready(Err(io::Error::new( + io::ErrorKind::NotFound, + "WASI read operation buffer is missing", + ))); + }; + let result = + unsafe { take_read(operation.0, buffer.as_mut_ptr() as u32, buffer.len() as u32) }; + match result { + HOST_PENDING => Poll::Pending, + value if value >= 0 => { + let length = value as usize; + if length > buffer.len() { + buffers.remove(&operation); + return Poll::Ready(Err(io::Error::new( + io::ErrorKind::InvalidData, + "host read completion exceeds the submitted capacity", + ))); + } + let mut buffer = buffers.remove(&operation).unwrap(); + buffer.truncate(length); + Poll::Ready(Ok(buffer)) + } + value => { + buffers.remove(&operation); + Poll::Ready(Err(host_error("take_read", value))) + } + } + } + + fn submit_write( + &self, + handle: HostSocketHandle, + operation: HostOperationId, + source: &[u8], + ) -> io::Result<()> { + let length = u32::try_from(source.len()).map_err(|_| { + io::Error::new(io::ErrorKind::InvalidInput, "write buffer is too large") + })?; + status("start_write", unsafe { + start_write(handle.0, operation.0, source.as_ptr() as u32, length) + }) + } + + fn take_write(&self, operation: HostOperationId) -> Poll> { + match unsafe { take_write(operation.0) } { + HOST_PENDING => Poll::Pending, + 0 => Poll::Ready(Ok(())), + value => Poll::Ready(Err(host_error("take_write", value))), + } + } +} diff --git a/easytier-core/src/wasi/adapter/socket/udp.rs b/easytier-core/src/wasi/adapter/socket/udp.rs new file mode 100644 index 00000000..cffa0527 --- /dev/null +++ b/easytier-core/src/wasi/adapter/socket/udp.rs @@ -0,0 +1,179 @@ +use std::{collections::HashMap, io, sync::Mutex, task::Poll}; + +use crate::socket::udp::{UdpSocketRecvMeta, UdpSocketSendMeta}; + +use crate::{ + host::{ + socket::udp::{HostUdpDatagram, HostUdpIo}, + socket::{HostOperationId, HostSocketHandle, HostSocketIo}, + }, + wasi::{ + imports::{ + HOST_PENDING, HOST_WOULD_BLOCK, cancel_operation, close, start_udp_recv, + start_udp_send_ready, take_udp_recv, take_udp_send_ready, try_udp_send, + }, + wire::{ + common::{host_error, status}, + socket::{UDP_METADATA_LEN, decode_udp_metadata, encode_udp_metadata}, + }, + }, +}; + +struct WasiUdpRecvBuffer { + data: Vec, + metadata: [u8; UDP_METADATA_LEN], +} + +#[derive(Default)] +pub struct WasiHostUdpIo { + recv_buffers: Mutex>, +} + +impl WasiHostUdpIo { + pub(super) fn forget_operation(&self, operation: HostOperationId) { + self.recv_buffers + .lock() + .expect("WASI UDP receive buffer registry poisoned") + .remove(&operation); + } +} + +impl HostSocketIo for WasiHostUdpIo { + fn cancel_operation(&self, operation: HostOperationId) -> io::Result<()> { + self.forget_operation(operation); + status("cancel_operation", unsafe { cancel_operation(operation.0) }) + } + + fn close(&self, handle: HostSocketHandle) -> io::Result<()> { + status("close", unsafe { close(handle.0) }) + } +} + +impl HostUdpIo for WasiHostUdpIo { + fn submit_recv( + &self, + handle: HostSocketHandle, + operation: HostOperationId, + capacity: usize, + ) -> io::Result<()> { + let capacity_u32 = u32::try_from(capacity).map_err(|_| { + io::Error::new( + io::ErrorKind::InvalidInput, + "UDP receive buffer is too large", + ) + })?; + status("start_udp_recv", unsafe { + start_udp_recv(handle.0, operation.0, capacity_u32) + })?; + self.recv_buffers + .lock() + .expect("WASI UDP receive buffer registry poisoned") + .insert( + operation, + WasiUdpRecvBuffer { + data: vec![0; capacity], + metadata: [0; UDP_METADATA_LEN], + }, + ); + Ok(()) + } + + fn take_recv(&self, operation: HostOperationId) -> Poll> { + let mut buffers = self + .recv_buffers + .lock() + .expect("WASI UDP receive buffer registry poisoned"); + let Some(buffer) = buffers.get_mut(&operation) else { + return Poll::Ready(Err(io::Error::new( + io::ErrorKind::NotFound, + "WASI UDP receive operation buffer is missing", + ))); + }; + let result = unsafe { + take_udp_recv( + operation.0, + buffer.data.as_mut_ptr() as u32, + buffer.data.len() as u32, + buffer.metadata.as_mut_ptr() as u32, + UDP_METADATA_LEN as u32, + ) + }; + match result { + HOST_PENDING => Poll::Pending, + value if value >= 0 => { + let length = value as usize; + let mut buffer = buffers.remove(&operation).unwrap(); + if length > buffer.data.len() { + return Poll::Ready(Err(io::Error::new( + io::ErrorKind::InvalidData, + "host UDP receive completion exceeds the submitted capacity", + ))); + } + buffer.data.truncate(length); + let (peer_addr, dst_ip, _) = match decode_udp_metadata(&buffer.metadata) { + Ok(metadata) => metadata, + Err(error) => return Poll::Ready(Err(error)), + }; + Poll::Ready(Ok(HostUdpDatagram { + data: buffer.data, + peer_addr, + meta: UdpSocketRecvMeta { dst_ip }, + })) + } + value => { + buffers.remove(&operation); + Poll::Ready(Err(host_error("take_udp_recv", value))) + } + } + } + + fn try_send( + &self, + handle: HostSocketHandle, + source: &[u8], + peer_addr: std::net::SocketAddr, + meta: UdpSocketSendMeta, + ) -> io::Result<()> { + if meta.src_ifindex.is_some() && !matches!(meta.src_ip, Some(std::net::IpAddr::V6(_))) { + return Err(io::Error::new( + io::ErrorKind::InvalidInput, + "UDP source interface index requires an IPv6 source address", + )); + } + let length = u32::try_from(source.len()).map_err(|_| { + io::Error::new(io::ErrorKind::InvalidInput, "UDP send buffer is too large") + })?; + let metadata = encode_udp_metadata(peer_addr, meta.src_ip, meta.src_ifindex); + match unsafe { + try_udp_send( + handle.0, + source.as_ptr() as u32, + length, + metadata.as_ptr() as u32, + UDP_METADATA_LEN as u32, + ) + } { + 0 => Ok(()), + HOST_WOULD_BLOCK => Err(io::ErrorKind::WouldBlock.into()), + value => Err(host_error("try_udp_send", value)), + } + } + + fn submit_send_ready( + &self, + handle: HostSocketHandle, + operation: HostOperationId, + ) -> io::Result<()> { + status("start_udp_send_ready", unsafe { + start_udp_send_ready(handle.0, operation.0) + }) + } + + fn take_send_ready(&self, operation: HostOperationId) -> Poll> { + match unsafe { take_udp_send_ready(operation.0) } { + HOST_PENDING => Poll::Pending, + 0 => Poll::Ready(Ok(())), + value => Poll::Ready(Err(host_error("take_udp_send_ready", value))), + } + } +} diff --git a/easytier-core/src/wasi/imports.rs b/easytier-core/src/wasi/imports.rs new file mode 100644 index 00000000..d5fb05bb --- /dev/null +++ b/easytier-core/src/wasi/imports.rs @@ -0,0 +1,139 @@ +//! Raw `easytier_host` imports, compiled only for the WASI guest. +//! +//! [`crate::wasi::abi`] declares the shared contract metadata; each import +//! below documents its own ownership and completion rules. Concrete adapters +//! call these functions directly. + +pub(crate) const HOST_PENDING: i32 = -1; +pub(crate) const HOST_WOULD_BLOCK: i32 = -5; + +#[link(wasm_import_module = "easytier_host")] +unsafe extern "C" { + /// Starts one TCP read into a host-owned pending operation. + /// + /// The host records at most `capacity` bytes for `operation` and must not + /// write guest memory until [`take_read`] supplies a destination buffer. + pub(crate) fn start_read(handle: u64, operation: u64, capacity: u32) -> i32; + + /// Copies a completed TCP read into `destination`, returning its byte count. + /// + /// Returns [`HOST_PENDING`] while the operation is incomplete. A completed + /// read, including EOF with a zero length, consumes `operation`. + pub(crate) fn take_read(operation: u64, destination: u32, capacity: u32) -> i32; + + /// Starts one TCP write after copying `source[..length]` from guest memory. + pub(crate) fn start_write(handle: u64, operation: u64, source: u32, length: u32) -> i32; + + /// Reports completion of a TCP write and consumes `operation` on success or error. + pub(crate) fn take_write(operation: u64) -> i32; + + /// Starts receipt of one UDP datagram of at most `capacity` bytes. + pub(crate) fn start_udp_recv(handle: u64, operation: u64, capacity: u32) -> i32; + + /// Copies one completed UDP datagram and its metadata into guest memory. + /// + /// A non-pending result consumes `operation`; `metadata` has exactly + /// `metadata_len` bytes allocated by core for the socket wire format. + pub(crate) fn take_udp_recv( + operation: u64, + destination: u32, + capacity: u32, + metadata: u32, + metadata_len: u32, + ) -> i32; + + /// Attempts to enqueue one complete UDP datagram after copying its bytes and metadata. + /// + /// [`HOST_WOULD_BLOCK`] means the datagram was not accepted and has no + /// side effects. Any other success means the host owns a complete copy. + pub(crate) fn try_udp_send( + handle: u64, + source: u32, + length: u32, + metadata: u32, + metadata_len: u32, + ) -> i32; + + /// Starts waiting until another UDP send attempt may succeed. + pub(crate) fn start_udp_send_ready(handle: u64, operation: u64) -> i32; + + /// Reports UDP write readiness; readiness never sends a datagram itself. + pub(crate) fn take_udp_send_ready(operation: u64) -> i32; + + /// Starts a TCP connection using an encoded `TcpConnectOptions` document. + pub(crate) fn start_tcp_connect(operation: u64, options: u32, options_len: u32) -> i32; + + /// Copies the completed TCP connection handle and addresses into `result`. + pub(crate) fn take_tcp_connect(operation: u64, result: u32, result_len: u32) -> i32; + + /// Starts a UDP bind using an encoded `UdpBindOptions` document. + pub(crate) fn start_udp_bind(operation: u64, options: u32, options_len: u32) -> i32; + + /// Copies the completed UDP socket handle and local address into `result`. + pub(crate) fn take_udp_bind(operation: u64, result: u32, result_len: u32) -> i32; + + /// Starts a TCP listener bind using an encoded `TcpListenOptions` document. + pub(crate) fn start_tcp_bind(operation: u64, options: u32, options_len: u32) -> i32; + + /// Copies the completed listener handle and local address into `result`. + pub(crate) fn take_tcp_bind(operation: u64, result: u32, result_len: u32) -> i32; + + /// Starts accepting one TCP stream from a listener handle. + pub(crate) fn start_tcp_accept(handle: u64, operation: u64) -> i32; + + /// Copies the accepted TCP stream handle and addresses into `result`. + pub(crate) fn take_tcp_accept(operation: u64, result: u32, result_len: u32) -> i32; + + /// Starts an address-record DNS lookup for an encoded [`crate::host::dns::DnsQuery`]. + pub(crate) fn start_dns_resolve(operation: u64, query: u32, query_len: u32) -> i32; + + /// Probes or copies the encoded DNS address result for `operation`. + /// + /// A zero-capacity call probes the required result length without consuming + /// it; a subsequent call with enough capacity copies and consumes it. + pub(crate) fn take_dns_resolve(operation: u64, result: u32, result_capacity: u32) -> i32; + + /// Starts a TXT-record DNS lookup for an encoded query. + pub(crate) fn start_dns_txt(operation: u64, query: u32, query_len: u32) -> i32; + + /// Probes or copies the encoded DNS TXT result using the DNS result protocol. + pub(crate) fn take_dns_txt(operation: u64, result: u32, result_capacity: u32) -> i32; + + /// Starts an SRV-record DNS lookup for an encoded query. + pub(crate) fn start_dns_srv(operation: u64, query: u32, query_len: u32) -> i32; + + /// Probes or copies the encoded DNS SRV result using the DNS result protocol. + pub(crate) fn take_dns_srv(operation: u64, result: u32, result_capacity: u32) -> i32; + + /// Starts finding the local address and source context needed to reach `remote_addr`. + pub(crate) fn start_local_addr_for_remote( + operation: u64, + remote_addr: u32, + remote_addr_len: u32, + context: u32, + context_len: u32, + ) -> i32; + + /// Copies the resolved local socket address into the fixed-size `result` buffer. + pub(crate) fn take_local_addr_for_remote(operation: u64, result: u32, result_len: u32) -> i32; + + /// Attempts to deliver one raw IP packet to a host packet sink. + /// + /// On success the host owns a complete copy. [`HOST_WOULD_BLOCK`] leaves + /// the packet unaccepted and requires core to wait for write readiness. + pub(crate) fn try_packet_write(handle: u64, packet: u32, packet_len: u32) -> i32; + + /// Starts waiting until another packet-sink admission attempt may succeed. + pub(crate) fn start_packet_write_ready(handle: u64, operation: u64) -> i32; + + /// Reports packet-sink write readiness; it never accepts a packet itself. + pub(crate) fn take_packet_write_ready(operation: u64) -> i32; + + /// Cancels a pending or completed-but-unread operation and releases host state. + /// + /// Cancellation must be idempotent when the operation is already absent. + pub(crate) fn cancel_operation(operation: u64) -> i32; + + /// Closes a host socket or listener handle. Closing an already-closed handle is valid. + pub(crate) fn close(handle: u64) -> i32; +} diff --git a/easytier-core/src/wasi/mod.rs b/easytier-core/src/wasi/mod.rs new file mode 100644 index 00000000..4377237c --- /dev/null +++ b/easytier-core/src/wasi/mod.rs @@ -0,0 +1,24 @@ +//! WASI runtime integration for EasyTier core. +//! +//! Portable Host capability seams live in [`crate::host`]. This module owns +//! the concrete WASI adapters, the `easytier_host` import contract, and the +//! guest lifecycle ABI used by an externally driven WASI runtime. + +pub mod abi; + +/// Concrete adapters a WASI guest uses to connect portable Host seams to the +/// [`abi::HOST_IMPORT_MODULE`] runtime contract. +#[cfg(target_os = "wasi")] +pub mod adapter; +#[cfg(target_os = "wasi")] +pub(crate) mod imports; +#[cfg(target_os = "wasi")] +pub(crate) mod runtime; +#[cfg(any(test, target_os = "wasi"))] +pub(crate) mod runtime_driver; +#[cfg(any(test, target_os = "wasi"))] +pub(crate) mod schema; +#[cfg(any(test, target_os = "wasi"))] +pub(crate) mod time; +#[cfg(any(test, target_os = "wasi"))] +pub(crate) mod wire; diff --git a/easytier-core/src/wasi/runtime.rs b/easytier-core/src/wasi/runtime.rs new file mode 100644 index 00000000..08a1aff5 --- /dev/null +++ b/easytier-core/src/wasi/runtime.rs @@ -0,0 +1,778 @@ +//! Runtime implementation and lifecycle exports for a WASI core instance. + +use crate::{ + config::toml::TomlConfig, connectivity::connector_host::HostConnectorEnvironmentSnapshot, +}; + +pub(super) type WasiCore = crate::instance::CoreInstance< + crate::connectivity::connector_host::ConnectorHost< + crate::wasi::adapter::socket::backend::WasiHostSocketBackend, + crate::wasi::adapter::environment::WasiHostConnectorEnvironmentIo, + >, +>; + +pub(super) struct WasiCoreRuntime { + socket_runtime: crate::host::socket::HostSocketRuntime, + core: std::sync::Arc, +} + +impl WasiCoreRuntime { + pub(super) fn core(&self) -> &std::sync::Arc { + &self.core + } + + pub(super) fn notify_host_completions(&self) { + self.socket_runtime.notify_completions(); + } +} + +pub(super) fn new_wasi_core_runtime( + config: TomlConfig, + process_runtime: std::sync::Arc, + environment_snapshot: HostConnectorEnvironmentSnapshot, + packet_sink: crate::host::packet::HostPacketSinkHandle, +) -> anyhow::Result { + use std::sync::Arc; + + use crate::host::{dns::HostDnsResolver, packet::HostPacketSink, socket::HostSocketRuntime}; + use crate::{ + connectivity::connector_host::new_connector_host, + instance::{CoreHostAdapters, CoreInstance}, + wasi::adapter::{ + dns::WasiHostDnsIo, environment::WasiHostConnectorEnvironmentIo, + packet::WasiHostPacketIo, socket::backend::WasiHostSocketBackend, + }, + }; + + let socket_runtime = HostSocketRuntime::new(); + let host = Arc::new(new_connector_host( + socket_runtime.clone(), + Arc::new(WasiHostSocketBackend::default()), + environment_snapshot, + Arc::new(WasiHostConnectorEnvironmentIo), + )); + let dns = Arc::new(HostDnsResolver::new( + socket_runtime.clone(), + Arc::new(WasiHostDnsIo), + )); + let packet_sink = Arc::new(HostPacketSink::new( + socket_runtime.clone(), + Arc::new(WasiHostPacketIo), + packet_sink, + )); + let adapters = CoreHostAdapters::new(host, dns, packet_sink, process_runtime); + let core = CoreInstance::from_toml(config, adapters)?; + + Ok(WasiCoreRuntime { + socket_runtime, + core, + }) +} + +mod abi { + use std::{ + cell::RefCell, + collections::BTreeMap, + sync::{Arc, Mutex}, + }; + + use tokio::{runtime::Builder, task::JoinHandle}; + + use crate::{ + config::toml::{ConfigLoader as _, TomlConfig}, + foundation::time::{clear_domain, enter_domain, next_deadline_millis}, + host::packet::HostPacketSinkHandle, + instance::{ + CoreInstanceState, + manager::{InstanceFactory, ManagedInstance}, + }, + process_runtime::{CoreProcessRuntime, ProtectedTcpPortLease}, + wasi::runtime_driver::{RuntimeDriveOutcome, RuntimeDriver}, + }; + + use super::{WasiCoreRuntime, new_wasi_core_runtime}; + use crate::wasi::schema::WasiCoreInstanceCreateConfig; + + #[cfg(feature = "proxy-smoltcp-stack")] + mod data_plane; + + const MAX_CREATE_CONFIG_LEN: usize = 16 * 1024 * 1024; + const MAX_GUEST_BUFFER_LEN: usize = MAX_CREATE_CONFIG_LEN; + const INVALID_HANDLE: i32 = -1; + const INVALID_STATE: i32 = -2; + const INVALID_INPUT: i32 = -3; + const ASYNC_ERROR: i32 = -4; + const BUSY: i32 = -5; + + struct WasiAbiState { + next_handle: u64, + handles: BTreeMap, + buffers: BTreeMap>, + active_instance: bool, + global_error: String, + } + + struct WasiHandleState { + instance_id: uuid::Uuid, + error: String, + } + + impl Default for WasiAbiState { + fn default() -> Self { + Self { + next_handle: 0, + handles: BTreeMap::new(), + buffers: BTreeMap::new(), + active_instance: false, + global_error: String::new(), + } + } + } + + struct WasiContext { + factory: WasiInstanceFactory, + instances: RefCell>>, + abi: RefCell, + } + + impl Default for WasiContext { + fn default() -> Self { + let factory = WasiInstanceFactory { + process_runtime: CoreProcessRuntime::new(), + }; + Self { + factory, + instances: RefCell::new(BTreeMap::new()), + abi: RefCell::new(WasiAbiState::default()), + } + } + } + + thread_local! { + static CONTEXT: WasiContext = WasiContext::default(); + } + + struct WasiInstanceFactory { + process_runtime: std::sync::Arc, + } + + struct WasiCreateContext { + domain: u64, + environment: crate::connectivity::connector_host::HostConnectorEnvironmentSnapshot, + packet_sink: HostPacketSinkHandle, + } + + struct WasiInstance { + instance_id: uuid::Uuid, + domain: u64, + core: WasiCoreRuntime, + execution: Mutex, + _protected_tcp_port_leases: Vec, + } + + struct WasiExecution { + runtime: tokio::runtime::Runtime, + runtime_driver: RuntimeDriver, + drive_again: bool, + start_task: Option>>, + stop_task: Option>, + } + + impl ManagedInstance for WasiInstance { + fn instance_id(&self) -> uuid::Uuid { + self.instance_id + } + } + + impl InstanceFactory for WasiInstanceFactory { + type Instance = WasiInstance; + type CreateContext = WasiCreateContext; + type Error = anyhow::Error; + + fn create( + &self, + config: TomlConfig, + context: Self::CreateContext, + ) -> Result, Self::Error> { + let instance_id = config.get_id(); + let protected_tcp_ports = context + .environment + .protected_tcp_ports + .iter() + .copied() + .map(|port| self.process_runtime.protect_tcp_port(port)) + .collect(); + let runtime_driver = RuntimeDriver::default(); + let park_driver = runtime_driver.clone(); + let runtime = Builder::new_current_thread() + .enable_time() + .on_thread_park(move || park_driver.on_thread_park()) + .build()?; + let core = { + let _domain = enter_domain(context.domain); + let _runtime = runtime.enter(); + new_wasi_core_runtime( + config, + self.process_runtime.clone(), + context.environment, + context.packet_sink, + )? + }; + + Ok(Arc::new(WasiInstance { + instance_id, + domain: context.domain, + core, + execution: Mutex::new(WasiExecution { + runtime, + runtime_driver, + drive_again: false, + start_task: None, + stop_task: None, + }), + _protected_tcp_port_leases: protected_tcp_ports, + })) + } + } + + impl WasiAbiState { + fn allocate_handle(&mut self) -> Result { + if self.active_instance { + self.set_global_error("core instance lifecycle call is not reentrant"); + return Err(BUSY); + } + loop { + self.next_handle = self.next_handle.wrapping_add(1); + if self.next_handle != 0 && !self.handles.contains_key(&self.next_handle) { + return Ok(self.next_handle); + } + } + } + + fn set_global_error(&mut self, error: impl ToString) { + self.global_error = error.to_string(); + } + + fn set_handle_error(&mut self, handle: u64, error: impl ToString) { + if let Some(state) = self.handles.get_mut(&handle) { + state.error = error.to_string(); + } else { + self.set_global_error(error); + } + } + + fn error_for_handle(&self, handle: u64) -> &str { + self.handles + .get(&handle) + .map(|state| state.error.as_str()) + .unwrap_or(self.global_error.as_str()) + } + + fn begin_instance_call(&mut self, handle: u64) -> Result { + if self.active_instance { + self.set_handle_error(handle, "core instance lifecycle call is not reentrant"); + return Err(BUSY); + } + let Some(instance_id) = self.handles.get(&handle).map(|state| state.instance_id) else { + self.set_global_error(format!("unknown core instance handle: {handle}")); + return Err(INVALID_HANDLE); + }; + self.active_instance = true; + Ok(instance_id) + } + + fn begin_instance_drop(&mut self, handle: u64) -> Result { + self.begin_instance_call(handle) + } + + fn finish_instance_call(&mut self) { + debug_assert!(self.active_instance); + self.active_instance = false; + } + + fn finish_instance_drop(&mut self, handle: u64) { + debug_assert!(self.active_instance); + self.handles.remove(&handle); + self.active_instance = false; + } + + fn read_buffer(&self, pointer: u32, length: usize) -> anyhow::Result> { + let buffer = self + .buffers + .get(&pointer) + .ok_or_else(|| anyhow::anyhow!("unknown guest buffer: {pointer}"))?; + if length > buffer.len() { + anyhow::bail!( + "guest buffer length {length} exceeds allocation {}", + buffer.len() + ); + } + Ok(buffer[..length].to_vec()) + } + } + + impl WasiInstance { + fn start(&self) -> anyhow::Result<()> { + let mut execution = self.execution.lock().unwrap(); + if execution.start_task.is_some() + || execution.stop_task.is_some() + || self.core.core().state() != CoreInstanceState::Created + { + anyhow::bail!("core instance cannot schedule start from its current state"); + } + let instance = self.core.core().clone(); + execution.start_task = Some( + execution + .runtime + .spawn(async move { instance.start().await }), + ); + Ok(()) + } + + fn stop(&self) { + let mut execution = self.execution.lock().unwrap(); + if execution.stop_task.is_some() + || self.core.core().state() == CoreInstanceState::Stopped + { + return; + } + let instance = self.core.core().clone(); + execution.stop_task = Some(execution.runtime.spawn(async move { + instance.stop().await; + })); + } + + fn drive(&self) -> anyhow::Result<()> { + let _domain = enter_domain(self.domain); + let mut execution = self.execution.lock().unwrap(); + execution.drive_again = execution.runtime_driver.drive(&execution.runtime) + == RuntimeDriveOutcome::BudgetExhausted; + + if execution + .start_task + .as_ref() + .is_some_and(JoinHandle::is_finished) + { + let task = execution.start_task.take().unwrap(); + match execution.runtime.block_on(task) { + Ok(Ok(())) => {} + Ok(Err(error)) => return Err(error), + Err(error) => anyhow::bail!("core instance start task failed: {error}"), + } + } + if execution + .stop_task + .as_ref() + .is_some_and(JoinHandle::is_finished) + { + let task = execution.stop_task.take().unwrap(); + if let Err(error) = execution.runtime.block_on(task) { + anyhow::bail!("core instance stop task failed: {error}"); + } + } + Ok(()) + } + + fn state_code(&self) -> i32 { + let execution = self.execution.lock().unwrap(); + if execution.stop_task.is_some() { + return 3; + } + if execution.start_task.is_some() { + return 1; + } + match self.core.core().state() { + CoreInstanceState::Created => 0, + CoreInstanceState::Starting => 1, + CoreInstanceState::Running => 2, + CoreInstanceState::Stopping => 3, + CoreInstanceState::Stopped => 4, + } + } + + fn next_wait_millis(&self) -> Option { + let execution = self.execution.lock().unwrap(); + if execution.drive_again { + Some(0) + } else { + next_deadline_millis(self.domain) + } + } + + fn send_packet(&self, packet: Vec) { + let packet_plane = self.core.core().packet_plane(); + self.execution.lock().unwrap().runtime.spawn(async move { + if let Err(error) = packet_plane.send_ip_packet(packet).await { + tracing::warn!(?error, "host packet ingress failed"); + } + }); + } + } + + fn decode_create_config(encoded: &[u8]) -> anyhow::Result { + if encoded.is_empty() || encoded.len() > MAX_CREATE_CONFIG_LEN { + anyhow::bail!("invalid host core instance config buffer"); + } + let config: WasiCoreInstanceCreateConfig = serde_json::from_slice(encoded)?; + config.validate()?; + Ok(config) + } + + fn with_instance( + handle: u64, + operation: impl FnOnce(&WasiInstance) -> anyhow::Result, + ) -> i32 { + let instance_id = + match CONTEXT.with(|context| context.abi.borrow_mut().begin_instance_call(handle)) { + Ok(instance_id) => instance_id, + Err(status) => return status, + }; + let instance = manager_get(instance_id); + let status = match instance { + Some(instance) => match operation(&instance) { + Ok(status) => status, + Err(error) => { + CONTEXT.with(|context| { + context.abi.borrow_mut().set_handle_error(handle, error); + }); + ASYNC_ERROR + } + }, + None => { + CONTEXT.with(|context| { + context.abi.borrow_mut().set_handle_error( + handle, + format!("core instance {instance_id} is not registered"), + ); + }); + INVALID_STATE + } + }; + CONTEXT.with(|context| context.abi.borrow_mut().finish_instance_call()); + status + } + + fn with_abi_state(operation: impl FnOnce(&WasiAbiState) -> T) -> T { + CONTEXT.with(|context| operation(&context.abi.borrow())) + } + + fn with_abi_state_mut(operation: impl FnOnce(&mut WasiAbiState) -> T) -> T { + CONTEXT.with(|context| operation(&mut context.abi.borrow_mut())) + } + + fn set_abi_error(error: impl ToString) { + with_abi_state_mut(|state| state.set_global_error(error)); + } + + fn set_instance_error(handle: u64, error: impl ToString) { + with_abi_state_mut(|state| state.set_handle_error(handle, error)); + } + + fn manager_get(instance_id: uuid::Uuid) -> Option> { + CONTEXT.with(|context| context.instances.borrow().get(&instance_id).cloned()) + } + + fn manager_remove(instance_id: uuid::Uuid) -> Option> { + CONTEXT.with(|context| context.instances.borrow_mut().remove(&instance_id)) + } + + fn manager_create( + config: TomlConfig, + context: WasiCreateContext, + ) -> Result< + std::sync::Arc, + crate::instance::manager::InstanceCreateError, + > { + CONTEXT.with(|wasi| { + let instance = wasi + .factory + .create(config, context) + .map_err(crate::instance::manager::InstanceCreateError::Factory)?; + let instance_id = instance.instance_id(); + let mut instances = wasi.instances.borrow_mut(); + match instances.entry(instance_id) { + std::collections::btree_map::Entry::Vacant(entry) => { + entry.insert(instance.clone()); + Ok(instance) + } + std::collections::btree_map::Entry::Occupied(_) => Err( + crate::instance::manager::InstanceCreateError::AlreadyExists { instance_id }, + ), + } + }) + } + + fn read_guest_buffer(pointer: u32, length: u32, maximum: usize) -> anyhow::Result> { + let length = usize::try_from(length).expect("u32 fits usize on wasm32"); + if pointer == 0 || length == 0 || length > maximum { + anyhow::bail!("invalid guest buffer reference"); + } + with_abi_state(|state| state.read_buffer(pointer, length)) + } + + #[unsafe(no_mangle)] + /// Allocates a guest-owned ABI buffer and returns its linear-memory offset. + /// + /// The runtime may write at most `length` bytes through wasm memory, then + /// pass the pointer to another lifecycle export and eventually free it. + pub extern "C" fn easytier_buffer_alloc(length: u32) -> u32 { + let length = usize::try_from(length).expect("u32 fits usize on wasm32"); + if length == 0 || length > MAX_GUEST_BUFFER_LEN { + set_abi_error("invalid guest buffer length"); + return 0; + } + let mut buffer = vec![0_u8; length].into_boxed_slice(); + let pointer = buffer.as_mut_ptr() as u32; + if pointer == 0 { + set_abi_error("guest buffer allocation failed"); + return 0; + } + with_abi_state_mut(|state| { + if state.buffers.contains_key(&pointer) { + state.set_global_error("guest buffer allocation collided with a live buffer"); + 0 + } else { + state.buffers.insert(pointer, buffer); + pointer + } + }) + } + + #[unsafe(no_mangle)] + /// Releases a buffer previously returned by [`easytier_buffer_alloc`]. + pub extern "C" fn easytier_buffer_free(pointer: u32) -> i32 { + with_abi_state_mut(|state| { + if state.buffers.remove(&pointer).is_some() { + 0 + } else { + state.set_global_error(format!("unknown guest buffer: {pointer}")); + INVALID_INPUT + } + }) + } + + #[unsafe(no_mangle)] + /// Creates one core instance from a versioned envelope containing TOML. + /// + /// `config_pointer` must name a live ABI buffer and `packet_sink_handle` + /// identifies the host sink used for locally delivered raw IP packets. + /// Returns zero on failure; retrieve the reason through the error exports. + pub extern "C" fn easytier_instance_create( + config_pointer: u32, + config_length: u32, + packet_sink_handle: u64, + ) -> u64 { + let encoded = match read_guest_buffer(config_pointer, config_length, MAX_CREATE_CONFIG_LEN) + { + Ok(encoded) => encoded, + Err(error) => { + set_abi_error(error); + return 0; + } + }; + let create_config = match decode_create_config(&encoded) { + Ok(config) => config, + Err(error) => { + set_abi_error(error); + return 0; + } + }; + let config = match create_config.parse_config() { + Ok(config) => config, + Err(error) => { + set_abi_error(error); + return 0; + } + }; + let handle = match with_abi_state_mut(WasiAbiState::allocate_handle) { + Ok(handle) => handle, + Err(_) => return 0, + }; + let instance = manager_create( + config, + WasiCreateContext { + domain: handle, + environment: create_config.environment, + packet_sink: HostPacketSinkHandle(packet_sink_handle), + }, + ); + let instance = match instance { + Ok(instance) => instance, + Err(error) => { + clear_domain(handle); + set_abi_error(error); + return 0; + } + }; + + with_abi_state_mut(|state| { + let previous = state.handles.insert( + handle, + WasiHandleState { + instance_id: instance.instance_id(), + error: String::new(), + }, + ); + debug_assert!(previous.is_none()); + handle + }) + } + + #[unsafe(no_mangle)] + /// Schedules instance startup. Completion is advanced by subsequent drive calls. + pub extern "C" fn easytier_instance_start(handle: u64) -> i32 { + with_instance(handle, |instance| match instance.start() { + Ok(()) => Ok(0), + Err(error) => { + set_instance_error(handle, error); + Ok(INVALID_STATE) + } + }) + } + + #[unsafe(no_mangle)] + /// Requests graceful instance shutdown. Completion is advanced by drive calls. + pub extern "C" fn easytier_instance_stop(handle: u64) -> i32 { + with_instance(handle, |instance| { + instance.stop(); + Ok(0) + }) + } + + #[unsafe(no_mangle)] + /// Runs one bounded turn of the instance's externally driven Tokio runtime. + /// + /// The return value is the current lifecycle state code, or a negative ABI + /// status on failure. Call after a timer deadline or host completion. + pub extern "C" fn easytier_instance_drive(handle: u64) -> i32 { + with_instance(handle, |instance| { + instance.drive()?; + Ok(instance.state_code()) + }) + } + + #[unsafe(no_mangle)] + /// Wakes tasks whose host I/O operation may have completed. + /// + /// The runtime calls this after finishing one or more `easytier_host` + /// operations; it does not itself consume a host completion. + pub extern "C" fn easytier_instance_notify_completions(handle: u64) -> i32 { + with_instance(handle, |instance| { + instance.core.notify_host_completions(); + Ok(0) + }) + } + + #[unsafe(no_mangle)] + /// Returns the current lifecycle state code without running the instance. + pub extern "C" fn easytier_instance_state(handle: u64) -> i32 { + with_instance(handle, |instance| Ok(instance.state_code())) + } + + /// Returns milliseconds until the next required drive, rounded up. A zero + /// means lifecycle work remains locally runnable. `i64::MAX` means core is + /// waiting only for a host completion; negative values are ABI status codes. + #[unsafe(no_mangle)] + pub extern "C" fn easytier_instance_next_deadline_millis(handle: u64) -> i64 { + let instance_id = match with_abi_state_mut(|state| state.begin_instance_call(handle)) { + Ok(instance_id) => instance_id, + Err(status) => return i64::from(status), + }; + let result = match manager_get(instance_id) { + Some(instance) => instance + .next_wait_millis() + .map(|millis| i64::try_from(millis).unwrap_or(i64::MAX)) + .unwrap_or(i64::MAX), + None => { + set_instance_error( + handle, + format!("core instance {instance_id} is not registered"), + ); + i64::from(INVALID_STATE) + } + }; + with_abi_state_mut(WasiAbiState::finish_instance_call); + result + } + + #[unsafe(no_mangle)] + /// Copies a raw IP packet from a guest ABI buffer into EasyTier ingress. + /// + /// Packet processing is asynchronous; success only means the ingress task + /// was scheduled. The caller retains and may free the source buffer once + /// this function returns. + pub extern "C" fn easytier_instance_send_packet( + handle: u64, + packet_pointer: u32, + packet_length: u32, + ) -> i32 { + with_instance(handle, |instance| { + let packet = match read_guest_buffer(packet_pointer, packet_length, 1024 * 1024) { + Ok(packet) => packet, + Err(error) => { + set_instance_error(handle, error); + return Ok(INVALID_INPUT); + } + }; + instance.send_packet(packet); + Ok(0) + }) + } + + #[unsafe(no_mangle)] + /// Destroys an instance and releases its lifecycle, timer, and runtime state. + pub extern "C" fn easytier_instance_drop(handle: u64) -> i32 { + let instance_id = match with_abi_state_mut(|state| state.begin_instance_drop(handle)) { + Ok(instance_id) => instance_id, + Err(status) => return status, + }; + let Some(instance) = manager_remove(instance_id) else { + set_instance_error( + handle, + format!("core instance {instance_id} is not registered"), + ); + with_abi_state_mut(WasiAbiState::finish_instance_call); + return INVALID_STATE; + }; + let domain = instance.domain; + { + let _domain = enter_domain(domain); + drop(instance); + } + clear_domain(domain); + with_abi_state_mut(|state| state.finish_instance_drop(handle)); + 0 + } + + #[unsafe(no_mangle)] + /// Returns the byte length of the most recent lifecycle error for `handle`. + pub extern "C" fn easytier_instance_error_len(handle: u64) -> u32 { + with_abi_state(|state| { + u32::try_from(state.error_for_handle(handle).len()).unwrap_or(u32::MAX) + }) + } + + #[unsafe(no_mangle)] + /// Copies the most recent lifecycle error into a caller-owned ABI buffer. + /// + /// `capacity` must contain the whole error; on success returns the copied + /// byte count, otherwise a negative ABI status code. + pub extern "C" fn easytier_instance_error_copy( + handle: u64, + destination: u32, + capacity: u32, + ) -> i32 { + with_abi_state_mut(|state| { + let error = state.error_for_handle(handle).as_bytes().to_vec(); + let capacity = usize::try_from(capacity).expect("u32 fits usize on wasm32"); + let Some(destination) = state.buffers.get_mut(&destination) else { + return INVALID_INPUT; + }; + if capacity > destination.len() || capacity < error.len() { + return INVALID_INPUT; + } + destination[..error.len()].copy_from_slice(&error); + i32::try_from(error.len()).unwrap_or(INVALID_INPUT) + }) + } +} diff --git a/easytier-core/src/wasi/runtime/abi/data_plane.rs b/easytier-core/src/wasi/runtime/abi/data_plane.rs new file mode 100644 index 00000000..02fde008 --- /dev/null +++ b/easytier-core/src/wasi/runtime/abi/data_plane.rs @@ -0,0 +1,737 @@ +//! Public data-plane guest exports backed by the instance operation broker. + +use std::{net::SocketAddr, time::Duration}; + +use crate::{ + gateway::{ + DataPlaneError, DataPlaneErrorKind, DataPlaneOperationId, DataPlaneOperationKind, + DataPlaneOperationResult, DataPlaneResourceId, DataPlaneSession, + }, + wasi::{ + abi::{ + DATA_PLANE_ABI_VERSION, DATA_PLANE_CAPABILITY, DATA_PLANE_TCP_CAPABILITY, + DATA_PLANE_UDP_CAPABILITY, + }, + wire::{ + data_plane::{ + COMPLETION_LEN, TCP_ACCEPT_RESULT_LEN, TCP_BIND_RESULT_LEN, TCP_CONNECT_RESULT_LEN, + TCP_READ_METADATA_LEN, UDP_BIND_RESULT_LEN, UDP_RECEIVE_METADATA_LEN, + decode_ipv4_socket_address, encode_completion, encode_resource_and_address, + encode_stream_addresses, encode_tcp_read_metadata, encode_udp_receive_metadata, + error_status, normalize_call_status, + }, + socket::SOCKET_ADDRESS_LEN, + }, + }, +}; + +use super::{WasiInstance, set_instance_error, with_abi_state, with_abi_state_mut, with_instance}; + +const OPERATION_ID_LEN: usize = 8; +const MAX_WRITE_LEN: usize = 1024 * 1024; + +type WasiDataPlaneSession = DataPlaneSession< + crate::connectivity::connector_host::ConnectorHost< + crate::wasi::adapter::socket::backend::WasiHostSocketBackend, + crate::wasi::adapter::environment::WasiHostConnectorEnvironmentIo, + >, +>; + +impl WasiInstance { + fn data_plane_session(&self) -> std::sync::Arc { + self.core.core().data_plane_session() + } + + fn submit_data_plane( + &self, + submit: impl FnOnce( + &std::sync::Arc, + ) -> Result, + ) -> Result { + let execution = self.execution.lock().unwrap(); + let _domain = crate::foundation::time::enter_domain(self.domain); + let _runtime = execution.runtime.enter(); + submit(&self.data_plane_session()) + } +} + +fn error(kind: DataPlaneErrorKind, message: impl Into) -> DataPlaneError { + DataPlaneError::new(kind, message) +} + +fn invalid_input(message: impl Into) -> DataPlaneError { + error(DataPlaneErrorKind::Io, message) +} + +fn timeout(timeout_ms: u64) -> Option { + (timeout_ms != u64::MAX).then(|| Duration::from_millis(timeout_ms)) +} + +fn operation_id(raw: u64) -> Result { + DataPlaneOperationId::from_raw(raw) + .ok_or_else(|| error(DataPlaneErrorKind::HandleClosed, "invalid operation ID")) +} + +fn resource_id(raw: u64) -> Result { + DataPlaneResourceId::from_raw(raw) + .ok_or_else(|| error(DataPlaneErrorKind::HandleClosed, "invalid resource ID")) +} + +fn validate_local_port(raw: u32) -> Result { + u16::try_from(raw).map_err(|_| invalid_input(format!("invalid local port {raw}"))) +} + +fn data_plane_call( + handle: u64, + operation: impl FnOnce(&WasiInstance) -> Result, +) -> i32 { + let mut entered = false; + let status = with_instance(handle, |instance| { + entered = true; + Ok(match operation(instance) { + Ok(status) => status, + Err(error) => { + let status = error_status(error.kind()); + set_instance_error(handle, error.message()); + status + } + }) + }); + normalize_call_status(entered, status) +} + +fn validate_output(pointer: u32, capacity: usize, required: usize) -> Result<(), DataPlaneError> { + if required == 0 && capacity == 0 && pointer == 0 { + return Ok(()); + } + if pointer == 0 { + return Err(invalid_input("guest output buffer pointer is zero")); + } + with_abi_state(|state| { + let buffer = state + .buffers + .get(&pointer) + .ok_or_else(|| invalid_input(format!("unknown guest buffer: {pointer}")))?; + if capacity > buffer.len() { + return Err(invalid_input(format!( + "guest output capacity {capacity} exceeds allocation {}", + buffer.len() + ))); + } + if required > capacity { + return Err(error( + DataPlaneErrorKind::BufferTooSmall, + format!("result requires {required} bytes, buffer has {capacity}"), + )); + } + Ok(()) + }) +} + +fn validate_fixed_output(pointer: u32, required: usize) -> Result<(), DataPlaneError> { + validate_output(pointer, required, required) +} + +fn write_output(pointer: u32, bytes: &[u8]) -> Result<(), DataPlaneError> { + if bytes.is_empty() && pointer == 0 { + return Ok(()); + } + with_abi_state_mut(|state| { + let buffer = state + .buffers + .get_mut(&pointer) + .ok_or_else(|| invalid_input(format!("unknown guest buffer: {pointer}")))?; + if bytes.len() > buffer.len() { + return Err(error( + DataPlaneErrorKind::BufferTooSmall, + format!( + "result requires {} bytes, allocation has {}", + bytes.len(), + buffer.len() + ), + )); + } + buffer[..bytes.len()].copy_from_slice(bytes); + Ok(()) + }) +} + +fn read_input(pointer: u32, length: u32, maximum: usize) -> Result, DataPlaneError> { + let length = usize::try_from(length).expect("u32 fits usize on wasm32"); + if length == 0 { + return Ok(Vec::new()); + } + if pointer == 0 || length > maximum { + return Err(invalid_input("invalid guest input buffer reference")); + } + with_abi_state(|state| { + state + .read_buffer(pointer, length) + .map_err(DataPlaneError::from) + }) +} + +fn read_ipv4_address(pointer: u32) -> Result { + let encoded = read_input(pointer, SOCKET_ADDRESS_LEN as u32, SOCKET_ADDRESS_LEN)?; + decode_ipv4_socket_address(&encoded).map_err(DataPlaneError::from) +} + +fn write_operation_id(pointer: u32, operation: DataPlaneOperationId) -> Result<(), DataPlaneError> { + write_output(pointer, &operation.get().to_be_bytes()) +} + +fn submit_operation( + handle: u64, + output: u32, + submit: impl FnOnce(&WasiInstance) -> Result, +) -> i32 { + data_plane_call(handle, |instance| { + validate_fixed_output(output, OPERATION_ID_LEN)?; + let operation = submit(instance)?; + if let Err(error) = write_operation_id(output, operation) { + instance.data_plane_session().free_operation(operation); + return Err(error); + } + Ok(0) + }) +} + +fn take_result( + session: &WasiDataPlaneSession, + operation: DataPlaneOperationId, + expected: DataPlaneOperationKind, + take: impl FnOnce(&DataPlaneOperationResult) -> Result, +) -> Result { + let actual = session.operation_kind(operation)?; + if actual != expected { + return Err(invalid_input(format!( + "operation kind mismatch: expected {expected:?}, got {actual:?}" + ))); + } + + let mut extraction_error = None; + let result = session.take_result_with(operation, |outcome| { + Some(match outcome { + Ok(result) => match take(result) { + Ok(value) => Ok(value), + Err(error) => { + extraction_error = Some(error); + return None; + } + }, + Err(kind) => Err(error( + *kind, + format!("data-plane operation failed with {kind:?}"), + )), + }) + })?; + if let Some(error) = extraction_error { + return Err(error); + } + result.ok_or_else(|| invalid_input("data-plane result could not be consumed"))? +} + +fn require_ipv4(address: SocketAddr) -> Result { + address.is_ipv4().then_some(address).ok_or_else(|| { + error( + DataPlaneErrorKind::AddressFamilyUnsupported, + "data-plane ABI v2 supports IPv4 only", + ) + }) +} + +#[unsafe(no_mangle)] +pub extern "C" fn easytier_data_plane_abi_version() -> u32 { + DATA_PLANE_ABI_VERSION +} + +#[unsafe(no_mangle)] +pub extern "C" fn easytier_data_plane_capabilities() -> u64 { + DATA_PLANE_CAPABILITY | DATA_PLANE_TCP_CAPABILITY | DATA_PLANE_UDP_CAPABILITY +} + +#[unsafe(no_mangle)] +pub extern "C" fn easytier_data_plane_tcp_connect_submit( + handle: u64, + peer_address: u32, + timeout_ms: u64, + output_operation: u32, +) -> i32 { + let peer_address = match read_ipv4_address(peer_address) { + Ok(address) => address, + Err(error) => { + set_instance_error(handle, error.message()); + return error_status(error.kind()); + } + }; + submit_operation(handle, output_operation, |instance| { + instance.submit_data_plane(|session| { + session.submit_tcp_connect(peer_address, timeout(timeout_ms)) + }) + }) +} + +#[unsafe(no_mangle)] +pub extern "C" fn easytier_data_plane_tcp_bind_submit( + handle: u64, + local_port: u32, + timeout_ms: u64, + output_operation: u32, +) -> i32 { + let local_port = match validate_local_port(local_port) { + Ok(port) => port, + Err(error) => { + set_instance_error(handle, error.message()); + return error_status(error.kind()); + } + }; + submit_operation(handle, output_operation, |instance| { + instance + .submit_data_plane(|session| session.submit_tcp_bind(local_port, timeout(timeout_ms))) + }) +} + +#[unsafe(no_mangle)] +pub extern "C" fn easytier_data_plane_tcp_accept_submit( + handle: u64, + listener: u64, + timeout_ms: u64, + output_operation: u32, +) -> i32 { + let listener = match resource_id(listener) { + Ok(listener) => listener, + Err(error) => { + set_instance_error(handle, error.message()); + return error_status(error.kind()); + } + }; + submit_operation(handle, output_operation, |instance| { + instance + .submit_data_plane(|session| session.submit_tcp_accept(listener, timeout(timeout_ms))) + }) +} + +#[unsafe(no_mangle)] +pub extern "C" fn easytier_data_plane_tcp_read_submit( + handle: u64, + stream: u64, + max_len: u32, + timeout_ms: u64, + output_operation: u32, +) -> i32 { + let stream = match resource_id(stream) { + Ok(stream) => stream, + Err(error) => { + set_instance_error(handle, error.message()); + return error_status(error.kind()); + } + }; + submit_operation(handle, output_operation, |instance| { + instance.submit_data_plane(|session| { + session.submit_tcp_read(stream, max_len as usize, timeout(timeout_ms)) + }) + }) +} + +#[unsafe(no_mangle)] +pub extern "C" fn easytier_data_plane_tcp_write_submit( + handle: u64, + stream: u64, + data_pointer: u32, + data_length: u32, + timeout_ms: u64, + output_operation: u32, +) -> i32 { + let stream = match resource_id(stream) { + Ok(stream) => stream, + Err(error) => { + set_instance_error(handle, error.message()); + return error_status(error.kind()); + } + }; + let data = match read_input(data_pointer, data_length, MAX_WRITE_LEN) { + Ok(data) => data, + Err(error) => { + set_instance_error(handle, error.message()); + return error_status(error.kind()); + } + }; + submit_operation(handle, output_operation, |instance| { + instance.submit_data_plane(|session| { + session.submit_tcp_write(stream, data, timeout(timeout_ms)) + }) + }) +} + +#[unsafe(no_mangle)] +pub extern "C" fn easytier_data_plane_udp_bind_submit( + handle: u64, + local_port: u32, + timeout_ms: u64, + output_operation: u32, +) -> i32 { + let local_port = match validate_local_port(local_port) { + Ok(port) => port, + Err(error) => { + set_instance_error(handle, error.message()); + return error_status(error.kind()); + } + }; + submit_operation(handle, output_operation, |instance| { + instance + .submit_data_plane(|session| session.submit_udp_bind(local_port, timeout(timeout_ms))) + }) +} + +#[unsafe(no_mangle)] +pub extern "C" fn easytier_data_plane_udp_receive_submit( + handle: u64, + socket: u64, + max_len: u32, + timeout_ms: u64, + output_operation: u32, +) -> i32 { + let socket = match resource_id(socket) { + Ok(socket) => socket, + Err(error) => { + set_instance_error(handle, error.message()); + return error_status(error.kind()); + } + }; + submit_operation(handle, output_operation, |instance| { + instance.submit_data_plane(|session| { + session.submit_udp_receive(socket, max_len as usize, timeout(timeout_ms)) + }) + }) +} + +#[unsafe(no_mangle)] +pub extern "C" fn easytier_data_plane_udp_send_submit( + handle: u64, + socket: u64, + peer_address: u32, + data_pointer: u32, + data_length: u32, + timeout_ms: u64, + output_operation: u32, +) -> i32 { + let socket = match resource_id(socket) { + Ok(socket) => socket, + Err(error) => { + set_instance_error(handle, error.message()); + return error_status(error.kind()); + } + }; + let peer_address = match read_ipv4_address(peer_address) { + Ok(address) => address, + Err(error) => { + set_instance_error(handle, error.message()); + return error_status(error.kind()); + } + }; + let data = match read_input(data_pointer, data_length, MAX_WRITE_LEN) { + Ok(data) => data, + Err(error) => { + set_instance_error(handle, error.message()); + return error_status(error.kind()); + } + }; + submit_operation(handle, output_operation, |instance| { + instance.submit_data_plane(|session| { + session.submit_udp_send(socket, peer_address, data, timeout(timeout_ms)) + }) + }) +} + +#[unsafe(no_mangle)] +pub extern "C" fn easytier_data_plane_operation_cancel(handle: u64, operation: u64) -> i32 { + data_plane_call(handle, |instance| { + instance + .data_plane_session() + .cancel_operation(operation_id(operation)?); + Ok(0) + }) +} + +#[unsafe(no_mangle)] +pub extern "C" fn easytier_data_plane_operation_free(handle: u64, operation: u64) -> i32 { + data_plane_call(handle, |instance| { + instance + .data_plane_session() + .free_operation(operation_id(operation)?); + Ok(0) + }) +} + +#[unsafe(no_mangle)] +pub extern "C" fn easytier_data_plane_resource_close(handle: u64, resource: u64) -> i32 { + data_plane_call(handle, |instance| { + instance + .data_plane_session() + .close_resource(resource_id(resource)?); + Ok(0) + }) +} + +#[unsafe(no_mangle)] +pub extern "C" fn easytier_data_plane_completion_drain( + handle: u64, + output: u32, + capacity: u32, +) -> i32 { + data_plane_call(handle, |instance| { + let capacity = usize::try_from(capacity).expect("u32 fits usize on wasm32"); + let required = capacity + .checked_mul(COMPLETION_LEN) + .ok_or_else(|| invalid_input("completion output size overflow"))?; + validate_output(output, required, required)?; + let completions = instance.data_plane_session().drain_completions(capacity); + let mut encoded = Vec::with_capacity(completions.len() * COMPLETION_LEN); + for completion in completions { + encoded.extend_from_slice(&encode_completion(completion)); + } + write_output(output, &encoded)?; + i32::try_from(encoded.len() / COMPLETION_LEN) + .map_err(|_| invalid_input("completion count exceeds i32")) + }) +} + +#[unsafe(no_mangle)] +pub extern "C" fn easytier_data_plane_result_size(handle: u64, operation: u64) -> i32 { + data_plane_call(handle, |instance| { + let size = instance + .data_plane_session() + .result_payload_bytes(operation_id(operation)?)?; + i32::try_from(size).map_err(|_| invalid_input("data-plane result size exceeds i32")) + }) +} + +#[unsafe(no_mangle)] +pub extern "C" fn easytier_data_plane_tcp_connect_result_take( + handle: u64, + operation: u64, + output: u32, +) -> i32 { + data_plane_call(handle, |instance| { + validate_fixed_output(output, TCP_CONNECT_RESULT_LEN)?; + let session = instance.data_plane_session(); + let wire = take_result( + &session, + operation_id(operation)?, + DataPlaneOperationKind::TcpConnect, + |result| match result { + DataPlaneOperationResult::TcpConnected { + stream, + local_addr, + peer_addr, + } => Ok(encode_stream_addresses( + stream.get(), + require_ipv4(*local_addr)?, + require_ipv4(*peer_addr)?, + )), + _ => Err(invalid_input("TCP connect result variant mismatch")), + }, + )?; + write_output(output, &wire)?; + Ok(0) + }) +} + +#[unsafe(no_mangle)] +pub extern "C" fn easytier_data_plane_tcp_bind_result_take( + handle: u64, + operation: u64, + output: u32, +) -> i32 { + data_plane_call(handle, |instance| { + validate_fixed_output(output, TCP_BIND_RESULT_LEN)?; + let session = instance.data_plane_session(); + let wire = take_result( + &session, + operation_id(operation)?, + DataPlaneOperationKind::TcpBind, + |result| match result { + DataPlaneOperationResult::TcpBound { + listener, + local_addr, + } => Ok(encode_resource_and_address( + listener.get(), + require_ipv4(*local_addr)?, + )), + _ => Err(invalid_input("TCP bind result variant mismatch")), + }, + )?; + write_output(output, &wire)?; + Ok(0) + }) +} + +#[unsafe(no_mangle)] +pub extern "C" fn easytier_data_plane_tcp_accept_result_take( + handle: u64, + operation: u64, + output: u32, +) -> i32 { + data_plane_call(handle, |instance| { + validate_fixed_output(output, TCP_ACCEPT_RESULT_LEN)?; + let session = instance.data_plane_session(); + let wire = take_result( + &session, + operation_id(operation)?, + DataPlaneOperationKind::TcpAccept, + |result| match result { + DataPlaneOperationResult::TcpAccepted { + stream, + local_addr, + peer_addr, + } => Ok(encode_stream_addresses( + stream.get(), + require_ipv4(*local_addr)?, + require_ipv4(*peer_addr)?, + )), + _ => Err(invalid_input("TCP accept result variant mismatch")), + }, + )?; + write_output(output, &wire)?; + Ok(0) + }) +} + +#[unsafe(no_mangle)] +pub extern "C" fn easytier_data_plane_tcp_read_result_take( + handle: u64, + operation: u64, + data_output: u32, + data_capacity: u32, + metadata_output: u32, +) -> i32 { + data_plane_call(handle, |instance| { + let session = instance.data_plane_session(); + let operation = operation_id(operation)?; + let required = session.result_payload_bytes(operation)?; + validate_output(data_output, data_capacity as usize, required)?; + validate_fixed_output(metadata_output, TCP_READ_METADATA_LEN)?; + take_result( + &session, + operation, + DataPlaneOperationKind::TcpRead, + |result| match result { + DataPlaneOperationResult::TcpRead { data, eof } => { + write_output(data_output, data)?; + write_output(metadata_output, &encode_tcp_read_metadata(*eof))?; + i32::try_from(data.len()) + .map_err(|_| invalid_input("TCP read result exceeds i32")) + } + _ => Err(invalid_input("TCP read result variant mismatch")), + }, + ) + }) +} + +#[unsafe(no_mangle)] +pub extern "C" fn easytier_data_plane_tcp_write_result_take(handle: u64, operation: u64) -> i32 { + data_plane_call(handle, |instance| { + let session = instance.data_plane_session(); + let len = take_result( + &session, + operation_id(operation)?, + DataPlaneOperationKind::TcpWrite, + |result| match result { + DataPlaneOperationResult::TcpWritten { len } => Ok(*len), + _ => Err(invalid_input("TCP write result variant mismatch")), + }, + )?; + i32::try_from(len).map_err(|_| invalid_input("TCP write result exceeds i32")) + }) +} + +#[unsafe(no_mangle)] +pub extern "C" fn easytier_data_plane_udp_bind_result_take( + handle: u64, + operation: u64, + output: u32, +) -> i32 { + data_plane_call(handle, |instance| { + validate_fixed_output(output, UDP_BIND_RESULT_LEN)?; + let session = instance.data_plane_session(); + let wire = take_result( + &session, + operation_id(operation)?, + DataPlaneOperationKind::UdpBind, + |result| match result { + DataPlaneOperationResult::UdpBound { socket, local_addr } => Ok( + encode_resource_and_address(socket.get(), require_ipv4(*local_addr)?), + ), + _ => Err(invalid_input("UDP bind result variant mismatch")), + }, + )?; + write_output(output, &wire)?; + Ok(0) + }) +} + +#[unsafe(no_mangle)] +pub extern "C" fn easytier_data_plane_udp_receive_result_take( + handle: u64, + operation: u64, + data_output: u32, + data_capacity: u32, + metadata_output: u32, +) -> i32 { + data_plane_call(handle, |instance| { + let session = instance.data_plane_session(); + let operation = operation_id(operation)?; + let required = session.result_payload_bytes(operation)?; + validate_output(data_output, data_capacity as usize, required)?; + validate_fixed_output(metadata_output, UDP_RECEIVE_METADATA_LEN)?; + take_result( + &session, + operation, + DataPlaneOperationKind::UdpReceive, + |result| match result { + DataPlaneOperationResult::UdpReceived { + data, + peer_addr, + truncated, + } => { + write_output(data_output, data)?; + write_output( + metadata_output, + &encode_udp_receive_metadata(require_ipv4(*peer_addr)?, *truncated), + )?; + i32::try_from(data.len()) + .map_err(|_| invalid_input("UDP receive result exceeds i32")) + } + _ => Err(invalid_input("UDP receive result variant mismatch")), + }, + ) + }) +} + +#[unsafe(no_mangle)] +pub extern "C" fn easytier_data_plane_udp_send_result_take(handle: u64, operation: u64) -> i32 { + data_plane_call(handle, |instance| { + let session = instance.data_plane_session(); + let len = take_result( + &session, + operation_id(operation)?, + DataPlaneOperationKind::UdpSend, + |result| match result { + DataPlaneOperationResult::UdpSent { len } => Ok(*len), + _ => Err(invalid_input("UDP send result variant mismatch")), + }, + )?; + i32::try_from(len).map_err(|_| invalid_input("UDP send result exceeds i32")) + }) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn infinite_timeout_sentinel_is_distinct_from_zero() { + assert_eq!(timeout(u64::MAX), None); + assert_eq!(timeout(0), Some(Duration::ZERO)); + } +} diff --git a/easytier-core/src/wasi/runtime_driver.rs b/easytier-core/src/wasi/runtime_driver.rs new file mode 100644 index 00000000..1c821b7a --- /dev/null +++ b/easytier-core/src/wasi/runtime_driver.rs @@ -0,0 +1,151 @@ +//! Bounded current-thread Tokio turns for externally driven runtimes. + +use std::{ + future::{Future, poll_fn}, + sync::{ + Arc, Mutex, + atomic::{AtomicBool, Ordering}, + }, + task::{Poll, Waker}, + time::Duration, +}; + +use tokio::runtime::Runtime; + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub(super) enum RuntimeDriveOutcome { + Quiescent, + BudgetExhausted, +} + +#[derive(Clone, Default)] +pub(super) struct RuntimeDriver { + state: Arc, +} + +#[derive(Default)] +struct RuntimeDriverState { + active: AtomicBool, + quiescent: AtomicBool, + waker: Mutex>, +} + +impl RuntimeDriver { + pub(super) fn on_thread_park(&self) { + if !self.state.active.load(Ordering::SeqCst) { + return; + } + self.state.quiescent.store(true, Ordering::SeqCst); + if let Some(waker) = self.state.waker.lock().unwrap().take() { + waker.wake(); + } + } + + pub(super) fn drive(&self, runtime: &Runtime) -> RuntimeDriveOutcome { + // First give the timer driver a non-blocking turn. The quiescence hook + // stays disabled here so an expired timer can wake its task. + runtime.block_on(async { + tokio::time::sleep(Duration::ZERO).await; + }); + + let _active = RuntimeDriverGuard::activate(self.state.as_ref()); + runtime.block_on(async { + let budget = tokio::time::sleep(Duration::ZERO); + tokio::pin!(budget); + poll_fn(|context| { + if self.state.poll_quiescent(context.waker()) { + return Poll::Ready(RuntimeDriveOutcome::Quiescent); + } + if budget.as_mut().poll(context).is_ready() { + return Poll::Ready(RuntimeDriveOutcome::BudgetExhausted); + } + Poll::Pending + }) + .await + }) + } +} + +impl RuntimeDriverState { + fn poll_quiescent(&self, waker: &Waker) -> bool { + if self.quiescent.load(Ordering::SeqCst) { + return true; + } + *self.waker.lock().unwrap() = Some(waker.clone()); + self.quiescent.load(Ordering::SeqCst) + } +} + +struct RuntimeDriverGuard<'a> { + state: &'a RuntimeDriverState, +} + +impl<'a> RuntimeDriverGuard<'a> { + fn activate(state: &'a RuntimeDriverState) -> Self { + state.quiescent.store(false, Ordering::SeqCst); + *state.waker.lock().unwrap() = None; + state.active.store(true, Ordering::SeqCst); + Self { state } + } +} + +impl Drop for RuntimeDriverGuard<'_> { + fn drop(&mut self) { + self.state.active.store(false, Ordering::SeqCst); + self.state.quiescent.store(false, Ordering::SeqCst); + *self.state.waker.lock().unwrap() = None; + } +} + +#[cfg(test)] +mod tests { + use std::{future::poll_fn, sync::Arc, task::Poll}; + + use tokio::{runtime::Builder, sync::Notify}; + + use super::{RuntimeDriveOutcome, RuntimeDriver}; + + fn runtime(driver: &RuntimeDriver) -> tokio::runtime::Runtime { + let park_driver = driver.clone(); + Builder::new_current_thread() + .enable_time() + .event_interval(3) + .on_thread_park(move || park_driver.on_thread_park()) + .build() + .unwrap() + } + + #[test] + fn reports_budget_exhaustion_for_a_continuously_runnable_task() { + let driver = RuntimeDriver::default(); + let runtime = runtime(&driver); + let task = runtime.spawn(poll_fn(|context| { + context.waker().wake_by_ref(); + Poll::<()>::Pending + })); + + assert_eq!(driver.drive(&runtime), RuntimeDriveOutcome::BudgetExhausted); + + task.abort(); + while driver.drive(&runtime) == RuntimeDriveOutcome::BudgetExhausted {} + assert!(task.is_finished()); + } + + #[test] + fn reports_quiescence_while_waiting_for_an_external_wake() { + let driver = RuntimeDriver::default(); + let runtime = runtime(&driver); + let notify = Arc::new(Notify::new()); + let task_notify = notify.clone(); + let task = runtime.spawn(async move { + task_notify.notified().await; + }); + + assert_eq!(driver.drive(&runtime), RuntimeDriveOutcome::Quiescent); + assert!(!task.is_finished()); + + notify.notify_one(); + while driver.drive(&runtime) == RuntimeDriveOutcome::BudgetExhausted {} + assert!(task.is_finished()); + } +} diff --git a/easytier-core/src/wasi/schema.rs b/easytier-core/src/wasi/schema.rs new file mode 100644 index 00000000..b9fa8329 --- /dev/null +++ b/easytier-core/src/wasi/schema.rs @@ -0,0 +1,34 @@ +//! Versioned, serialized inputs accepted by the WASI instance lifecycle ABI. + +use serde::{Deserialize, Serialize}; + +use crate::{ + config::toml::TomlConfig, connectivity::connector_host::HostConnectorEnvironmentSnapshot, +}; + +pub(crate) const WASI_CORE_INSTANCE_CONFIG_VERSION: u32 = + crate::wasi::abi::CORE_INSTANCE_CONFIG_VERSION; + +/// Versioned payload accepted by host-driven instance frontends. +#[derive(Debug, Clone, Serialize, Deserialize)] +pub(crate) struct WasiCoreInstanceCreateConfig { + pub version: u32, + pub config: String, + pub environment: HostConnectorEnvironmentSnapshot, +} + +impl WasiCoreInstanceCreateConfig { + pub fn validate(&self) -> anyhow::Result<()> { + if self.version != WASI_CORE_INSTANCE_CONFIG_VERSION { + anyhow::bail!( + "unsupported host core instance config version: {}", + self.version + ); + } + Ok(()) + } + + pub fn parse_config(&self) -> anyhow::Result { + TomlConfig::new_from_str_with_source("WASI create config", &self.config) + } +} diff --git a/easytier-core/src/wasi/time.rs b/easytier-core/src/wasi/time.rs new file mode 100644 index 00000000..7cf1d899 --- /dev/null +++ b/easytier-core/src/wasi/time.rs @@ -0,0 +1,318 @@ +//! Deadline-tracking Tokio time implementation for externally driven WASI runtimes. + +mod tracked { + use std::{ + cell::{Cell, RefCell}, + collections::BTreeMap, + future::{Future, IntoFuture}, + pin::Pin, + task::{Context, Poll}, + }; + + pub use tokio::time::{Duration, Instant, MissedTickBehavior, error}; + + thread_local! { + static CURRENT_DOMAIN: Cell> = const { Cell::new(None) }; + static DEADLINES: RefCell = RefCell::new(DeadlineRegistry::default()); + } + + #[derive(Default)] + struct DeadlineRegistry { + next_token: u64, + entries: BTreeMap, + } + + impl DeadlineRegistry { + fn insert(&mut self, domain: u64, deadline: Instant) -> u64 { + loop { + self.next_token = self.next_token.wrapping_add(1); + if self.next_token != 0 && !self.entries.contains_key(&self.next_token) { + self.entries.insert(self.next_token, (domain, deadline)); + return self.next_token; + } + } + } + } + + pub(crate) struct TimerDomainGuard(Option); + + impl Drop for TimerDomainGuard { + fn drop(&mut self) { + CURRENT_DOMAIN.set(self.0); + } + } + + pub(crate) fn enter_domain(domain: u64) -> TimerDomainGuard { + TimerDomainGuard(CURRENT_DOMAIN.replace(Some(domain))) + } + + pub(crate) fn clear_domain(domain: u64) { + DEADLINES.with_borrow_mut(|registry| { + registry + .entries + .retain(|_, (entry_domain, _)| *entry_domain != domain); + }); + } + + pub(crate) fn next_deadline_millis(domain: u64) -> Option { + let deadline = DEADLINES.with_borrow(|registry| { + registry + .entries + .values() + .filter_map(|(entry_domain, deadline)| { + (*entry_domain == domain).then_some(*deadline) + }) + .min() + })?; + let duration = deadline.saturating_duration_since(Instant::now()); + let nanos = duration.as_nanos(); + Some(u64::try_from(nanos.div_ceil(1_000_000)).unwrap_or(u64::MAX)) + } + + struct Registration { + token: Option, + deadline: Instant, + } + + impl Registration { + fn new(deadline: Instant) -> Self { + let mut registration = Self { + token: None, + deadline, + }; + registration.ensure(); + registration + } + + fn ensure(&mut self) { + if self.token.is_some() { + return; + } + let Some(domain) = CURRENT_DOMAIN.get() else { + return; + }; + self.token = + Some(DEADLINES.with_borrow_mut(|registry| registry.insert(domain, self.deadline))); + } + + fn reset(&mut self, deadline: Instant) { + self.remove(); + self.deadline = deadline; + self.ensure(); + } + + fn remove(&mut self) { + if let Some(token) = self.token.take() { + DEADLINES.with_borrow_mut(|registry| { + registry.entries.remove(&token); + }); + } + } + } + + impl Drop for Registration { + fn drop(&mut self) { + self.remove(); + } + } + + pub struct Sleep { + inner: Pin>, + registration: Registration, + } + + impl Sleep { + pub fn reset(mut self: Pin<&mut Self>, deadline: Instant) { + let this = self.as_mut().get_mut(); + this.inner.as_mut().reset(deadline); + this.registration.reset(deadline); + } + } + + impl Future for Sleep { + type Output = (); + + fn poll(mut self: Pin<&mut Self>, context: &mut Context<'_>) -> Poll { + let this = self.as_mut().get_mut(); + this.registration.ensure(); + let result = this.inner.as_mut().poll(context); + if result.is_ready() { + this.registration.remove(); + } + result + } + } + + pub fn sleep(duration: Duration) -> Sleep { + sleep_until(Instant::now() + duration) + } + + pub fn sleep_until(deadline: Instant) -> Sleep { + Sleep { + inner: Box::pin(tokio::time::sleep_until(deadline)), + registration: Registration::new(deadline), + } + } + + pub struct Interval { + inner: Pin>, + period: Duration, + next_deadline: Instant, + missed_tick_behavior: MissedTickBehavior, + registration: Registration, + } + + impl Interval { + pub async fn tick(&mut self) -> Instant { + self.registration.ensure(); + let tick = self.next_deadline; + self.inner.as_mut().await; + let now = Instant::now(); + self.next_deadline = if now > tick + Duration::from_millis(5) { + next_interval_deadline(self.missed_tick_behavior, tick, now, self.period) + } else { + tick + self.period + }; + self.inner.as_mut().reset(self.next_deadline); + self.registration.reset(self.next_deadline); + tick + } + } + + pub(super) fn next_interval_deadline( + behavior: MissedTickBehavior, + tick: Instant, + now: Instant, + period: Duration, + ) -> Instant { + match behavior { + MissedTickBehavior::Burst => tick + period, + MissedTickBehavior::Delay => now + period, + MissedTickBehavior::Skip => { + now + period + - Duration::from_nanos( + ((now - tick).as_nanos() % period.as_nanos()) + .try_into() + .expect("too much time has elapsed since the interval tick"), + ) + } + } + } + + pub fn interval(period: Duration) -> Interval { + interval_at(Instant::now(), period) + } + + pub fn interval_at(start: Instant, period: Duration) -> Interval { + assert!(period > Duration::ZERO, "`period` must be non-zero."); + Interval { + inner: Box::pin(tokio::time::sleep_until(start)), + period, + next_deadline: start, + missed_tick_behavior: MissedTickBehavior::Burst, + registration: Registration::new(start), + } + } + + pub async fn timeout(duration: Duration, future: F) -> Result + where + F: IntoFuture, + { + timeout_at(Instant::now() + duration, future).await + } + + pub async fn timeout_at(deadline: Instant, future: F) -> Result + where + F: IntoFuture, + { + let _registration = Registration::new(deadline); + tokio::time::timeout_at(deadline, future).await + } +} + +pub use tracked::{Duration, Instant, Interval, error, interval, sleep, timeout}; + +pub(crate) use tracked::{clear_domain, enter_domain, next_deadline_millis}; + +#[cfg(test)] +mod tests { + use super::tracked::MissedTickBehavior; + use super::*; + + #[tokio::test] + async fn tracks_reset_completion_and_drop_per_domain() { + let _domain = enter_domain(7); + { + let sleep = sleep(Duration::from_millis(50)); + tokio::pin!(sleep); + assert!(matches!(next_deadline_millis(7), Some(1..=50))); + assert_eq!(next_deadline_millis(8), None); + + sleep + .as_mut() + .reset(Instant::now() + Duration::from_millis(20)); + assert!(matches!(next_deadline_millis(7), Some(1..=20))); + } + assert_eq!(next_deadline_millis(7), None); + + let pending = sleep(Duration::from_secs(1)); + assert!(matches!(next_deadline_millis(7), Some(999..=1000))); + drop(pending); + assert_eq!(next_deadline_millis(7), None); + + sleep(Duration::ZERO).await; + assert_eq!(next_deadline_millis(7), None); + } + + #[tokio::test] + async fn tracks_interval_and_timeout_lifetimes() { + let _domain = enter_domain(9); + let mut ticker = interval(Duration::from_millis(40)); + assert_eq!(next_deadline_millis(9), Some(0)); + ticker.tick().await; + assert!(matches!(next_deadline_millis(9), Some(1..=40))); + drop(ticker); + assert_eq!(next_deadline_millis(9), None); + + let timeout = tokio::spawn(timeout( + Duration::from_millis(60), + std::future::pending::<()>(), + )); + tokio::task::yield_now().await; + assert!(matches!(next_deadline_millis(9), Some(1..=60))); + timeout.abort(); + let _ = timeout.await; + assert_eq!(next_deadline_millis(9), None); + } + + #[tokio::test] + async fn clears_all_deadlines_for_a_domain() { + let _domain = enter_domain(11); + let _timer = sleep(Duration::from_secs(1)); + assert!(next_deadline_millis(11).is_some()); + + clear_domain(11); + + assert_eq!(next_deadline_millis(11), None); + } + + #[test] + fn computes_missed_interval_deadlines_like_tokio() { + let tick = Instant::now(); + let now = tick + Duration::from_millis(250); + let period = Duration::from_millis(100); + + assert_eq!( + tracked::next_interval_deadline(MissedTickBehavior::Burst, tick, now, period), + tick + period + ); + assert_eq!( + tracked::next_interval_deadline(MissedTickBehavior::Delay, tick, now, period), + now + period + ); + assert_eq!( + tracked::next_interval_deadline(MissedTickBehavior::Skip, tick, now, period), + tick + Duration::from_millis(300) + ); + } +} diff --git a/easytier-core/src/wasi/wire/common.rs b/easytier-core/src/wasi/wire/common.rs new file mode 100644 index 00000000..65311039 --- /dev/null +++ b/easytier-core/src/wasi/wire/common.rs @@ -0,0 +1,53 @@ +use std::io; + +pub(crate) fn status(operation: &str, result: i32) -> io::Result<()> { + if result == 0 { + Ok(()) + } else { + Err(host_error(operation, result)) + } +} + +pub(crate) fn host_error(operation: &str, code: i32) -> io::Error { + io::Error::other(format!("host {operation} failed with code {code}")) +} + +pub(crate) fn tcp_connect_error(code: i32) -> io::Error { + let kind = match code { + -6 => io::ErrorKind::ConnectionRefused, + -7 => io::ErrorKind::ConnectionAborted, + -8 => io::ErrorKind::ConnectionReset, + -9 => io::ErrorKind::NotConnected, + _ => return host_error("take_tcp_connect", code), + }; + io::Error::new(kind, format!("host TCP connect failed with code {code}")) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn tcp_connect_status_preserves_error_kinds() { + assert_eq!( + tcp_connect_error(-6).kind(), + io::ErrorKind::ConnectionRefused + ); + assert_eq!( + tcp_connect_error(-7).kind(), + io::ErrorKind::ConnectionAborted + ); + assert_eq!(tcp_connect_error(-8).kind(), io::ErrorKind::ConnectionReset); + assert_eq!(tcp_connect_error(-9).kind(), io::ErrorKind::NotConnected); + assert_eq!(tcp_connect_error(-3).kind(), io::ErrorKind::Other); + } + + #[test] + fn status_preserves_success_and_host_error_context() { + assert!(status("start_read", 0).is_ok()); + assert_eq!( + status("start_read", -3).unwrap_err().kind(), + io::ErrorKind::Other + ); + } +} diff --git a/easytier-core/src/wasi/wire/data_plane.rs b/easytier-core/src/wasi/wire/data_plane.rs new file mode 100644 index 00000000..fa62c8c1 --- /dev/null +++ b/easytier-core/src/wasi/wire/data_plane.rs @@ -0,0 +1,161 @@ +//! Explicit big-endian wire records returned by data-plane guest exports. + +use std::{io, net::SocketAddr}; + +use crate::gateway::{DataPlaneCompletionDescriptor, DataPlaneErrorKind}; + +use super::socket::{SOCKET_ADDRESS_LEN, decode_socket_address, encode_socket_address}; + +pub(crate) const COMPLETION_LEN: usize = 12; +pub(crate) const TCP_CONNECT_RESULT_LEN: usize = 8 + SOCKET_ADDRESS_LEN * 2; +pub(crate) const TCP_BIND_RESULT_LEN: usize = 8 + SOCKET_ADDRESS_LEN; +pub(crate) const TCP_ACCEPT_RESULT_LEN: usize = TCP_CONNECT_RESULT_LEN; +pub(crate) const UDP_BIND_RESULT_LEN: usize = TCP_BIND_RESULT_LEN; +pub(crate) const TCP_READ_METADATA_LEN: usize = 1; +pub(crate) const UDP_RECEIVE_METADATA_LEN: usize = SOCKET_ADDRESS_LEN + 1; + +pub(crate) fn decode_ipv4_socket_address(wire: &[u8]) -> io::Result { + let wire = <&[u8; SOCKET_ADDRESS_LEN]>::try_from(wire).map_err(|_| { + io::Error::new(io::ErrorKind::InvalidInput, "invalid socket address length") + })?; + let address = decode_socket_address(wire)?; + if !address.is_ipv4() { + return Err(io::Error::new( + io::ErrorKind::Unsupported, + "data-plane ABI v2 supports IPv4 only", + )); + } + Ok(address) +} + +pub(crate) fn encode_completion(completion: DataPlaneCompletionDescriptor) -> [u8; COMPLETION_LEN] { + let mut wire = [0; COMPLETION_LEN]; + wire[..8].copy_from_slice(&completion.operation_id.get().to_be_bytes()); + wire[8..10].copy_from_slice(&(completion.kind as u16).to_be_bytes()); + wire[10..12].copy_from_slice(&completion.status.code().to_be_bytes()); + wire +} + +pub(crate) fn encode_resource_and_address( + resource: u64, + address: SocketAddr, +) -> [u8; TCP_BIND_RESULT_LEN] { + let mut wire = [0; TCP_BIND_RESULT_LEN]; + wire[..8].copy_from_slice(&resource.to_be_bytes()); + wire[8..].copy_from_slice(&encode_socket_address(address)); + wire +} + +pub(crate) fn encode_stream_addresses( + stream: u64, + local_addr: SocketAddr, + peer_addr: SocketAddr, +) -> [u8; TCP_CONNECT_RESULT_LEN] { + let mut wire = [0; TCP_CONNECT_RESULT_LEN]; + wire[..8].copy_from_slice(&stream.to_be_bytes()); + wire[8..8 + SOCKET_ADDRESS_LEN].copy_from_slice(&encode_socket_address(local_addr)); + wire[8 + SOCKET_ADDRESS_LEN..].copy_from_slice(&encode_socket_address(peer_addr)); + wire +} + +pub(crate) fn encode_tcp_read_metadata(eof: bool) -> [u8; TCP_READ_METADATA_LEN] { + [u8::from(eof)] +} + +pub(crate) fn encode_udp_receive_metadata( + peer_addr: SocketAddr, + truncated: bool, +) -> [u8; UDP_RECEIVE_METADATA_LEN] { + let mut wire = [0; UDP_RECEIVE_METADATA_LEN]; + wire[..SOCKET_ADDRESS_LEN].copy_from_slice(&encode_socket_address(peer_addr)); + wire[SOCKET_ADDRESS_LEN] = u8::from(truncated); + wire +} + +pub(crate) fn error_status(kind: DataPlaneErrorKind) -> i32 { + -(kind as i32) +} + +pub(crate) fn normalize_call_status(entered: bool, status: i32) -> i32 { + if entered { + status + } else { + error_status(DataPlaneErrorKind::HandleClosed) + } +} + +#[cfg(test)] +mod tests { + use crate::gateway::{DataPlaneCompletionStatus, DataPlaneOperationId, DataPlaneOperationKind}; + + use super::*; + + #[test] + fn completion_record_has_no_native_padding() { + let completion = DataPlaneCompletionDescriptor { + operation_id: DataPlaneOperationId::from_raw(0x0102_0304_0506_0708).unwrap(), + kind: DataPlaneOperationKind::UdpReceive, + status: DataPlaneCompletionStatus::Error(DataPlaneErrorKind::BufferTooSmall), + }; + assert_eq!( + encode_completion(completion), + [1, 2, 3, 4, 5, 6, 7, 8, 0, 7, 0, 13] + ); + } + + #[test] + fn operation_result_records_have_stable_layouts() { + let local_addr = "192.0.2.1:1234".parse().unwrap(); + let peer_addr = "198.51.100.2:4321".parse().unwrap(); + let resource = 0x0102_0304_0506_0708; + + let resource_and_address = encode_resource_and_address(resource, local_addr); + assert_eq!(resource_and_address.len(), UDP_BIND_RESULT_LEN); + assert_eq!(&resource_and_address[..8], &resource.to_be_bytes()); + assert_eq!( + &resource_and_address[8..], + &encode_socket_address(local_addr) + ); + + let stream_addresses = encode_stream_addresses(resource, local_addr, peer_addr); + assert_eq!(stream_addresses.len(), TCP_ACCEPT_RESULT_LEN); + assert_eq!(&stream_addresses[..8], &resource.to_be_bytes()); + assert_eq!( + &stream_addresses[8..8 + SOCKET_ADDRESS_LEN], + &encode_socket_address(local_addr) + ); + assert_eq!( + &stream_addresses[8 + SOCKET_ADDRESS_LEN..], + &encode_socket_address(peer_addr) + ); + + assert_eq!(encode_tcp_read_metadata(false), [0]); + assert_eq!(encode_tcp_read_metadata(true), [1]); + + let udp_metadata = encode_udp_receive_metadata(peer_addr, true); + assert_eq!( + &udp_metadata[..SOCKET_ADDRESS_LEN], + &encode_socket_address(peer_addr) + ); + assert_eq!(udp_metadata[SOCKET_ADDRESS_LEN], 1); + } + + #[test] + fn public_address_decoder_rejects_ipv6() { + let address = "[2001:db8::1]:80".parse().unwrap(); + let error = decode_ipv4_socket_address(&encode_socket_address(address)).unwrap_err(); + assert_eq!(error.kind(), io::ErrorKind::Unsupported); + } + + #[test] + fn lifecycle_failures_use_data_plane_status_codes() { + assert_eq!( + normalize_call_status(false, -1), + error_status(DataPlaneErrorKind::HandleClosed) + ); + assert_eq!( + normalize_call_status(true, error_status(DataPlaneErrorKind::Cancelled)), + error_status(DataPlaneErrorKind::Cancelled) + ); + } +} diff --git a/easytier-core/src/wasi/wire/dns.rs b/easytier-core/src/wasi/wire/dns.rs new file mode 100644 index 00000000..80077652 --- /dev/null +++ b/easytier-core/src/wasi/wire/dns.rs @@ -0,0 +1,248 @@ +use std::{io, net::IpAddr}; + +use crate::host::dns::{DnsQuery, DnsSrvRecord}; +use crate::socket::IpVersion; + +const DNS_WIRE_VERSION: u8 = 1; + +pub(crate) fn encode_query(query: &DnsQuery) -> io::Result> { + let host = query.host.as_bytes(); + let netns = query + .context + .netns + .as_ref() + .map(|netns| netns.token().as_bytes()); + let host_len = encoded_len("DNS host", host.len())?; + let netns_len = encoded_len("DNS netns token", netns.map_or(0, <[u8]>::len))?; + let mut encoded = Vec::with_capacity(16 + host.len() + netns.map_or(0, <[u8]>::len)); + encoded.push(DNS_WIRE_VERSION); + encoded.push(match query.context.ip_version { + IpVersion::V4 => 4, + IpVersion::V6 => 6, + IpVersion::Both => 0, + }); + encoded.push(u8::from(query.context.socket_mark.is_some())); + encoded.extend_from_slice(&query.context.socket_mark.unwrap_or_default().to_be_bytes()); + encoded.push(u8::from(netns.is_some())); + encoded.extend_from_slice(&netns_len.to_be_bytes()); + if let Some(netns) = netns { + encoded.extend_from_slice(netns); + } + encoded.extend_from_slice(&host_len.to_be_bytes()); + encoded.extend_from_slice(host); + Ok(encoded) +} + +pub(crate) fn decode_addresses(encoded: &[u8]) -> io::Result> { + let mut decoder = Decoder::new(encoded); + let count = decoder.take_count("DNS address count")?; + let mut addresses = Vec::with_capacity(count.min(64)); + for _ in 0..count { + let family = decoder.take_u8("DNS address family")?; + let address = match family { + 4 => IpAddr::V4(decoder.take_array::<4>("IPv4 address")?.into()), + 6 => IpAddr::V6(decoder.take_array::<16>("IPv6 address")?.into()), + _ => return Err(invalid_data("invalid DNS address family")), + }; + addresses.push(address); + } + decoder.finish()?; + Ok(addresses) +} + +pub(crate) fn decode_txt(encoded: &[u8]) -> io::Result { + let mut decoder = Decoder::new(encoded); + let length = decoder.take_count("DNS TXT length")?; + let text = decoder.take(length, "DNS TXT value")?; + decoder.finish()?; + String::from_utf8(text.to_vec()).map_err(|_| invalid_data("DNS TXT is not UTF-8")) +} + +pub(crate) fn decode_srv(encoded: &[u8]) -> io::Result> { + let mut decoder = Decoder::new(encoded); + let count = decoder.take_count("DNS SRV count")?; + let mut records = Vec::with_capacity(count.min(64)); + for _ in 0..count { + let priority = decoder.take_u16("DNS SRV priority")?; + let weight = decoder.take_u16("DNS SRV weight")?; + let port = decoder.take_u16("DNS SRV port")?; + let target_len = decoder.take_count("DNS SRV target length")?; + let target = String::from_utf8(decoder.take(target_len, "DNS SRV target")?.to_vec()) + .map_err(|_| invalid_data("DNS SRV target is not UTF-8"))?; + records.push(DnsSrvRecord { + priority, + weight, + port, + target, + }); + } + decoder.finish()?; + Ok(records) +} + +fn encoded_len(description: &str, length: usize) -> io::Result { + u32::try_from(length).map_err(|_| { + io::Error::new( + io::ErrorKind::InvalidInput, + format!("{description} is too long"), + ) + }) +} + +fn invalid_data(message: &'static str) -> io::Error { + io::Error::new(io::ErrorKind::InvalidData, message) +} + +struct Decoder<'a> { + encoded: &'a [u8], + offset: usize, +} + +impl<'a> Decoder<'a> { + fn new(encoded: &'a [u8]) -> Self { + Self { encoded, offset: 0 } + } + + fn take(&mut self, length: usize, description: &'static str) -> io::Result<&'a [u8]> { + let end = self + .offset + .checked_add(length) + .ok_or_else(|| invalid_data("DNS result length overflow"))?; + let value = self + .encoded + .get(self.offset..end) + .ok_or_else(|| invalid_data(description))?; + self.offset = end; + Ok(value) + } + + fn take_u8(&mut self, description: &'static str) -> io::Result { + Ok(self.take(1, description)?[0]) + } + + fn take_u16(&mut self, description: &'static str) -> io::Result { + Ok(u16::from_be_bytes(self.take_array::<2>(description)?)) + } + + fn take_count(&mut self, description: &'static str) -> io::Result { + usize::try_from(u32::from_be_bytes(self.take_array::<4>(description)?)) + .map_err(|_| invalid_data("DNS result count exceeds guest usize")) + } + + fn take_array(&mut self, description: &'static str) -> io::Result<[u8; N]> { + self.take(N, description)? + .try_into() + .map_err(|_| invalid_data(description)) + } + + fn finish(self) -> io::Result<()> { + if self.offset == self.encoded.len() { + Ok(()) + } else { + Err(invalid_data("DNS result has trailing bytes")) + } + } +} + +#[cfg(test)] +mod tests { + use crate::socket::{NetNamespace, SocketContext}; + + use super::*; + + #[test] + fn query_encoding_has_stable_versioned_layout() { + let query = DnsQuery::new( + "peer.example", + SocketContext { + ip_version: IpVersion::V6, + socket_mark: Some(0x01020304), + netns: Some(NetNamespace::new("netns0")), + }, + ); + let encoded = encode_query(&query).unwrap(); + + let mut expected = vec![DNS_WIRE_VERSION, 6, 1, 1, 2, 3, 4, 1]; + expected.extend_from_slice(&6_u32.to_be_bytes()); + expected.extend_from_slice(b"netns0"); + expected.extend_from_slice(&12_u32.to_be_bytes()); + expected.extend_from_slice(b"peer.example"); + assert_eq!(encoded, expected); + + let without_optional = encode_query(&DnsQuery::new( + "v4.example", + SocketContext { + ip_version: IpVersion::V4, + socket_mark: None, + netns: None, + }, + )) + .unwrap(); + assert_eq!( + &without_optional[..12], + &[1, 4, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0] + ); + assert_eq!(&without_optional[12..16], &10_u32.to_be_bytes()); + assert_eq!(&without_optional[16..], b"v4.example"); + } + + #[test] + fn decodes_owned_address_txt_and_srv_results() { + let mut addresses = 2_u32.to_be_bytes().to_vec(); + addresses.push(4); + addresses.extend_from_slice(&[192, 0, 2, 1]); + addresses.push(6); + addresses.extend_from_slice( + &"2001:db8::1" + .parse::() + .unwrap() + .octets(), + ); + assert_eq!( + decode_addresses(&addresses).unwrap(), + vec![ + "192.0.2.1".parse::().unwrap(), + "2001:db8::1".parse::().unwrap(), + ] + ); + + let mut txt = 12_u32.to_be_bytes().to_vec(); + txt.extend_from_slice(b"tcp://peer:1"); + assert_eq!(decode_txt(&txt).unwrap(), "tcp://peer:1"); + + let mut srv = 1_u32.to_be_bytes().to_vec(); + srv.extend_from_slice(&10_u16.to_be_bytes()); + srv.extend_from_slice(&20_u16.to_be_bytes()); + srv.extend_from_slice(&11010_u16.to_be_bytes()); + srv.extend_from_slice(&13_u32.to_be_bytes()); + srv.extend_from_slice(b"peer.example."); + assert_eq!( + decode_srv(&srv).unwrap(), + vec![DnsSrvRecord { + priority: 10, + weight: 20, + port: 11010, + target: "peer.example.".to_owned(), + }] + ); + } + + #[test] + fn rejects_malformed_results() { + assert!(decode_addresses(&[0, 0, 0]).is_err()); + assert!(decode_addresses(&[0, 0, 0, 1, 9]).is_err()); + assert!(decode_addresses(&[0, 0, 0, 0, 1]).is_err()); + + let mut invalid_txt = 1_u32.to_be_bytes().to_vec(); + invalid_txt.push(0xff); + assert!(decode_txt(&invalid_txt).is_err()); + + let mut truncated_srv = 1_u32.to_be_bytes().to_vec(); + truncated_srv.extend_from_slice(&10_u16.to_be_bytes()); + assert!(decode_srv(&truncated_srv).is_err()); + + let mut trailing_srv = 0_u32.to_be_bytes().to_vec(); + trailing_srv.push(0); + assert!(decode_srv(&trailing_srv).is_err()); + } +} diff --git a/easytier-core/src/wasi/wire/mod.rs b/easytier-core/src/wasi/wire/mod.rs new file mode 100644 index 00000000..36cab443 --- /dev/null +++ b/easytier-core/src/wasi/wire/mod.rs @@ -0,0 +1,8 @@ +//! WASI host ABI codecs shared by concrete adapters. + +pub(crate) mod common; +#[cfg(feature = "proxy-smoltcp-stack")] +pub(crate) mod data_plane; +pub(crate) mod dns; +pub(crate) mod options; +pub(crate) mod socket; diff --git a/easytier-core/src/wasi/wire/options.rs b/easytier-core/src/wasi/wire/options.rs new file mode 100644 index 00000000..e933427f --- /dev/null +++ b/easytier-core/src/wasi/wire/options.rs @@ -0,0 +1,394 @@ +use std::io; + +use crate::host::socket::{ + HostSocketHandle, + factory::{HostTcpConnectResult, HostUdpBindResult}, + listener::HostTcpBindResult, +}; +use crate::socket::{ + IpVersion, SocketContext, + tcp::{TcpConnectOptions, TcpListenOptions, TcpListenPurpose, TcpSocketPurpose}, + udp::{UdpBindOptions, UdpSocketPurpose}, +}; + +use super::socket::{SOCKET_ADDRESS_LEN, decode_socket_address, encode_socket_address}; + +const OPTIONS_VERSION: u8 = 2; +pub(crate) const TCP_SOCKET_RESULT_LEN: usize = 8 + SOCKET_ADDRESS_LEN * 2; +pub(crate) const BOUND_SOCKET_RESULT_LEN: usize = 8 + SOCKET_ADDRESS_LEN; + +pub(crate) fn encode_tcp_connect_options(options: &TcpConnectOptions) -> io::Result> { + let mut encoded = Vec::with_capacity( + 75 + context_variable_len(&options.bind.context) + + bind_device_len(&options.bind.bind_device), + ); + encoded.push(OPTIONS_VERSION); + encoded.extend_from_slice(&encode_socket_address(options.remote_addr)); + encode_optional_address(&mut encoded, options.bind.local_addr); + encode_context(&mut encoded, &options.bind.context)?; + encoded.push(match options.bind.reuse_addr { + None => 0, + Some(false) => 1, + Some(true) => 2, + }); + encoded.push(u8::from(options.bind.reuse_port)); + encoded.push(u8::from(options.bind.only_v6)); + encoded.push(match options.purpose { + TcpSocketPurpose::DirectConnect => 0, + TcpSocketPurpose::FakeTcp => 1, + TcpSocketPurpose::HolePunch => 2, + TcpSocketPurpose::ManualConnect => 3, + TcpSocketPurpose::ProxyNat => 4, + TcpSocketPurpose::StunProbe => 5, + TcpSocketPurpose::Socks5 => 6, + TcpSocketPurpose::PortForward => 7, + TcpSocketPurpose::DataPlane => 8, + }); + encode_bind_device(&mut encoded, &options.bind.bind_device)?; + Ok(encoded) +} + +pub(crate) fn encode_udp_bind_options(options: &UdpBindOptions) -> io::Result> { + let mut encoded = Vec::with_capacity( + 48 + context_variable_len(&options.context) + bind_device_len(&options.bind_device), + ); + encoded.push(OPTIONS_VERSION); + encode_optional_address(&mut encoded, options.local_addr); + encode_context(&mut encoded, &options.context)?; + encoded.push(u8::from(options.reuse_addr)); + encoded.push(u8::from(options.reuse_port)); + encoded.push(u8::from(options.only_v6)); + encoded.push(match options.purpose { + UdpSocketPurpose::HolePunchControl => 0, + UdpSocketPurpose::HolePunchCandidate => 1, + UdpSocketPurpose::DirectConnect => 2, + UdpSocketPurpose::PortBoundListener => 3, + UdpSocketPurpose::ProxyNat => 4, + UdpSocketPurpose::StunProbe => 5, + UdpSocketPurpose::Socks5 => 6, + UdpSocketPurpose::PortForward => 7, + UdpSocketPurpose::PortLease => 8, + }); + encode_bind_device(&mut encoded, &options.bind_device)?; + Ok(encoded) +} + +pub(crate) fn encode_tcp_listen_options(options: &TcpListenOptions) -> io::Result> { + let mut encoded = Vec::with_capacity( + 48 + context_variable_len(&options.bind.context) + + bind_device_len(&options.bind.bind_device), + ); + encoded.push(OPTIONS_VERSION); + encode_optional_address(&mut encoded, options.bind.local_addr); + encode_context(&mut encoded, &options.bind.context)?; + encoded.push(match options.bind.reuse_addr { + None => 0, + Some(false) => 1, + Some(true) => 2, + }); + encoded.push(u8::from(options.bind.reuse_port)); + encoded.push(u8::from(options.bind.only_v6)); + encoded.push(match options.purpose { + TcpListenPurpose::DirectConnect => 0, + TcpListenPurpose::HolePunch => 1, + TcpListenPurpose::ManualConnect => 2, + TcpListenPurpose::ProxyNat => 3, + TcpListenPurpose::Socks5 => 4, + TcpListenPurpose::PortForward => 5, + TcpListenPurpose::PortLease => 6, + }); + encode_bind_device(&mut encoded, &options.bind.bind_device)?; + Ok(encoded) +} + +pub(crate) fn decode_tcp_socket_result( + encoded: &[u8; TCP_SOCKET_RESULT_LEN], +) -> io::Result { + let local = <[u8; SOCKET_ADDRESS_LEN]>::try_from(&encoded[8..8 + SOCKET_ADDRESS_LEN]).unwrap(); + let peer = <[u8; SOCKET_ADDRESS_LEN]>::try_from(&encoded[8 + SOCKET_ADDRESS_LEN..]).unwrap(); + Ok(HostTcpConnectResult { + handle: HostSocketHandle(u64::from_be_bytes(encoded[..8].try_into().unwrap())), + local_addr: decode_socket_address(&local)?, + peer_addr: decode_socket_address(&peer)?, + transport_label: None, + }) +} + +pub(crate) fn decode_udp_bind_result( + encoded: &[u8; BOUND_SOCKET_RESULT_LEN], +) -> io::Result { + Ok(HostUdpBindResult { + handle: decode_bound_handle(encoded), + local_addr: decode_bound_address(encoded)?, + }) +} + +pub(crate) fn decode_tcp_bind_result( + encoded: &[u8; BOUND_SOCKET_RESULT_LEN], +) -> io::Result { + Ok(HostTcpBindResult { + handle: decode_bound_handle(encoded), + local_addr: decode_bound_address(encoded)?, + }) +} + +fn decode_bound_handle(encoded: &[u8; BOUND_SOCKET_RESULT_LEN]) -> HostSocketHandle { + HostSocketHandle(u64::from_be_bytes(encoded[..8].try_into().unwrap())) +} + +fn decode_bound_address( + encoded: &[u8; BOUND_SOCKET_RESULT_LEN], +) -> io::Result { + let address = <[u8; SOCKET_ADDRESS_LEN]>::try_from(&encoded[8..]).unwrap(); + decode_socket_address(&address) +} + +fn encode_optional_address(encoded: &mut Vec, address: Option) { + encoded.extend_from_slice( + &address + .map(encode_socket_address) + .unwrap_or([0; SOCKET_ADDRESS_LEN]), + ); +} + +fn encode_mark(encoded: &mut Vec, mark: Option) { + encoded.push(u8::from(mark.is_some())); + encoded.extend_from_slice(&mark.unwrap_or_default().to_be_bytes()); +} + +fn encode_context(encoded: &mut Vec, context: &SocketContext) -> io::Result<()> { + encoded.push(match context.ip_version { + IpVersion::V4 => 0, + IpVersion::V6 => 1, + IpVersion::Both => 2, + }); + encode_mark(encoded, context.socket_mark); + let netns = context.netns.as_ref().map(|netns| netns.token().as_bytes()); + encoded.push(u8::from(netns.is_some())); + let length = u32::try_from(netns.map_or(0, <[u8]>::len)) + .map_err(|_| io::Error::new(io::ErrorKind::InvalidInput, "netns token is too long"))?; + encoded.extend_from_slice(&length.to_be_bytes()); + if let Some(netns) = netns { + encoded.extend_from_slice(netns); + } + Ok(()) +} + +pub(crate) fn encode_socket_context(context: &SocketContext) -> io::Result> { + let mut encoded = Vec::with_capacity(11 + context_variable_len(context)); + encode_context(&mut encoded, context)?; + Ok(encoded) +} + +fn encode_bind_device(encoded: &mut Vec, device: &Option) -> io::Result<()> { + let bytes = device.as_deref().unwrap_or_default().as_bytes(); + encoded.push(u8::from(device.is_some())); + let length = u32::try_from(bytes.len()) + .map_err(|_| io::Error::new(io::ErrorKind::InvalidInput, "bind device is too long"))?; + encoded.extend_from_slice(&length.to_be_bytes()); + encoded.extend_from_slice(bytes); + Ok(()) +} + +fn bind_device_len(device: &Option) -> usize { + device.as_ref().map_or(0, String::len) +} + +fn context_variable_len(context: &SocketContext) -> usize { + context + .netns + .as_ref() + .map_or(0, |netns| netns.token().len()) +} + +#[cfg(test)] +mod tests { + use crate::socket::{NetNamespace, tcp::TcpBindOptions}; + + use super::*; + + #[test] + fn encodes_tcp_connect_options_with_stable_offsets() { + let options = TcpConnectOptions { + remote_addr: "192.0.2.2:11013".parse().unwrap(), + bind: TcpBindOptions::default() + .with_local_addr(Some("[2001:db8::1]:22026".parse().unwrap())) + .with_socket_mark(Some(0x01020304)) + .with_bind_device(Some("device0".to_owned())) + .with_reuse_addr(true) + .with_reuse_port(true) + .with_only_v6(true), + purpose: TcpSocketPurpose::ManualConnect, + }; + let encoded = encode_tcp_connect_options(&options).unwrap(); + assert_eq!(encoded.len(), 82); + assert_eq!(encoded[0], OPTIONS_VERSION); + assert_eq!(&encoded[55..66], &[2, 1, 1, 2, 3, 4, 0, 0, 0, 0, 0]); + assert_eq!(&encoded[66..70], &[2, 1, 1, 3]); + assert_eq!(encoded[70], 1); + assert_eq!(&encoded[71..75], &7_u32.to_be_bytes()); + assert_eq!(&encoded[75..], b"device0"); + } + + #[test] + fn bind_device_presence_distinguishes_none_empty_and_named() { + let remote = "192.0.2.2:11013".parse().unwrap(); + let none = encode_tcp_connect_options(&TcpConnectOptions::direct_connect(remote)).unwrap(); + let empty = encode_tcp_connect_options( + &TcpConnectOptions::direct_connect(remote) + .with_bind(TcpBindOptions::default().with_bind_device(Some(String::new()))), + ) + .unwrap(); + assert_eq!(&none[70..75], &[0, 0, 0, 0, 0]); + assert_eq!(&empty[70..75], &[1, 0, 0, 0, 0]); + + let udp_none = encode_udp_bind_options(&UdpBindOptions::direct_connect()).unwrap(); + let udp_empty = encode_udp_bind_options( + &UdpBindOptions::direct_connect().with_bind_device(Some(String::new())), + ) + .unwrap(); + assert_eq!(&udp_none[43..48], &[0, 0, 0, 0, 0]); + assert_eq!(&udp_empty[43..48], &[1, 0, 0, 0, 0]); + + let listen_none = encode_tcp_listen_options(&TcpListenOptions::direct_connect( + "192.0.2.1:11013".parse().unwrap(), + )) + .unwrap(); + let listen_empty = encode_tcp_listen_options( + &TcpListenOptions::direct_connect("192.0.2.1:11013".parse().unwrap()) + .with_bind(TcpBindOptions::default().with_bind_device(Some(String::new()))), + ) + .unwrap(); + assert_eq!(&listen_none[43..48], &[0, 0, 0, 0, 0]); + assert_eq!(&listen_empty[43..48], &[1, 0, 0, 0, 0]); + } + + #[test] + fn encodes_socket_context_netns_and_zero_mark() { + let bind = TcpBindOptions::default().with_context( + SocketContext::default() + .with_ip_version(IpVersion::V6) + .with_socket_mark(Some(0)) + .with_netns(Some(NetNamespace::new("instance-a"))), + ); + let encoded = encode_tcp_connect_options( + &TcpConnectOptions::direct_connect("[2001:db8::2]:11013".parse().unwrap()) + .with_bind(bind), + ) + .unwrap(); + + assert_eq!(encoded[55], 1); + assert_eq!(&encoded[56..61], &[1, 0, 0, 0, 0]); + assert_eq!(encoded[61], 1); + assert_eq!(&encoded[62..66], &10_u32.to_be_bytes()); + assert_eq!(&encoded[66..76], b"instance-a"); + } + + #[test] + fn encodes_standalone_socket_context() { + let context = SocketContext::default().with_ip_version(IpVersion::V6); + + assert_eq!( + encode_socket_context(&context).unwrap(), + vec![1, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0] + ); + } + + #[test] + fn encodes_udp_proxy_nat_purpose() { + let encoded = encode_udp_bind_options(&UdpBindOptions::proxy_nat()).unwrap(); + assert_eq!(encoded[42], 4); + } + + #[test] + fn encodes_udp_socks5_purpose() { + let encoded = encode_udp_bind_options(&UdpBindOptions::socks5()).unwrap(); + assert_eq!(encoded[42], 6); + } + + #[test] + fn encodes_tcp_proxy_nat_purposes() { + let remote = "192.0.2.2:11013".parse().unwrap(); + let connect = encode_tcp_connect_options(&TcpConnectOptions::proxy_nat(remote)).unwrap(); + assert_eq!(connect[69], 4); + + let local = "0.0.0.0:0".parse().unwrap(); + let listen = encode_tcp_listen_options(&TcpListenOptions::proxy_nat(local)).unwrap(); + assert_eq!(listen[42], 3); + } + + #[test] + fn encodes_stun_probe_purposes() { + let remote = "192.0.2.2:3478".parse().unwrap(); + let local = "0.0.0.0:0".parse().unwrap(); + let tcp = + encode_tcp_connect_options(&TcpConnectOptions::stun_probe(remote, local)).unwrap(); + assert_eq!(tcp[69], 5); + + let udp = encode_udp_bind_options(&UdpBindOptions::stun_probe()).unwrap(); + assert_eq!(udp[42], 5); + } + + #[test] + fn encodes_gateway_purposes_with_stable_values() { + let remote = "192.0.2.2:443".parse().unwrap(); + assert_eq!( + encode_tcp_connect_options(&TcpConnectOptions::socks5(remote)).unwrap()[69], + 6 + ); + assert_eq!( + encode_tcp_connect_options(&TcpConnectOptions::port_forward(remote)).unwrap()[69], + 7 + ); + assert_eq!( + encode_tcp_connect_options(&TcpConnectOptions::data_plane(remote)).unwrap()[69], + 8 + ); + + let local = "0.0.0.0:0".parse().unwrap(); + assert_eq!( + encode_tcp_listen_options(&TcpListenOptions::socks5(local)).unwrap()[42], + 4 + ); + assert_eq!( + encode_tcp_listen_options(&TcpListenOptions::port_forward(local)).unwrap()[42], + 5 + ); + assert_eq!( + encode_tcp_listen_options(&TcpListenOptions::port_lease(local)).unwrap()[42], + 6 + ); + assert_eq!( + encode_udp_bind_options(&UdpBindOptions::port_forward(local)).unwrap()[42], + 7 + ); + assert_eq!( + encode_udp_bind_options(&UdpBindOptions::port_lease(local)).unwrap()[42], + 8 + ); + } + + #[test] + fn decodes_fixed_socket_results() { + let mut tcp = [0_u8; TCP_SOCKET_RESULT_LEN]; + tcp[..8].copy_from_slice(&41_u64.to_be_bytes()); + tcp[8..35].copy_from_slice(&encode_socket_address("192.0.2.1:40100".parse().unwrap())); + tcp[35..].copy_from_slice(&encode_socket_address("192.0.2.2:11013".parse().unwrap())); + let result = decode_tcp_socket_result(&tcp).unwrap(); + assert_eq!(result.handle, HostSocketHandle(41)); + assert_eq!(result.local_addr, "192.0.2.1:40100".parse().unwrap()); + assert_eq!(result.peer_addr, "192.0.2.2:11013".parse().unwrap()); + + let mut bound = [0_u8; BOUND_SOCKET_RESULT_LEN]; + bound[..8].copy_from_slice(&42_u64.to_be_bytes()); + bound[8..].copy_from_slice(&encode_socket_address("[::]:22026".parse().unwrap())); + assert_eq!( + decode_udp_bind_result(&bound).unwrap().handle, + HostSocketHandle(42) + ); + assert_eq!( + decode_tcp_bind_result(&bound).unwrap().local_addr, + "[::]:22026".parse().unwrap() + ); + } +} diff --git a/easytier-core/src/wasi/wire/socket.rs b/easytier-core/src/wasi/wire/socket.rs new file mode 100644 index 00000000..e05201e6 --- /dev/null +++ b/easytier-core/src/wasi/wire/socket.rs @@ -0,0 +1,207 @@ +use std::{ + io, + net::{IpAddr, Ipv4Addr, Ipv6Addr, SocketAddr, SocketAddrV4, SocketAddrV6}, +}; + +pub(crate) const SOCKET_ADDRESS_LEN: usize = 27; +pub(crate) const UDP_METADATA_LEN: usize = 48; + +const V4_FAMILY: u8 = 4; +const V6_FAMILY: u8 = 6; +const ADDRESS_FAMILY: usize = 0; +const ADDRESS_BYTES: std::ops::Range = 1..17; +const PORT_BYTES: std::ops::Range = 17..19; +const FLOWINFO_BYTES: std::ops::Range = 19..23; +const SCOPE_ID_BYTES: std::ops::Range = 23..27; +const OPTIONAL_IP_FAMILY: usize = 27; +const OPTIONAL_IP_BYTES: std::ops::Range = 28..44; +const OPTIONAL_IFINDEX_BYTES: std::ops::Range = 44..48; + +pub(crate) fn encode_udp_metadata( + peer_addr: SocketAddr, + optional_ip: Option, + optional_ifindex: Option, +) -> [u8; UDP_METADATA_LEN] { + let mut wire = [0_u8; UDP_METADATA_LEN]; + wire[..SOCKET_ADDRESS_LEN].copy_from_slice(&encode_socket_address(peer_addr)); + + match optional_ip { + None => {} + Some(IpAddr::V4(ip)) => { + wire[OPTIONAL_IP_FAMILY] = V4_FAMILY; + wire[OPTIONAL_IP_BYTES.start..OPTIONAL_IP_BYTES.start + 4] + .copy_from_slice(&ip.octets()); + } + Some(IpAddr::V6(ip)) => { + wire[OPTIONAL_IP_FAMILY] = V6_FAMILY; + wire[OPTIONAL_IP_BYTES].copy_from_slice(&ip.octets()); + } + } + if let Some(ifindex) = optional_ifindex { + wire[OPTIONAL_IFINDEX_BYTES].copy_from_slice(&ifindex.to_be_bytes()); + } + wire +} + +pub(crate) fn encode_socket_address(addr: SocketAddr) -> [u8; SOCKET_ADDRESS_LEN] { + let mut wire = [0_u8; SOCKET_ADDRESS_LEN]; + match addr { + SocketAddr::V4(addr) => { + wire[ADDRESS_FAMILY] = V4_FAMILY; + wire[ADDRESS_BYTES.start..ADDRESS_BYTES.start + 4].copy_from_slice(&addr.ip().octets()); + wire[PORT_BYTES].copy_from_slice(&addr.port().to_be_bytes()); + } + SocketAddr::V6(addr) => { + wire[ADDRESS_FAMILY] = V6_FAMILY; + wire[ADDRESS_BYTES].copy_from_slice(&addr.ip().octets()); + wire[PORT_BYTES].copy_from_slice(&addr.port().to_be_bytes()); + wire[FLOWINFO_BYTES].copy_from_slice(&addr.flowinfo().to_be_bytes()); + wire[SCOPE_ID_BYTES].copy_from_slice(&addr.scope_id().to_be_bytes()); + } + } + wire +} + +pub(crate) fn decode_udp_metadata( + wire: &[u8; UDP_METADATA_LEN], +) -> io::Result<(SocketAddr, Option, Option)> { + let address = <[u8; SOCKET_ADDRESS_LEN]>::try_from(&wire[..SOCKET_ADDRESS_LEN]).unwrap(); + let peer_addr = decode_socket_address(&address)?; + + let optional_ip = match wire[OPTIONAL_IP_FAMILY] { + 0 => { + require_zero(&wire[OPTIONAL_IP_BYTES], "absent optional IP")?; + require_zero( + &wire[OPTIONAL_IFINDEX_BYTES], + "absent optional IP interface index", + )?; + None + } + V4_FAMILY => { + require_zero( + &wire[OPTIONAL_IP_BYTES.start + 4..OPTIONAL_IP_BYTES.end], + "optional IPv4 padding", + )?; + require_zero( + &wire[OPTIONAL_IFINDEX_BYTES], + "optional IPv4 interface index", + )?; + Some(IpAddr::V4(Ipv4Addr::from( + <[u8; 4]>::try_from(&wire[OPTIONAL_IP_BYTES.start..OPTIONAL_IP_BYTES.start + 4]) + .unwrap(), + ))) + } + V6_FAMILY => Some(IpAddr::V6(Ipv6Addr::from( + <[u8; 16]>::try_from(&wire[OPTIONAL_IP_BYTES]).unwrap(), + ))), + family => return Err(invalid_family("optional IP", family)), + }; + let optional_ifindex = + match u32::from_be_bytes(wire[OPTIONAL_IFINDEX_BYTES].try_into().unwrap()) { + 0 => None, + ifindex => Some(ifindex), + }; + Ok((peer_addr, optional_ip, optional_ifindex)) +} + +pub(crate) fn decode_socket_address(wire: &[u8; SOCKET_ADDRESS_LEN]) -> io::Result { + let port = u16::from_be_bytes(wire[PORT_BYTES].try_into().unwrap()); + match wire[ADDRESS_FAMILY] { + V4_FAMILY => { + require_zero( + &wire[ADDRESS_BYTES.start + 4..ADDRESS_BYTES.end], + "IPv4 padding", + )?; + require_zero(&wire[FLOWINFO_BYTES], "IPv4 flowinfo")?; + require_zero(&wire[SCOPE_ID_BYTES], "IPv4 scope ID")?; + Ok(SocketAddr::V4(SocketAddrV4::new( + Ipv4Addr::from( + <[u8; 4]>::try_from(&wire[ADDRESS_BYTES.start..ADDRESS_BYTES.start + 4]) + .unwrap(), + ), + port, + ))) + } + V6_FAMILY => Ok(SocketAddr::V6(SocketAddrV6::new( + Ipv6Addr::from(<[u8; 16]>::try_from(&wire[ADDRESS_BYTES]).unwrap()), + port, + u32::from_be_bytes(wire[FLOWINFO_BYTES].try_into().unwrap()), + u32::from_be_bytes(wire[SCOPE_ID_BYTES].try_into().unwrap()), + ))), + family => Err(invalid_family("peer address", family)), + } +} + +fn require_zero(bytes: &[u8], field: &str) -> io::Result<()> { + if bytes.iter().all(|byte| *byte == 0) { + Ok(()) + } else { + Err(io::Error::new( + io::ErrorKind::InvalidData, + format!("host UDP metadata has nonzero {field}"), + )) + } +} + +fn invalid_family(field: &str, family: u8) -> io::Error { + io::Error::new( + io::ErrorKind::InvalidData, + format!("host UDP metadata has invalid {field} family {family}"), + ) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn round_trips_ipv4_address_and_optional_source() { + let peer = SocketAddr::V4(SocketAddrV4::new(Ipv4Addr::new(192, 0, 2, 1), 11013)); + let source = Some(IpAddr::V4(Ipv4Addr::new(198, 51, 100, 2))); + let expected = [ + 0x04, 0xc0, 0x00, 0x02, 0x01, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, + 0x00, 0x00, 0x00, 0x2b, 0x05, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x04, + 0xc6, 0x33, 0x64, 0x02, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, + 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, + ]; + assert_eq!(encode_udp_metadata(peer, source, None), expected); + assert_eq!( + decode_udp_metadata(&expected).unwrap(), + (peer, source, None) + ); + } + + #[test] + fn round_trips_ipv6_flow_scope_and_optional_destination() { + let peer = SocketAddr::V6(SocketAddrV6::new( + "2001:db8::1".parse().unwrap(), + 22026, + 7, + 11, + )); + let destination = Some(IpAddr::V6("2001:db8::2".parse().unwrap())); + assert_eq!( + decode_udp_metadata(&encode_udp_metadata(peer, destination, Some(17))).unwrap(), + (peer, destination, Some(17)) + ); + } + + #[test] + fn rejects_noncanonical_or_unknown_families() { + let mut wire = encode_udp_metadata("192.0.2.1:11013".parse().unwrap(), None, None); + wire[ADDRESS_BYTES.start + 4] = 1; + assert!(decode_udp_metadata(&wire).is_err()); + + wire = encode_udp_metadata("192.0.2.1:11013".parse().unwrap(), None, None); + wire[OPTIONAL_IP_FAMILY] = 9; + assert!(decode_udp_metadata(&wire).is_err()); + + wire = encode_udp_metadata( + "192.0.2.1:11013".parse().unwrap(), + Some("192.0.2.2".parse().unwrap()), + None, + ); + wire[OPTIONAL_IFINDEX_BYTES.end - 1] = 1; + assert!(decode_udp_metadata(&wire).is_err()); + } +} diff --git a/easytier-core/testdata/wasi_core_instance_create.json b/easytier-core/testdata/wasi_core_instance_create.json new file mode 100644 index 00000000..e1806b4e --- /dev/null +++ b/easytier-core/testdata/wasi_core_instance_create.json @@ -0,0 +1,14 @@ +{ + "version": 14, + "config": "instance_id = \"018f4fb1-7a2c-7d1f-9d89-935b0ad7e135\"\ninstance_name = \"wasi-core\"\n\n[network_identity]\nnetwork_name = \"default\"\nnetwork_secret = \"test\"\n\n[flags]\ndisable_p2p = true\nenable_encryption = false\nbind_device = false\n", + "environment": { + "public_ipv4": null, + "interface_ipv4s": [], + "public_ipv6": null, + "interface_ipv6s": [], + "mapped_listeners": [], + "local_ips": [], + "protected_tcp_ports": [], + "preferred_ipv6_sources": [] + } +} diff --git a/easytier-core/tests/feature_profiles.rs b/easytier-core/tests/feature_profiles.rs new file mode 100644 index 00000000..718670f3 --- /dev/null +++ b/easytier-core/tests/feature_profiles.rs @@ -0,0 +1,21 @@ +#![cfg(all( + feature = "test-utils", + feature = "proxy-cidr-monitor", + not(feature = "wrapped-transport") +))] + +use easytier_core::{config::toml::TomlConfig, instance::CoreInstanceConfig}; + +#[test] +fn manual_routes_require_cidr_monitor_not_wrapped_transport() { + let toml = TomlConfig::new_from_str( + r#" + instance_name = "manual-routes-cidr-monitor" + routes = ["192.0.2.0/24"] + "#, + ) + .unwrap(); + let config = CoreInstanceConfig::from_toml(&toml).unwrap(); + + config.validate_build_capabilities_for_test().unwrap(); +} diff --git a/easytier-gui/src-tauri/Cargo.toml b/easytier-gui/src-tauri/Cargo.toml index 0a79975e..bdcefdec 100644 --- a/easytier-gui/src-tauri/Cargo.toml +++ b/easytier-gui/src-tauri/Cargo.toml @@ -25,6 +25,7 @@ serde = { version = "1", features = ["derive"] } serde_json = "1" easytier = { path = "../../easytier" } +easytier-core = { path = "../../easytier-core" } tokio = { version = "1", features = ["full"] } anyhow = "1.0" chrono = { version = "0.4.37", features = ["serde"] } diff --git a/easytier-gui/src-tauri/src/lib.rs b/easytier-gui/src-tauri/src/lib.rs index 876f201f..ce4bd86b 100644 --- a/easytier-gui/src-tauri/src/lib.rs +++ b/easytier-gui/src-tauri/src/lib.rs @@ -4,30 +4,35 @@ mod elevate; use anyhow::Context; +#[cfg(target_os = "android")] +use easytier::instance::factory::subscribe_native_instance_event; use easytier::proto::api::manage::{ CollectNetworkInfoResponse, ValidateConfigResponse, WebClientService, WebClientServiceClientFactory, }; -use easytier::rpc_service::remote_client::{ - GetNetworkMetasResponse, ListNetworkInstanceIdsJsonResp, ListNetworkProps, RemoteClientManager, - Storage, -}; use easytier::web_client::{self, WebClient}; use easytier::{ + common::config::{NetworkConfig, NetworkConfigExt}, common::{ config::{ ConfigLoader, ConfigSource, FileLoggerConfig, LoggingConfigBuilder, TomlConfigLoader, }, log, }, - instance_manager::NetworkInstanceManager, - launcher::NetworkConfig, + instance::factory::{NativeInstanceManager, native_instance_manager}, + proto::rpc::standalone::{runtime_rpc_dialer, runtime_rpc_listener}, rpc_service::ApiRpcServer, - tunnel::TunnelListener, - tunnel::ring::RingTunnelListener, - tunnel::tcp::TcpTunnelListener, utils::panic::setup_panic_handler, }; +use easytier_core::management::config_source_to_rpc; +use easytier_core::management::remote_client::{ + GetNetworkMetasResponse, ListNetworkInstanceIdsJsonResp, ListNetworkProps, RemoteClientManager, + Storage, +}; +use easytier_core::{ + connectivity::protocol::raw::TunnelDialer as _, process_runtime::CoreProcessRuntime, + socket::SocketListener, tunnel::Tunnel, +}; use std::ops::Deref; use std::sync::Arc; use tokio::sync::{Mutex, RwLock, RwLockReadGuard}; @@ -38,7 +43,7 @@ use tauri::{AppHandle, Emitter, Manager as _}; #[cfg(not(target_os = "android"))] use tauri::tray::{MouseButton, MouseButtonState, TrayIconBuilder, TrayIconEvent}; -static INSTANCE_MANAGER: once_cell::sync::Lazy>>> = +static INSTANCE_MANAGER: once_cell::sync::Lazy>>> = once_cell::sync::Lazy::new(|| RwLock::new(None)); static RPC_RING_UUID: once_cell::sync::Lazy = @@ -47,7 +52,7 @@ static RPC_RING_UUID: once_cell::sync::Lazy = static CLIENT_MANAGER: once_cell::sync::Lazy>> = once_cell::sync::Lazy::new(|| RwLock::new(None)); -type BoxedTunnelListener = Box; +type BoxedTunnelListener = Box>>; #[derive(Clone, Copy, PartialEq, Eq)] enum RpcServerKind { @@ -165,7 +170,7 @@ async fn set_tun_fd(fd: i32) -> Result<(), String> { .next() { instance_manager - .set_tun_fd(&uuid, fd) + .attach_tun_fd(uuid, fd) .map_err(|e| e.to_string())?; } Ok(()) @@ -372,6 +377,21 @@ fn normalize_normal_mode_rpc_portal(portal: &str) -> Result<(url::Url, url::Url) Ok((bind_url, connect_url)) } +async fn resolve_rpc_bind_url(url: &url::Url) -> Result { + if url.scheme() != "tcp" { + return Err(format!("RPC portal requires tcp URL: {url}")); + } + let host = url + .host_str() + .ok_or_else(|| format!("RPC portal has no host: {url}"))?; + let port = url.port().unwrap_or(11010); + tokio::net::lookup_host((host, port)) + .await + .map_err(|error| format!("failed to resolve RPC portal {url}: {error}"))? + .next() + .ok_or_else(|| format!("RPC portal has no resolved address: {url}")) +} + #[tauri::command] async fn init_rpc_connection( _app: AppHandle, @@ -390,11 +410,12 @@ async fn init_rpc_connection( .map_err(|_| "Failed to acquire lock for rpc server")?; let mut client_url = url.clone(); + let mut local_process_runtime = None; if is_normal_mode { let instance_manager = if let Some(im) = instance_manager_guard.take() { im } else { - Arc::new(NetworkInstanceManager::new()) + Arc::new(native_instance_manager()) }; let portal = url.and_then(|s| { @@ -422,12 +443,14 @@ async fn init_rpc_connection( *rpc_server_guard = None; let tunnel: BoxedTunnelListener = match desired_kind { - RpcServerKind::Ring => Box::new(RingTunnelListener::new( - format!("ring://{}", RPC_RING_UUID.deref()).parse().unwrap(), - )), - RpcServerKind::Tcp => Box::new(TcpTunnelListener::new( - bind_url.clone().expect("tcp rpc must have bind url"), - )), + RpcServerKind::Ring => instance_manager + .process_runtime() + .bind_ring_tunnel(*RPC_RING_UUID.deref()) + .map_err(|error| error.to_string())?, + RpcServerKind::Tcp => { + let bind_url = bind_url.as_ref().expect("tcp rpc must have bind url"); + Box::new(runtime_rpc_listener(resolve_rpc_bind_url(bind_url).await?)) + } }; let rpc_server = ApiRpcServer::from_tunnel(tunnel, instance_manager.clone()) @@ -442,6 +465,7 @@ async fn init_rpc_connection( }); } + local_process_runtime = Some(instance_manager.process_runtime()); *instance_manager_guard = Some(instance_manager); client_url = connect_url.map(|u| u.to_string()); } else { @@ -450,7 +474,7 @@ async fn init_rpc_connection( let client_manager = tokio::time::timeout( std::time::Duration::from_millis(1000), - manager::GUIClientManager::new(client_url), + manager::GUIClientManager::new(client_url, local_process_runtime), ) .await .map_err(|_| "connect remote rpc timed out".to_string())? @@ -462,7 +486,8 @@ async fn init_rpc_connection( drop(WEB_CLIENT.write().await.take()); if let Some(instance_manager) = instance_manager_guard.take() { instance_manager - .retain_network_instance(vec![]) + .retain_network_instances(&[]) + .await .map_err(|e| e.to_string())?; drop(instance_manager); } @@ -600,19 +625,13 @@ mod manager { use super::*; use async_trait::async_trait; use dashmap::{DashMap, DashSet}; - use easytier::common::global_ctx::GlobalCtx; - use easytier::common::stun::MockStunInfoCollector; - use easytier::launcher::NetworkConfig; + use easytier::common::config::{NetworkConfig, NetworkConfigExt}; use easytier::proto::api::logger::{LoggerRpc, LoggerRpcClientFactory, SetLoggerConfigRequest}; use easytier::proto::api::manage::RunNetworkInstanceRequest; - use easytier::proto::common::NatType; - use easytier::proto::rpc_impl::bidirect::BidirectRpcManager; + use easytier::proto::rpc::bidirect::BidirectRpcManager; use easytier::proto::rpc_types::controller::BaseController; - use easytier::rpc_service::logger::LoggerRpcService; - use easytier::rpc_service::remote_client::PersistentConfig; - use easytier::tunnel::TunnelConnector; - use easytier::tunnel::ring::RingTunnelConnector; use easytier::web_client::WebClientHooks; + use easytier_core::management::remote_client::PersistentConfig; pub(super) struct GuiHooks { pub(super) app: AppHandle, @@ -877,27 +896,16 @@ mod manager { pub(super) rpc_manager: BidirectRpcManager, } impl GUIClientManager { - pub async fn new(rpc_url: Option) -> Result { - let global_ctx = Arc::new(GlobalCtx::new(TomlConfigLoader::default())); - global_ctx.replace_stun_info_collector(Box::new(MockStunInfoCollector { - udp_nat_type: NatType::Unknown, - })); - let mut flags = global_ctx.get_flags(); - flags.bind_device = false; - global_ctx.set_flags(flags); + pub async fn new( + rpc_url: Option, + local_process_runtime: Option>, + ) -> Result { let tunnel = if let Some(url) = rpc_url { - let mut connector = easytier::connector::create_connector_by_url( - &url, - &global_ctx, - easytier::tunnel::IpVersion::Both, - ) - .await?; - connector.connect().await? + runtime_rpc_dialer(url.parse()?).connect().await? } else { - let mut connector = RingTunnelConnector::new( - format!("ring://{}", RPC_RING_UUID.deref()).parse().unwrap(), - ); - connector.connect().await? + local_process_runtime + .context("local RPC requires a core process runtime")? + .connect_ring_tunnel(*RPC_RING_UUID.deref())? }; let rpc_manager = BidirectRpcManager::new(); @@ -936,7 +944,7 @@ mod manager { &self, app: &AppHandle, web_only: bool, - ) -> Result<(), easytier::rpc_service::remote_client::RemoteClientError> + ) -> Result<(), easytier_core::management::remote_client::RemoteClientError> { let inst_ids: Vec = if web_only { self.get_enabled_instances_with_web_like_tun_ids().collect() @@ -1011,11 +1019,8 @@ mod manager { #[cfg(target_os = "android")] if let Some(instance_manager) = super::INSTANCE_MANAGER.read().await.as_ref() { let instance_uuid = *instance_id; - if let Some(instance_ref) = instance_manager - .iter() - .find(|item| *item.key() == instance_uuid) - { - if let Some(mut event_receiver) = instance_ref.value().subscribe_event() { + if let Some(instance) = instance_manager.instance(instance_uuid) { + if let Some(mut event_receiver) = subscribe_native_instance_event(&instance) { let app_clone = app.clone(); let instance_id_clone = *instance_id; tokio::spawn(async move { @@ -1090,7 +1095,7 @@ mod manager { .set_logger_config( BaseController::default(), SetLoggerConfigRequest { - level: LoggerRpcService::string_to_log_level(&level).into(), + level: easytier_core::management::parse_log_level(&level).into(), }, ) .await?; @@ -1139,7 +1144,7 @@ mod manager { inst_id: None, config: Some(config), overwrite: false, - source: source.to_runtime_source().to_rpc(), + source: config_source_to_rpc(source.to_runtime_source()), }, ) .await?; diff --git a/easytier-gui/src/auto-imports.d.ts b/easytier-gui/src/auto-imports.d.ts index 25d72636..b246d2dd 100644 --- a/easytier-gui/src/auto-imports.d.ts +++ b/easytier-gui/src/auto-imports.d.ts @@ -52,6 +52,7 @@ declare global { const mapWritableState: typeof import('pinia')['mapWritableState'] const markRaw: typeof import('vue')['markRaw'] const nextTick: typeof import('vue')['nextTick'] + const normalizeConfigSource: typeof import('./composables/config_source')['normalizeConfigSource'] const onActivated: typeof import('vue')['onActivated'] const onBeforeMount: typeof import('vue')['onBeforeMount'] const onBeforeRouteLeave: typeof import('vue-router')['onBeforeRouteLeave'] @@ -177,6 +178,7 @@ declare module 'vue' { readonly mapWritableState: UnwrapRef readonly markRaw: UnwrapRef readonly nextTick: UnwrapRef + readonly normalizeConfigSource: UnwrapRef readonly onActivated: UnwrapRef readonly onBeforeMount: UnwrapRef readonly onBeforeRouteLeave: UnwrapRef diff --git a/easytier-proto/Cargo.toml b/easytier-proto/Cargo.toml new file mode 100644 index 00000000..e69b5d35 --- /dev/null +++ b/easytier-proto/Cargo.toml @@ -0,0 +1,101 @@ +[package] +name = "easytier-proto" +description = "EasyTier protobuf and generated RPC types." +homepage = "https://github.com/EasyTier/EasyTier" +repository = "https://github.com/EasyTier/EasyTier" +version = "2.6.4" +edition.workspace = true +rust-version.workspace = true +authors = ["kkrainbow"] +keywords = ["vpn", "p2p", "network", "easytier"] +categories = ["network-programming"] +license-file = "../LICENSE" +build = "build/main.rs" + +[dependencies] +anyhow = { version = "1.0", optional = true } +async-trait = { version = "0.1.74", optional = true } +auto_impl = { version = "1.1.0", optional = true } +base64 = { version = "0.22", optional = true } +bytes = { version = "1.5.0", optional = true } +chrono = { version = "0.4.37", features = ["serde"], optional = true } +cidr = { version = "0.3.1", features = ["serde"], optional = true } +hmac = { version = "0.12.1", optional = true } +prost = { version = "0.14.3", optional = true } +prost-types = { version = "0.14.3", optional = true } +prost-wkt-types = { version = "0.7.1", optional = true } +pbjson = { version = "0.9.0", optional = true } +serde = { version = "1.0", features = ["derive"], optional = true } +serde_json = { version = "1", optional = true } +sha2 = { version = "0.10.8", optional = true } +thiserror = { version = "1.0", optional = true } +tokio = { version = "1", default-features = false, features = ["time"], optional = true } +url = { version = "2.5", features = ["serde"], optional = true } +uuid = { version = "1.5.0", features = ["serde"], optional = true } +x25519-dalek = { version = "2.0", features = ["static_secrets"], optional = true } + +[build-dependencies] +indoc = "2.0" +pbjson-build = "0.9.0" +proc-macro2 = "1" +prost-build = "0.14.3" +quote = "1" + +[target.'cfg(windows)'.build-dependencies] +reqwest = { version = "0.12.12", features = ["blocking"] } +zip = "4.0.0" + +[dev-dependencies] +tokio = { version = "1", default-features = false, features = [ + "macros", + "rt", +] } +uuid = { version = "1.5.0", features = [ + "v4", + "fast-rng", +] } + +[features] +default = ["full"] +full = [ + "api", + "core", + "faketcp", + "magic-dns", + "quic", + "utils", + "websocket", + "wireguard", + "zstd", + "json-rpc", +] +api = ["core"] +core = [ + "dep:anyhow", + "dep:async-trait", + "dep:auto_impl", + "dep:base64", + "dep:bytes", + "dep:chrono", + "dep:cidr", + "dep:hmac", + "dep:pbjson", + "dep:prost", + "dep:prost-types", + "dep:serde", + "dep:serde_json", + "dep:sha2", + "dep:thiserror", + "dep:tokio", + "dep:url", + "dep:uuid", + "dep:x25519-dalek", +] +faketcp = [] +magic-dns = [] +quic = [] +utils = [] +websocket = [] +wireguard = [] +zstd = [] +json-rpc = ["dep:prost-wkt-types"] diff --git a/easytier-proto/build/main.rs b/easytier-proto/build/main.rs new file mode 100644 index 00000000..b8b75c27 --- /dev/null +++ b/easytier-proto/build/main.rs @@ -0,0 +1,142 @@ +mod rpc; + +use crate::rpc::ServiceGenerator; +use std::{env, path::PathBuf}; + +#[cfg(target_os = "windows")] +use std::io::Cursor; + +#[cfg(target_os = "windows")] +fn check_protoc_exist() -> Option { + let path = env::var_os("PROTOC").map(PathBuf::from); + if path.is_some() && path.as_ref().unwrap().exists() { + return path; + } + + let path = env::var_os("PATH").unwrap_or_default(); + for p in env::split_paths(&path) { + let p = p.join("protoc.exe"); + if p.exists() && p.is_file() { + return Some(p); + } + } + + None +} + +#[cfg(target_os = "windows")] +fn get_cargo_target_dir() -> Result> { + let out_dir = PathBuf::from(env::var("OUT_DIR")?); + let profile = env::var("PROFILE")?; + let mut target_dir = None; + let mut sub_path = out_dir.as_path(); + while let Some(parent) = sub_path.parent() { + if parent.ends_with(&profile) { + target_dir = Some(parent); + break; + } + sub_path = parent; + } + let target_dir = target_dir.ok_or("not found")?; + Ok(target_dir.to_path_buf()) +} + +#[cfg(target_os = "windows")] +fn download_protoc() -> PathBuf { + let out_dir = get_cargo_target_dir().unwrap().join("protobuf"); + let fname = out_dir.join("bin/protoc.exe"); + if fname.exists() { + println!("cargo:info=use existing protoc: {:?}", fname); + return fname; + } + + println!("cargo:info=need download protoc, please wait..."); + + let url = "https://github.com/protocolbuffers/protobuf/releases/download/v26.0-rc1/protoc-26.0-rc-1-win64.zip"; + let response = reqwest::blocking::get(url).unwrap(); + println!("{:?}", response); + let mut content = response + .bytes() + .map(|v| v.to_vec()) + .map(Cursor::new) + .map(zip::ZipArchive::new) + .unwrap() + .unwrap(); + content.extract(out_dir).unwrap(); + + fname +} + +#[cfg(target_os = "windows")] +fn ensure_protoc_for_windows() { + let protoc_path = if let Some(path) = check_protoc_exist() { + println!("cargo:info=use os existing protoc: {:?}", path); + path + } else { + download_protoc() + }; + + unsafe { + env::set_var("PROTOC", protoc_path); + } +} + +fn main() -> Result<(), Box> { + #[cfg(target_os = "windows")] + ensure_protoc_for_windows(); + + let proto_files_reflect = ["proto/peer_rpc.proto", "proto/common.proto"]; + + let proto_files = [ + "proto/core_peer.proto", + "proto/core_config.proto", + "proto/error.proto", + "proto/tests.proto", + "proto/api_instance.proto", + "proto/api_logger.proto", + "proto/api_config.proto", + "proto/api_manage.proto", + "proto/web.proto", + "proto/magic_dns.proto", + "proto/acl.proto", + ]; + + for proto_file in proto_files.iter().chain(proto_files_reflect.iter()) { + println!("cargo:rerun-if-changed={proto_file}"); + } + + let out = PathBuf::from(env::var("OUT_DIR")?); + let descriptor = out.join("descriptors.bin"); + + let mut config = prost_build::Config::new(); + if env::var_os("CARGO_FEATURE_JSON_RPC").is_some() { + config + .extern_path(".google.protobuf.Any", "::prost_wkt_types::Any") + .extern_path(".google.protobuf.Timestamp", "::prost_wkt_types::Timestamp") + .extern_path(".google.protobuf.Value", "::prost_wkt_types::Value"); + } else { + config + .extern_path(".google.protobuf.Any", "::prost_types::Any") + .extern_path(".google.protobuf.Timestamp", "::prost_types::Timestamp") + .extern_path(".google.protobuf.Value", "::prost_types::Value"); + } + config + .file_descriptor_set_path(&descriptor) + .service_generator(Box::new(ServiceGenerator::default())) + .btree_map(["."]) + .skip_debug([".common.Ipv4Addr", ".common.Ipv6Addr", ".common.UUID"]); + + config.compile_protos(&proto_files, &["proto/"])?; + + config.file_descriptor_set_path(out.join("file_descriptor_set.bin")); + config.compile_protos(&proto_files_reflect, &["proto/"])?; + + let descriptor = std::fs::read(descriptor)?; + pbjson_build::Builder::new() + .register_descriptors(&descriptor)? + .preserve_proto_field_names() + .btree_map(["."]) + .build(&["."])?; + + Ok(()) +} diff --git a/easytier/build/rpc.rs b/easytier-proto/build/rpc.rs similarity index 99% rename from easytier/build/rpc.rs rename to easytier-proto/build/rpc.rs index e8d4472d..8b579822 100644 --- a/easytier/build/rpc.rs +++ b/easytier-proto/build/rpc.rs @@ -155,6 +155,7 @@ impl Service { #(#methods)* + #[cfg(feature = "json-rpc")] async fn json_call_method( &self, ctrl: Self::Controller, diff --git a/easytier/src/proto/acl.proto b/easytier-proto/proto/acl.proto similarity index 100% rename from easytier/src/proto/acl.proto rename to easytier-proto/proto/acl.proto diff --git a/easytier/src/proto/api_config.proto b/easytier-proto/proto/api_config.proto similarity index 100% rename from easytier/src/proto/api_config.proto rename to easytier-proto/proto/api_config.proto diff --git a/easytier/src/proto/api_instance.proto b/easytier-proto/proto/api_instance.proto similarity index 100% rename from easytier/src/proto/api_instance.proto rename to easytier-proto/proto/api_instance.proto diff --git a/easytier/src/proto/api_logger.proto b/easytier-proto/proto/api_logger.proto similarity index 100% rename from easytier/src/proto/api_logger.proto rename to easytier-proto/proto/api_logger.proto diff --git a/easytier/src/proto/api_manage.proto b/easytier-proto/proto/api_manage.proto similarity index 100% rename from easytier/src/proto/api_manage.proto rename to easytier-proto/proto/api_manage.proto diff --git a/easytier/src/proto/common.proto b/easytier-proto/proto/common.proto similarity index 100% rename from easytier/src/proto/common.proto rename to easytier-proto/proto/common.proto diff --git a/easytier-proto/proto/core_config.proto b/easytier-proto/proto/core_config.proto new file mode 100644 index 00000000..14d70936 --- /dev/null +++ b/easytier-proto/proto/core_config.proto @@ -0,0 +1,56 @@ +syntax = "proto3"; + +import "common.proto"; + +package core_config; + +message CoreConfig { + NodeConfig node = 1; + RouteConfig routes = 2; + PeerPolicyConfig peer_policy = 3; + TrafficConfig traffic = 4; +} + +message NodeConfig { + optional uint32 peer_id = 1; + optional common.UUID instance_id = 2; + optional string hostname = 3; + string network_name = 4; +} + +message RouteConfig { + optional IpPrefix ipv4 = 1; + optional IpPrefix ipv6 = 2; + repeated IpPrefix advertised_routes = 3; + repeated ProxyNetworkConfig proxy_networks = 4; + repeated ForeignNetworkConfig foreign_networks = 5; +} + +message IpPrefix { + common.IpAddr address = 1; + uint32 prefix_len = 2; +} + +message ProxyNetworkConfig { + IpPrefix real = 1; + optional IpPrefix mapped = 2; +} + +message ForeignNetworkConfig { + string name = 1; + repeated IpPrefix cidrs = 2; +} + +message PeerPolicyConfig { + optional bool p2p_enabled = 1; + optional bool relay_peer_rpc = 2; + optional bool relay_data = 3; + optional bool latency_first = 4; + optional bool encryption_required = 5; +} + +message TrafficConfig { + optional uint32 mtu = 1; + optional uint64 instance_recv_bps_limit = 2; + optional uint64 foreign_relay_bps_limit = 3; +} diff --git a/easytier-proto/proto/core_peer.proto b/easytier-proto/proto/core_peer.proto new file mode 100644 index 00000000..289966ee --- /dev/null +++ b/easytier-proto/proto/core_peer.proto @@ -0,0 +1,77 @@ +syntax = "proto3"; + +import "common.proto"; +import "peer_rpc.proto"; + +package core.peer; + +message PeerConnStats { + uint64 rx_bytes = 1; + uint64 tx_bytes = 2; + + uint64 rx_packets = 3; + uint64 tx_packets = 4; + + uint64 latency_us = 5; +} + +message PeerConnInfo { + string conn_id = 1; + uint32 my_peer_id = 2; + uint32 peer_id = 3; + repeated string features = 4; + common.TunnelInfo tunnel = 5; + PeerConnStats stats = 6; + float loss_rate = 7; + bool is_client = 8; + string network_name = 9; + bool is_closed = 10; + bytes noise_local_static_pubkey = 11; + bytes noise_remote_static_pubkey = 12; + peer_rpc.SecureAuthLevel secure_auth_level = 13; + peer_rpc.PeerIdentityType peer_identity_type = 14; +} + +message PeerInfo { + uint32 peer_id = 1; + repeated PeerConnInfo conns = 2; + common.UUID default_conn_id = 3; + repeated common.UUID directly_connected_conns = 4; +} + +message Route { + uint32 peer_id = 1; + common.Ipv4Inet ipv4_addr = 2; + + uint32 next_hop_peer_id = 3; + int32 cost = 4; + int32 path_latency = 11; + + repeated string proxy_cidrs = 5; + string hostname = 6; + common.StunInfo stun_info = 7; + string inst_id = 8; + string version = 9; + common.PeerFeatureFlag feature_flag = 10; + + optional uint32 next_hop_peer_id_latency_first = 12; + optional int32 cost_latency_first = 13; + optional int32 path_latency_latency_first = 14; + + common.Ipv6Inet ipv6_addr = 15; + common.Ipv6Inet public_ipv6_addr = 16; + common.Ipv6Inet ipv6_public_addr_prefix = 17; +} + +message PublicIpv6LeaseInfo { + uint32 peer_id = 1; + string inst_id = 2; + common.Ipv6Inet leased_addr = 3; + int64 valid_until_unix_seconds = 4; + bool reused = 5; +} + +message ListPublicIpv6InfoResponse { + common.Ipv6Inet provider_prefix = 1; + repeated PublicIpv6LeaseInfo provider_leases = 2; +} diff --git a/easytier/src/proto/error.proto b/easytier-proto/proto/error.proto similarity index 100% rename from easytier/src/proto/error.proto rename to easytier-proto/proto/error.proto diff --git a/easytier/src/proto/magic_dns.proto b/easytier-proto/proto/magic_dns.proto similarity index 100% rename from easytier/src/proto/magic_dns.proto rename to easytier-proto/proto/magic_dns.proto diff --git a/easytier/src/proto/peer_rpc.proto b/easytier-proto/proto/peer_rpc.proto similarity index 100% rename from easytier/src/proto/peer_rpc.proto rename to easytier-proto/proto/peer_rpc.proto diff --git a/easytier/src/proto/tests.proto b/easytier-proto/proto/tests.proto similarity index 100% rename from easytier/src/proto/tests.proto rename to easytier-proto/proto/tests.proto diff --git a/easytier/src/proto/web.proto b/easytier-proto/proto/web.proto similarity index 100% rename from easytier/src/proto/web.proto rename to easytier-proto/proto/web.proto diff --git a/easytier/src/proto/acl.rs b/easytier-proto/src/acl.rs similarity index 98% rename from easytier/src/proto/acl.rs rename to easytier-proto/src/acl.rs index a4cf5854..4d086b93 100644 --- a/easytier/src/proto/acl.rs +++ b/easytier-proto/src/acl.rs @@ -23,6 +23,7 @@ impl GroupInfo { } } +#[cfg(feature = "api")] impl Display for ConnTrackEntry { fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { let src = self @@ -87,6 +88,7 @@ impl Display for StatItem { } } +#[cfg(feature = "api")] impl Display for AclStats { fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { writeln!(f, "AclStats:")?; diff --git a/easytier/src/proto/api.rs b/easytier-proto/src/api.rs similarity index 74% rename from easytier/src/proto/api.rs rename to easytier-proto/src/api.rs index 010ed25b..fcb6d97b 100644 --- a/easytier/src/proto/api.rs +++ b/easytier-proto/src/api.rs @@ -1,5 +1,6 @@ pub mod config { include!(concat!(env!("OUT_DIR"), "/api.config.rs")); + #[cfg(feature = "json-rpc")] include!(concat!(env!("OUT_DIR"), "/api.config.serde.rs")); pub struct Patchable { @@ -7,15 +8,6 @@ pub mod config { pub value: Option, } - impl From for Patchable { - fn from(patch: PortForwardPatch) -> Self { - Patchable { - action: ConfigPatchAction::try_from(patch.action).ok(), - value: patch.cfg.map(Into::into), - } - } - } - impl From for Patchable { fn from(value: RoutePatch) -> Self { Patchable { @@ -78,9 +70,111 @@ pub mod config { } pub mod instance { + use std::fmt::{Display, Formatter}; + include!(concat!(env!("OUT_DIR"), "/api.instance.rs")); + #[cfg(feature = "json-rpc")] include!(concat!(env!("OUT_DIR"), "/api.instance.serde.rs")); + impl From for PeerConnStats { + fn from(value: crate::core_peer::peer::PeerConnStats) -> Self { + Self { + rx_bytes: value.rx_bytes, + tx_bytes: value.tx_bytes, + rx_packets: value.rx_packets, + tx_packets: value.tx_packets, + latency_us: value.latency_us, + } + } + } + + impl From for PeerConnInfo { + fn from(value: crate::core_peer::peer::PeerConnInfo) -> Self { + Self { + conn_id: value.conn_id, + my_peer_id: value.my_peer_id, + peer_id: value.peer_id, + features: value.features, + tunnel: value.tunnel, + stats: value.stats.map(Into::into), + loss_rate: value.loss_rate, + is_client: value.is_client, + network_name: value.network_name, + is_closed: value.is_closed, + noise_local_static_pubkey: value.noise_local_static_pubkey, + noise_remote_static_pubkey: value.noise_remote_static_pubkey, + secure_auth_level: value.secure_auth_level, + peer_identity_type: value.peer_identity_type, + } + } + } + + impl From for PeerInfo { + fn from(value: crate::core_peer::peer::PeerInfo) -> Self { + Self { + peer_id: value.peer_id, + conns: value.conns.into_iter().map(Into::into).collect(), + default_conn_id: value.default_conn_id, + directly_connected_conns: value.directly_connected_conns, + } + } + } + + impl From for Route { + fn from(value: crate::core_peer::peer::Route) -> Self { + Self { + peer_id: value.peer_id, + ipv4_addr: value.ipv4_addr, + next_hop_peer_id: value.next_hop_peer_id, + cost: value.cost, + path_latency: value.path_latency, + proxy_cidrs: value.proxy_cidrs, + hostname: value.hostname, + stun_info: value.stun_info, + inst_id: value.inst_id, + version: value.version, + feature_flag: value.feature_flag, + next_hop_peer_id_latency_first: value.next_hop_peer_id_latency_first, + cost_latency_first: value.cost_latency_first, + path_latency_latency_first: value.path_latency_latency_first, + ipv6_addr: value.ipv6_addr, + public_ipv6_addr: value.public_ipv6_addr, + ipv6_public_addr_prefix: value.ipv6_public_addr_prefix, + } + } + } + + impl From for PublicIpv6LeaseInfo { + fn from(value: crate::core_peer::peer::PublicIpv6LeaseInfo) -> Self { + Self { + peer_id: value.peer_id, + inst_id: value.inst_id, + leased_addr: value.leased_addr, + valid_until_unix_seconds: value.valid_until_unix_seconds, + reused: value.reused, + } + } + } + + impl From for ListPublicIpv6InfoResponse { + fn from(value: crate::core_peer::peer::ListPublicIpv6InfoResponse) -> Self { + Self { + provider_prefix: value.provider_prefix, + provider_leases: value.provider_leases.into_iter().map(Into::into).collect(), + } + } + } + + impl Display for PeerConnInfo { + fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result { + f.debug_struct("PeerConnInfo") + .field("my_peer_id", &self.my_peer_id) + .field("dst_peer_id", &self.peer_id) + .field("tunnel_info", &self.tunnel) + .finish() + } + } + impl PeerRoutePair { pub fn get_latency_ms(&self) -> Option { let mut ret = u64::MAX; @@ -232,11 +326,13 @@ pub mod instance { pub mod logger { include!(concat!(env!("OUT_DIR"), "/api.logger.rs")); + #[cfg(feature = "json-rpc")] include!(concat!(env!("OUT_DIR"), "/api.logger.serde.rs")); } pub mod manage { include!(concat!(env!("OUT_DIR"), "/api.manage.rs")); + #[cfg(feature = "json-rpc")] include!(concat!(env!("OUT_DIR"), "/api.manage.serde.rs")); } diff --git a/easytier/src/proto/common.rs b/easytier-proto/src/common.rs similarity index 93% rename from easytier/src/proto/common.rs rename to easytier-proto/src/common.rs index df955735..c2ad4d80 100644 --- a/easytier/src/proto/common.rs +++ b/easytier-proto/src/common.rs @@ -5,9 +5,8 @@ use std::{ fmt::{self, Display}, str::FromStr, }; -use strum::VariantArray; -use crate::tunnel::{IpScheme, packet_def::CompressorAlgo}; +const IP_SCHEMES: &[&str] = &["tcp", "udp", "wg", "quic", "ws", "wss", "faketcp"]; include!(concat!(env!("OUT_DIR"), "/common.rs")); include!(concat!(env!("OUT_DIR"), "/common.serde.rs")); @@ -16,7 +15,12 @@ pub trait TimestampExt { fn now() -> Self; } -impl TimestampExt for prost_wkt_types::Timestamp { +#[cfg(feature = "json-rpc")] +pub type RuntimeTimestamp = prost_wkt_types::Timestamp; +#[cfg(not(feature = "json-rpc"))] +pub type RuntimeTimestamp = prost_types::Timestamp; + +impl TimestampExt for RuntimeTimestamp { fn now() -> Self { SystemTime::now().into() } @@ -301,8 +305,7 @@ impl fmt::Display for Url { } fn split_tunnel_scheme(raw_scheme: &str) -> Option<(&str, &'static str, bool)> { - for scheme in IpScheme::VARIANTS { - let scheme: &'static str = scheme.into(); + for &scheme in IP_SCHEMES { if let Some(base) = raw_scheme.strip_suffix('6') && let Some(prefix) = base.strip_suffix(scheme) && (prefix.is_empty() || prefix.ends_with('-')) @@ -471,31 +474,6 @@ impl TunnelInfo { } } -impl TryFrom for CompressorAlgo { - type Error = anyhow::Error; - - fn try_from(value: CompressionAlgoPb) -> Result { - match value { - #[cfg(feature = "zstd")] - CompressionAlgoPb::Zstd => Ok(CompressorAlgo::ZstdDefault), - CompressionAlgoPb::None => Ok(CompressorAlgo::None), - _ => Err(anyhow::anyhow!("Invalid CompressionAlgoPb")), - } - } -} - -impl TryFrom for CompressionAlgoPb { - type Error = anyhow::Error; - - fn try_from(value: CompressorAlgo) -> Result { - match value { - #[cfg(feature = "zstd")] - CompressorAlgo::ZstdDefault => Ok(CompressionAlgoPb::Zstd), - CompressorAlgo::None => Ok(CompressionAlgoPb::None), - } - } -} - impl fmt::Debug for Ipv4Addr { fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result { let std_ipv4_addr = std::net::Ipv4Addr::from(*self); @@ -570,23 +548,13 @@ mod tests { } #[test] - fn normalize_all_enabled_ipv6_tunnel_urls() { + fn normalize_all_ipv6_tunnel_urls() { assert_ipv6_tunnel_normalization("tcp", 11010); assert_ipv6_tunnel_normalization("udp", 11010); - - #[cfg(feature = "wireguard")] assert_ipv6_tunnel_normalization("wg", 11011); - - #[cfg(feature = "quic")] assert_ipv6_tunnel_normalization("quic", 11012); - - #[cfg(feature = "websocket")] assert_ipv6_tunnel_normalization("ws", 80); - - #[cfg(feature = "websocket")] assert_ipv6_tunnel_normalization("wss", 443); - - #[cfg(feature = "faketcp")] assert_ipv6_tunnel_normalization("faketcp", 11013); } diff --git a/easytier-proto/src/core_config.rs b/easytier-proto/src/core_config.rs new file mode 100644 index 00000000..1076a787 --- /dev/null +++ b/easytier-proto/src/core_config.rs @@ -0,0 +1,3 @@ +include!(concat!(env!("OUT_DIR"), "/core_config.rs")); +#[cfg(feature = "json-rpc")] +include!(concat!(env!("OUT_DIR"), "/core_config.serde.rs")); diff --git a/easytier-proto/src/core_peer.rs b/easytier-proto/src/core_peer.rs new file mode 100644 index 00000000..18fca641 --- /dev/null +++ b/easytier-proto/src/core_peer.rs @@ -0,0 +1,5 @@ +pub mod peer { + include!(concat!(env!("OUT_DIR"), "/core.peer.rs")); + #[cfg(feature = "json-rpc")] + include!(concat!(env!("OUT_DIR"), "/core.peer.serde.rs")); +} diff --git a/easytier/src/proto/error.rs b/easytier-proto/src/error.rs similarity index 100% rename from easytier/src/proto/error.rs rename to easytier-proto/src/error.rs diff --git a/easytier-proto/src/lib.rs b/easytier-proto/src/lib.rs new file mode 100644 index 00000000..e99019fa --- /dev/null +++ b/easytier-proto/src/lib.rs @@ -0,0 +1,33 @@ +#[cfg(feature = "core")] +pub mod rpc_types; + +#[cfg(feature = "core")] +pub mod acl; +#[cfg(feature = "api")] +pub mod api; +#[cfg(feature = "core")] +pub mod common; +#[cfg(feature = "core")] +pub mod core_config; +#[cfg(feature = "core")] +pub mod core_peer; +#[cfg(feature = "core")] +pub mod error; +#[cfg(all(feature = "api", feature = "magic-dns"))] +pub mod magic_dns; +#[cfg(feature = "core")] +pub mod peer_rpc; +#[cfg(feature = "api")] +pub mod tests; +#[cfg(feature = "api")] +pub mod web; + +pub const DESCRIPTOR_POOL_BYTES: &[u8] = + include_bytes!(concat!(env!("OUT_DIR"), "/file_descriptor_set.bin")); + +pub const ALL_DESCRIPTOR_BYTES: &[u8] = + include_bytes!(concat!(env!("OUT_DIR"), "/descriptors.bin")); + +pub mod proto { + pub use crate::*; +} diff --git a/easytier/src/proto/magic_dns.rs b/easytier-proto/src/magic_dns.rs similarity index 79% rename from easytier/src/proto/magic_dns.rs rename to easytier-proto/src/magic_dns.rs index f78dd08a..bac0b1aa 100644 --- a/easytier/src/proto/magic_dns.rs +++ b/easytier-proto/src/magic_dns.rs @@ -1,2 +1,3 @@ include!(concat!(env!("OUT_DIR"), "/magic_dns.rs")); +#[cfg(feature = "json-rpc")] include!(concat!(env!("OUT_DIR"), "/magic_dns.serde.rs")); diff --git a/easytier/src/proto/peer_rpc.rs b/easytier-proto/src/peer_rpc.rs similarity index 72% rename from easytier/src/proto/peer_rpc.rs rename to easytier-proto/src/peer_rpc.rs index 142129c7..93a8dae1 100644 --- a/easytier/src/proto/peer_rpc.rs +++ b/easytier-proto/src/peer_rpc.rs @@ -1,10 +1,14 @@ use hmac::{Hmac, Mac}; use prost::Message; use sha2::Sha256; +#[cfg(feature = "api")] +use std::collections::BTreeMap; +use std::collections::BTreeSet; -use crate::common::PeerId; +type PeerId = u32; include!(concat!(env!("OUT_DIR"), "/peer_rpc.rs")); +#[cfg(feature = "json-rpc")] include!(concat!(env!("OUT_DIR"), "/peer_rpc.serde.rs")); impl PeerGroupInfo { @@ -103,6 +107,140 @@ impl From for sync_route_info_request::ConnInfo { } } +#[cfg(feature = "api")] +impl From> for PeerInfoForGlobalMap { + fn from(peers: Vec) -> Self { + let mut peer_map = BTreeMap::new(); + for peer in peers { + let Some(min_lat) = peer + .conns + .iter() + .map(|conn| conn.stats.as_ref().unwrap().latency_us) + .min() + else { + continue; + }; + + let dp_info = DirectConnectedPeerInfo { + latency_ms: std::cmp::max(1, (min_lat as u32 / 1000) as i32), + }; + + peer_map.insert(peer.peer_id, dp_info); + } + PeerInfoForGlobalMap { + direct_peers: peer_map, + } + } +} + +impl From for crate::core_peer::peer::Route { + fn from(val: RoutePeerInfo) -> Self { + let network_length = if val.network_length == 0 { + 24 + } else { + val.network_length + }; + + crate::core_peer::peer::Route { + peer_id: val.peer_id, + ipv4_addr: val.ipv4_addr.map(|ipv4_addr| crate::common::Ipv4Inet { + address: Some(ipv4_addr), + network_length, + }), + next_hop_peer_id: 0, + cost: 0, + path_latency: 0, + proxy_cidrs: val.proxy_cidrs.clone(), + hostname: val.hostname.unwrap_or_default(), + stun_info: { + let mut stun_info = crate::common::StunInfo::default(); + if let Ok(udp_nat_type) = crate::common::NatType::try_from(val.udp_nat_type) { + stun_info.set_udp_nat_type(udp_nat_type); + } + if let Ok(tcp_nat_type) = crate::common::NatType::try_from(val.tcp_nat_type) { + stun_info.set_tcp_nat_type(tcp_nat_type); + } + Some(stun_info) + }, + inst_id: val.inst_id.map(|x| x.to_string()).unwrap_or_default(), + version: val.easytier_version, + feature_flag: val.feature_flag, + + next_hop_peer_id_latency_first: None, + cost_latency_first: None, + path_latency_latency_first: None, + + ipv6_addr: val.ipv6_addr, + public_ipv6_addr: val.ipv6_public_addr_lease, + ipv6_public_addr_prefix: val.ipv6_public_addr_prefix, + } + } +} + +#[cfg(feature = "api")] +impl From for crate::api::instance::Route { + fn from(val: RoutePeerInfo) -> Self { + let network_length = if val.network_length == 0 { + 24 + } else { + val.network_length + }; + + crate::api::instance::Route { + peer_id: val.peer_id, + ipv4_addr: val.ipv4_addr.map(|ipv4_addr| crate::common::Ipv4Inet { + address: Some(ipv4_addr), + network_length, + }), + next_hop_peer_id: 0, + cost: 0, + path_latency: 0, + proxy_cidrs: val.proxy_cidrs.clone(), + hostname: val.hostname.unwrap_or_default(), + stun_info: { + let mut stun_info = crate::common::StunInfo::default(); + if let Ok(udp_nat_type) = crate::common::NatType::try_from(val.udp_nat_type) { + stun_info.set_udp_nat_type(udp_nat_type); + } + if let Ok(tcp_nat_type) = crate::common::NatType::try_from(val.tcp_nat_type) { + stun_info.set_tcp_nat_type(tcp_nat_type); + } + Some(stun_info) + }, + inst_id: val.inst_id.map(|x| x.to_string()).unwrap_or_default(), + version: val.easytier_version, + feature_flag: val.feature_flag, + + next_hop_peer_id_latency_first: None, + cost_latency_first: None, + path_latency_latency_first: None, + + ipv6_addr: val.ipv6_addr, + public_ipv6_addr: val.ipv6_public_addr_lease, + ipv6_public_addr_prefix: val.ipv6_public_addr_prefix, + } + } +} + +impl RouteConnBitmap { + pub fn get_bit(&self, idx: usize) -> bool { + let byte_idx = idx / 8; + let bit_idx = idx % 8; + let byte = self.bitmap[byte_idx]; + (byte >> bit_idx) & 1 == 1 + } + + pub fn get_connected_peers(&self, peer_idx: usize) -> BTreeSet { + let mut connected_peers = BTreeSet::new(); + for (idx, peer_id_version) in self.peer_ids.iter().enumerate() { + if self.get_bit(peer_idx * self.peer_ids.len() + idx) { + connected_peers.insert(peer_id_version.peer_id); + } + } + connected_peers + } +} + #[cfg(test)] mod tests { use super::*; diff --git a/easytier/src/proto/rpc_types/__rt.rs b/easytier-proto/src/rpc_types/__rt.rs similarity index 97% rename from easytier/src/proto/rpc_types/__rt.rs rename to easytier-proto/src/rpc_types/__rt.rs index 6c4e00bd..edb967e0 100644 --- a/easytier/src/proto/rpc_types/__rt.rs +++ b/easytier-proto/src/rpc_types/__rt.rs @@ -40,7 +40,6 @@ where I: prost::Message, O: prost::Message + Default, { - type Error = super::error::Error; let input_bytes = encode(input)?; let ret_msg = handler.call(ctrl, method, input_bytes).await?; decode(ret_msg) diff --git a/easytier/src/proto/rpc_types/controller.rs b/easytier-proto/src/rpc_types/controller.rs similarity index 100% rename from easytier/src/proto/rpc_types/controller.rs rename to easytier-proto/src/rpc_types/controller.rs diff --git a/easytier/src/proto/rpc_types/descriptor.rs b/easytier-proto/src/rpc_types/descriptor.rs similarity index 100% rename from easytier/src/proto/rpc_types/descriptor.rs rename to easytier-proto/src/rpc_types/descriptor.rs diff --git a/easytier/src/proto/rpc_types/error.rs b/easytier-proto/src/rpc_types/error.rs similarity index 95% rename from easytier/src/proto/rpc_types/error.rs rename to easytier-proto/src/rpc_types/error.rs index 8a4df29b..235d5c85 100644 --- a/easytier/src/proto/rpc_types/error.rs +++ b/easytier-proto/src/rpc_types/error.rs @@ -28,7 +28,7 @@ pub enum Error { Timeout(#[from] tokio::time::error::Elapsed), #[error("Tunnel error: {0}")] - TunnelError(#[from] crate::tunnel::TunnelError), + TunnelError(String), #[error("Shutdown")] Shutdown, diff --git a/easytier/src/proto/rpc_types/handler.rs b/easytier-proto/src/rpc_types/handler.rs similarity index 100% rename from easytier/src/proto/rpc_types/handler.rs rename to easytier-proto/src/rpc_types/handler.rs diff --git a/easytier/src/proto/rpc_types/mod.rs b/easytier-proto/src/rpc_types/mod.rs similarity index 100% rename from easytier/src/proto/rpc_types/mod.rs rename to easytier-proto/src/rpc_types/mod.rs diff --git a/easytier-proto/src/tests.rs b/easytier-proto/src/tests.rs new file mode 100644 index 00000000..e69f1034 --- /dev/null +++ b/easytier-proto/src/tests.rs @@ -0,0 +1,2 @@ +include!(concat!(env!("OUT_DIR"), "/tests.rs")); +include!(concat!(env!("OUT_DIR"), "/tests.serde.rs")); diff --git a/easytier/src/proto/web.rs b/easytier-proto/src/web.rs similarity index 100% rename from easytier/src/proto/web.rs rename to easytier-proto/src/web.rs diff --git a/easytier-web/Cargo.toml b/easytier-web/Cargo.toml index 6070ed6d..5744ff33 100644 --- a/easytier-web/Cargo.toml +++ b/easytier-web/Cargo.toml @@ -6,7 +6,8 @@ description = "Config server for easytier. easytier-core gets config from this a [dependencies] easytier = { path = "../easytier" } -tracing = { version = "0.1", features = ["log"] } +easytier-core = { path = "../easytier-core" } +tracing = "0.1" anyhow = { version = "1.0" } thiserror = "1.0" tokio = { version = "1", features = ["full"] } diff --git a/easytier-web/frontend-lib/scripts/codegen-proto.mjs b/easytier-web/frontend-lib/scripts/codegen-proto.mjs index 0a4315be..a5a446cc 100644 --- a/easytier-web/frontend-lib/scripts/codegen-proto.mjs +++ b/easytier-web/frontend-lib/scripts/codegen-proto.mjs @@ -6,7 +6,7 @@ import { fileURLToPath } from 'node:url' const require = createRequire(import.meta.url) const root = resolve(dirname(fileURLToPath(import.meta.url)), '..') -const protoRoot = resolve(root, '../../easytier/src/proto') +const protoRoot = resolve(root, '../../easytier-proto/proto') const generatedRoot = resolve(root, 'src/generated') const outDir = resolve(generatedRoot, 'proto') const nodeBinDir = resolve(root, 'node_modules/.bin') diff --git a/easytier-web/src/client_manager/managed_config.rs b/easytier-web/src/client_manager/managed_config.rs index 3320e6b7..e7be3822 100644 --- a/easytier-web/src/client_manager/managed_config.rs +++ b/easytier-web/src/client_manager/managed_config.rs @@ -11,7 +11,10 @@ use easytier::{ api::manage::{ConfigSource as RpcConfigSource, NetworkConfig, NetworkMeta}, common::Uuid as RpcUuid, }, - rpc_service::remote_client::{ListNetworkProps, PersistentConfig as _, Storage as _}, +}; +use easytier_core::management::config_source_from_rpc; +use easytier_core::management::remote_client::{ + ListNetworkProps, PersistentConfig as _, Storage as _, }; use super::storage::Storage; @@ -469,7 +472,7 @@ pub(super) async fn sync_running_config_sources( continue; }; - let Some(running_source) = ConfigSource::from_rpc(meta.source) else { + let Some(running_source) = config_source_from_rpc(meta.source) else { continue; }; let local_source = PersistedConfigSource::from_db(&local_cfg.source); @@ -503,9 +506,11 @@ mod tests { use std::collections::HashSet; use easytier::{ - common::config::{ConfigLoader as _, ConfigSource}, + common::config::{ConfigLoader as _, ConfigSource, NetworkConfigExt}, proto::api::manage::{ConfigSource as RpcConfigSource, NetworkConfig, NetworkMeta}, - rpc_service::remote_client::{ListNetworkProps, PersistentConfig as _, Storage as _}, + }; + use easytier_core::management::remote_client::{ + ListNetworkProps, PersistentConfig as _, Storage as _, }; use serde_json::json; diff --git a/easytier-web/src/client_manager/mod.rs b/easytier-web/src/client_manager/mod.rs index c0b59337..78eec408 100644 --- a/easytier-web/src/client_manager/mod.rs +++ b/easytier-web/src/client_manager/mod.rs @@ -10,13 +10,13 @@ use std::sync::{ use std::time::Duration; use dashmap::DashMap; -use easytier::{ - proto::{ - api::manage::WebClientService, rpc_types::controller::BaseController, web::HeartbeatRequest, - }, - rpc_service::remote_client::{self, RemoteClientManager}, - tunnel::TunnelListener, - web_client::security, +use easytier::proto::{ + api::manage::WebClientService, rpc_types::controller::BaseController, web::HeartbeatRequest, +}; +use easytier_core::{ + management::remote_client::{self, RemoteClientManager}, + socket::SocketListener, + tunnel::{Tunnel, web_security}, }; use maxminddb::geoip2; use session::{Location, Session}; @@ -105,11 +105,12 @@ impl ClientManager { } } - pub async fn add_listener( + pub async fn add_listener> + 'static>( &mut self, mut listener: L, - ) -> Result<(), anyhow::Error> { + ) -> Result { listener.listen().await?; + let local_url = listener.local_url(); self.listeners_cnt.fetch_add(1, Ordering::Relaxed); let sessions = self.client_sessions.clone(); let storage = self.storage.weak_ref(); @@ -120,7 +121,11 @@ impl ClientManager { let webhook_config = self.webhook_config.clone(); self.tasks.spawn(async move { while let Ok(tunnel) = listener.accept().await { - let (tunnel, secure) = match security::accept_or_upgrade_server_tunnel(tunnel).await { + let (tunnel, secure) = match web_security::accept_or_upgrade_server_tunnel( + tunnel, + ) + .await + { Ok(v) => v, Err(error) => { tracing::warn!(%error, "failed to accept secure tunnel, dropping connection"); @@ -150,7 +155,7 @@ impl ClientManager { listeners_cnt.fetch_sub(1, Ordering::Relaxed); }); - Ok(()) + Ok(local_url) } pub fn is_running(&self) -> bool { @@ -369,6 +374,7 @@ impl mod tests { use std::{ collections::VecDeque, + future::Future, sync::{ Arc, atomic::{AtomicBool, AtomicUsize, Ordering}, @@ -378,22 +384,18 @@ mod tests { use axum::{Json, Router, extract::State, routing::post}; use easytier::{ - common::MachineIdOptions, - instance_manager::NetworkInstanceManager, + common::{MachineIdOptions, config::NetworkConfigExt}, + instance::factory::{NativeInstanceManager, native_instance_manager}, proto::{ api::manage::{NetworkConfig, NetworkingMethod, PortForwardConfig}, common::CompressionAlgoPb, - }, - rpc_service::remote_client::Storage as RemoteStorage, - tunnel::{ - common::tests::wait_for_condition, - udp::{UdpTunnelConnector, UdpTunnelListener}, + rpc::standalone::{runtime_udp_tunnel_dialer, runtime_udp_tunnel_listener}, }, web_client::{WebClient, run_web_client}, }; + use easytier_core::management::remote_client::Storage as RemoteStorage; use serde_json::json; use sqlx::Executor; - use tokio::net::UdpSocket; use crate::{ FeatureFlags, client_manager::ClientManager, db::Db, webhook::ManagedNetworkConfig, @@ -401,6 +403,21 @@ mod tests { const MANAGED_CONFIG_TOKEN: &str = "managed-config-token"; + async fn wait_for_condition(mut condition: F, timeout: Duration) + where + F: FnMut() -> Fut, + Fut: Future, + { + let deadline = tokio::time::Instant::now() + timeout; + while !condition().await { + assert!( + tokio::time::Instant::now() < deadline, + "condition timed out" + ); + tokio::time::sleep(Duration::from_millis(50)).await; + } + } + #[derive(Debug, Clone)] struct TestWebhookState { validate_responses: Arc>>, @@ -510,12 +527,15 @@ mod tests { } async fn add_random_udp_listener(mgr: &mut ClientManager) -> std::net::SocketAddr { - let socket = Arc::new(UdpSocket::bind("127.0.0.1:0").await.unwrap()); - let addr = socket.local_addr().unwrap(); - let listener = - UdpTunnelListener::new_with_socket(format!("udp://{addr}").parse().unwrap(), socket); - mgr.add_listener(listener).await.unwrap(); - addr + let local_url = "udp://127.0.0.1:0".parse().unwrap(); + let listener = runtime_udp_tunnel_listener(local_url, "127.0.0.1:0".parse().unwrap()); + let local_url = mgr.add_listener(listener).await.unwrap(); + local_url + .socket_addrs(|| None) + .unwrap() + .into_iter() + .next() + .unwrap() } async fn wait_for_validated_user(mgr: &ClientManager, machine_id: uuid::Uuid) -> i32 { @@ -575,14 +595,14 @@ mod tests { } async fn wait_for_runtime_config( - manager: &NetworkInstanceManager, + manager: &NativeInstanceManager, inst_id: uuid::Uuid, predicate: impl Fn(&NetworkConfig) -> bool, ) -> NetworkConfig { tokio::time::timeout(Duration::from_secs(12), async { loop { if let Some(config) = manager - .get_instance_config(&inst_id) + .config(inst_id) .and_then(|config| NetworkConfig::new_from_config(&config).ok()) .filter(|config| predicate(config)) { @@ -598,7 +618,7 @@ mod tests { async fn start_web_client_for_test( config_server_addr: std::net::SocketAddr, machine_id: uuid::Uuid, - manager: Arc, + manager: Arc, ) -> WebClient { run_web_client( &format!("udp://{config_server_addr}/{MANAGED_CONFIG_TOKEN}"), @@ -813,7 +833,10 @@ mod tests { #[tokio::test] async fn test_client() { - let listener = UdpTunnelListener::new("udp://0.0.0.0:54333".parse().unwrap()); + let listener = runtime_udp_tunnel_listener( + "udp://127.0.0.1:0".parse().unwrap(), + "127.0.0.1:0".parse().unwrap(), + ); let mut mgr = ClientManager::new( Db::memory_db().await, None, @@ -823,7 +846,7 @@ mod tests { None, None, None, None, None, )), ); - mgr.add_listener(Box::new(listener)).await.unwrap(); + let listener_url = mgr.add_listener(listener).await.unwrap(); mgr.db() .inner() @@ -831,14 +854,14 @@ mod tests { .await .unwrap(); - let connector = UdpTunnelConnector::new("udp://127.0.0.1:54333".parse().unwrap()); + let connector = runtime_udp_tunnel_dialer(listener_url); let _c = WebClient::new( connector, "test", uuid::Uuid::new_v4(), "test", false, - Arc::new(NetworkInstanceManager::new()), + Arc::new(native_instance_manager()), None, ); @@ -892,7 +915,7 @@ mod tests { let machine_id = uuid::Uuid::new_v4(); let instance_id = uuid::Uuid::new_v4(); - let core_manager = Arc::new(NetworkInstanceManager::new()); + let core_manager = Arc::new(native_instance_manager()); let client = start_web_client_for_test(config_server_addr, machine_id, core_manager.clone()).await; @@ -939,7 +962,7 @@ mod tests { assert_updated_runtime_config(&updated, instance_id); assert_eq!( - core_manager.get_instance_network_config_source(&instance_id), + core_manager.config_source(instance_id), Some(easytier::common::config::ConfigSource::Web) ); assert_eq!( @@ -1003,7 +1026,7 @@ mod tests { assert_eq!(redelivered.enable_kcp_proxy, Some(true)); assert_eq!(redelivered.instance_recv_bps_limit, Some(654321)); assert_eq!( - core_manager.get_instance_network_config_source(&instance_id), + core_manager.config_source(instance_id), Some(easytier::common::config::ConfigSource::Web) ); assert_eq!( @@ -1018,7 +1041,7 @@ mod tests { // Reconnect path: a fresh core manager has no local runtime state, so // the new session must replay the managed config persisted in web DB. drop(client); - let reconnected_core_manager = Arc::new(NetworkInstanceManager::new()); + let reconnected_core_manager = Arc::new(native_instance_manager()); let _reconnected_client = start_web_client_for_test( config_server_addr, machine_id, @@ -1055,7 +1078,7 @@ mod tests { ); let config_server_addr = add_random_udp_listener(&mut mgr).await; let machine_id = uuid::Uuid::new_v4(); - let core_manager = Arc::new(NetworkInstanceManager::new()); + let core_manager = Arc::new(native_instance_manager()); let client = start_web_client_for_test(config_server_addr, machine_id, core_manager.clone()).await; diff --git a/easytier-web/src/client_manager/runtime_reconcile.rs b/easytier-web/src/client_manager/runtime_reconcile.rs index 42a063f2..f955721a 100644 --- a/easytier-web/src/client_manager/runtime_reconcile.rs +++ b/easytier-web/src/client_manager/runtime_reconcile.rs @@ -1,7 +1,8 @@ use anyhow::Context as _; use easytier::{ common::config::{ - ConfigLoader, EncryptionAlgorithm, PortForwardConfig as RuntimePortForwardConfig, + ConfigLoader, EncryptionAlgorithm, NetworkConfigExt, + PortForwardConfig as RuntimePortForwardConfig, }, proto::{ acl::Acl, @@ -802,8 +803,7 @@ mod tests { let current = config_with_port_forwards(vec![port_forward(23000, 5174)]); let mut desired = config_with_port_forwards(vec![port_forward(23000, 5174), port_forward(23007, 3389)]); - desired.hostname = - Some(easytier::common::config::TomlConfigLoader::default().get_hostname()); + desired.hostname = Some("desired-host".to_string()); let patch = web_source_runtime_patch(¤t, &desired).expect("build patch"); diff --git a/easytier-web/src/client_manager/session.rs b/easytier-web/src/client_manager/session.rs index 88075fcd..70591fec 100644 --- a/easytier-web/src/client_manager/session.rs +++ b/easytier-web/src/client_manager/session.rs @@ -6,18 +6,16 @@ use std::{ }; use anyhow::Context; -use easytier::{ - proto::{ - api::{ - config::{ConfigRpc, ConfigRpcClientFactory}, - manage::{WebClientService, WebClientServiceClientFactory}, - }, - rpc_impl::bidirect::BidirectRpcManager, - rpc_types::{self, controller::BaseController}, - web::{HeartbeatRequest, HeartbeatResponse, WebServerService, WebServerServiceServer}, +use easytier::proto::{ + api::{ + config::{ConfigRpc, ConfigRpcClientFactory}, + manage::{WebClientService, WebClientServiceClientFactory}, }, - tunnel::Tunnel, + rpc::bidirect::BidirectRpcManager, + rpc_types::{self, controller::BaseController}, + web::{HeartbeatRequest, HeartbeatResponse, WebServerService, WebServerServiceServer}, }; +use easytier_core::tunnel::Tunnel; use tokio::sync::{Notify, RwLock, broadcast}; use tokio_util::task::AbortOnDropHandle; @@ -565,7 +563,7 @@ impl WebServerService for SessionRpcService { _: easytier::proto::web::GetFeatureRequest, ) -> rpc_types::error::Result { Ok(easytier::proto::web::GetFeatureResponse { - support_encryption: easytier::web_client::security::web_secure_tunnel_supported(), + support_encryption: easytier_core::tunnel::web_security::web_secure_tunnel_supported(), }) } } diff --git a/easytier-web/src/client_manager/session/runtime_revision.rs b/easytier-web/src/client_manager/session/runtime_revision.rs index be91fb65..9e0c9008 100644 --- a/easytier-web/src/client_manager/session/runtime_revision.rs +++ b/easytier-web/src/client_manager/session/runtime_revision.rs @@ -1,16 +1,14 @@ use std::collections::{HashMap, HashSet}; -use easytier::{ - proto::{ - api::manage::{ - DeleteNetworkInstanceRequest, ListNetworkInstanceMetaRequest, - ListNetworkInstanceRequest, NetworkConfig, NetworkMeta, RunNetworkInstanceRequest, - }, - rpc_types::controller::BaseController, - web::HeartbeatRequest, +use easytier::proto::{ + api::manage::{ + DeleteNetworkInstanceRequest, ListNetworkInstanceMetaRequest, ListNetworkInstanceRequest, + NetworkConfig, NetworkMeta, RunNetworkInstanceRequest, }, - rpc_service::remote_client::{ListNetworkProps, Storage as _}, + rpc_types::controller::BaseController, + web::HeartbeatRequest, }; +use easytier_core::management::remote_client::{ListNetworkProps, Storage as _}; use tokio::sync::{RwLock, broadcast}; use super::{SessionConfigClient, SessionData, SessionRpcClient, SessionRpcService}; diff --git a/easytier-web/src/db/entity/user_running_network_configs.rs b/easytier-web/src/db/entity/user_running_network_configs.rs index 88a974a7..e1d6e0a0 100644 --- a/easytier-web/src/db/entity/user_running_network_configs.rs +++ b/easytier-web/src/db/entity/user_running_network_configs.rs @@ -1,9 +1,7 @@ //! `SeaORM` Entity, @generated by sea-orm-codegen 1.1.0 -use easytier::{ - common::config::ConfigSource, launcher::NetworkConfig, - rpc_service::remote_client::PersistentConfig, -}; +use easytier::common::config::{ConfigSource, NetworkConfig}; +use easytier_core::management::remote_client::PersistentConfig; use sea_orm::entity::prelude::*; use serde::{Deserialize, Serialize}; diff --git a/easytier-web/src/db/mod.rs b/easytier-web/src/db/mod.rs index a07de22c..01fe054c 100644 --- a/easytier-web/src/db/mod.rs +++ b/easytier-web/src/db/mod.rs @@ -2,11 +2,8 @@ #[allow(unused_imports)] pub mod entity; -use easytier::{ - common::config::ConfigSource, - launcher::NetworkConfig, - rpc_service::remote_client::{ListNetworkProps, Storage}, -}; +use easytier::common::config::{ConfigSource, NetworkConfig}; +use easytier_core::management::remote_client::{ListNetworkProps, Storage}; use entity::user_running_network_configs; use sea_orm::{ ColumnTrait as _, DatabaseConnection, DbErr, EntityTrait, QueryFilter as _, Set, @@ -385,11 +382,8 @@ impl Storage<(UserIdInDb, Uuid), user_running_network_configs::Model, DbErr> for #[cfg(test)] mod tests { - use easytier::{ - common::config::ConfigSource, - proto::api::manage::NetworkConfig, - rpc_service::remote_client::{PersistentConfig, Storage}, - }; + use easytier::{common::config::ConfigSource, proto::api::manage::NetworkConfig}; + use easytier_core::management::remote_client::{PersistentConfig, Storage}; use sea_orm::{ActiveModelTrait, ColumnTrait, EntityTrait, QueryFilter as _, Set}; use crate::db::{Db, ListNetworkProps, entity::user_running_network_configs}; diff --git a/easytier-web/src/main.rs b/easytier-web/src/main.rs index 6bddbf62..697c4d9a 100644 --- a/easytier-web/src/main.rs +++ b/easytier-web/src/main.rs @@ -16,12 +16,12 @@ use easytier::{ log, network::{local_ipv4, local_ipv6}, }, - tunnel::{TunnelListener, tcp::TcpTunnelListener, udp::UdpTunnelListener}, + proto::rpc::standalone::{runtime_rpc_listener, runtime_udp_tunnel_listener}, utils::panic::setup_panic_handler, }; +use easytier_core::{socket::SocketListener, tunnel::Tunnel}; use easytier::tunnel::IpScheme; -use easytier::utils::BoxExt; use mimalloc::MiMalloc; mod client_manager; @@ -230,11 +230,20 @@ impl LoggingConfigLoader for &Cli { } } -pub fn get_listener_by_url(scheme: IpScheme, l: &url::Url) -> Option> { +pub fn get_listener_by_url( + scheme: IpScheme, + l: &url::Url, +) -> Option>>> { Some(match scheme { - IpScheme::Tcp => TcpTunnelListener::new(l.clone()).boxed(), - IpScheme::Udp => UdpTunnelListener::new(l.clone()).boxed(), - IpScheme::Ws => WsTunnelListener::new(l.clone()).boxed(), + IpScheme::Tcp => { + let addr = l.socket_addrs(|| Some(11010)).ok()?.into_iter().next()?; + Box::new(runtime_rpc_listener(addr)) + } + IpScheme::Udp => { + let addr = l.socket_addrs(|| Some(11010)).ok()?.into_iter().next()?; + Box::new(runtime_udp_tunnel_listener(l.clone(), addr)) + } + IpScheme::Ws => Box::new(WsTunnelListener::new(l.clone())), _ => return None, }) } @@ -244,8 +253,8 @@ async fn get_dual_stack_listener( port: u16, ) -> Result< ( - Option>, - Option>, + Option>>>, + Option>>>, ), Error, > { diff --git a/easytier-web/src/restful/mod.rs b/easytier-web/src/restful/mod.rs index cf449d9b..71791405 100644 --- a/easytier-web/src/restful/mod.rs +++ b/easytier-web/src/restful/mod.rs @@ -16,8 +16,7 @@ use axum::{Extension, Json, Router, extract::State, routing::get}; use axum_login::tower_sessions::{ExpiredDeletion, SessionManagerLayer}; use axum_login::{AuthManagerLayerBuilder, AuthUser, login_required}; use axum_messages::MessagesManagerLayer; -use easytier::common::config::{ConfigLoader, TomlConfigLoader}; -use easytier::launcher::NetworkConfig; +use easytier::common::config::{ConfigLoader, NetworkConfig, NetworkConfigExt, TomlConfigLoader}; use easytier::proto::rpc_types; use network::NetworkApi; use sea_orm::DbErr; diff --git a/easytier-web/src/restful/network.rs b/easytier-web/src/restful/network.rs index d437f122..bd7fec0b 100644 --- a/easytier-web/src/restful/network.rs +++ b/easytier-web/src/restful/network.rs @@ -3,11 +3,12 @@ use axum::http::StatusCode; use axum::routing::{delete, post}; use axum::{Json, Router, extract::State, routing::get}; use axum_login::AuthUser; -use easytier::common::config::ConfigSource as RuntimeConfigSource; -use easytier::launcher::NetworkConfig; +use easytier::common::config::{ + ConfigSource as RuntimeConfigSource, NetworkConfig, config_source_from_rpc, +}; use easytier::proto::common::Void; use easytier::proto::{api::manage::*, web::*}; -use easytier::rpc_service::remote_client::{ +use easytier_core::management::remote_client::{ GetNetworkMetasResponse, ListNetworkInstanceIdsJsonResp, RemoteClientError, RemoteClientManager, }; use sea_orm::DbErr; @@ -321,7 +322,7 @@ impl NetworkApi { ) -> Result, HttpHandleError> { let source = payload .source - .and_then(RuntimeConfigSource::from_rpc) + .and_then(config_source_from_rpc) .unwrap_or(RuntimeConfigSource::Web); client_mgr .handle_run_network_instance_with_source( diff --git a/easytier/Cargo.toml b/easytier/Cargo.toml index 837fa033..3bf19ddd 100644 --- a/easytier/Cargo.toml +++ b/easytier/Cargo.toml @@ -19,70 +19,68 @@ build = "build/main.rs" name = "easytier-core" path = "src/easytier-core.rs" test = false +required-features = ["management"] [[bin]] name = "easytier-cli" path = "src/easytier-cli.rs" +required-features = ["management"] [lib] name = "easytier" path = "src/lib.rs" -[[bench]] -name = "tx_throughput" -harness = false - [[bench]] name = "packet_bytes_extraction" harness = false [dependencies] +easytier-core = { path = "../easytier-core", version = "2.6.4", default-features = false } +easytier-proto = { path = "../easytier-proto", version = "2.6.4", default-features = false, features = ["api", "core"] } git-version = "0.3.9" -tracing = { version = "0.1", features = ["log"] } -tracing-subscriber = { version = "0.3", features = [ - "env-filter", - "local-time", - "time", -] } -derivative = "2.2.0" +tracing = "0.1" +log = { version = "0.4", features = ["std"] } +tracing-subscriber = { version = "0.3", default-features = false, features = [ + "registry", +], optional = true } derive_more = { version = "2.1.1", features = ["full"] } console-subscriber = { version = "0.4.1", optional = true } indoc = "2.0.7" -regex = "1.8" paste = "1.0" thiserror = "1.0" auto_impl = "1.1.0" crossbeam = "0.8.4" arc-swap = "1.7" -time = "0.3" toml = "0.8.12" chrono = { version = "0.4.37", features = ["serde"] } guarden = "0.2" quanta = "0.12" -delegate = "0.13.5" - -itertools = "0.14.0" - strum = { version = "0.27.2", features = ["derive"] } gethostname = "0.5.0" futures = { version = "0.3", features = ["bilock", "unstable"] } -tokio = { version = "1", features = ["full"] } -tokio-stream = "0.1" +tokio = { version = "1", default-features = false, features = [ + "fs", + "io-util", + "macros", + "net", + "process", + "rt", + "signal", + "sync", + "time", +] } tokio-util = { version = "0.7.9", features = ["codec", "net", "io", "rt"] } -async-stream = "0.3.5" async-trait = "0.1.74" dashmap = "6.0" moka = { version = "0.12", features = ["future"] } -timedmap = "=1.0.1" - # for full-path zero-copy zerocopy = { version = "0.7.32", features = ["derive", "simd"] } bytes = "1.5.0" @@ -97,7 +95,6 @@ seahash = { version = "4.1.0", optional = true } rustls = { version = "0.23.0", features = [ "ring", "tls12" ], default-features = false, optional = true } -rcgen = { version = "0.12.1", optional = true } # for websocket tokio-websockets = { version = "0.13.2", git = "https://github.com/EasyTier/tokio-websockets", optional = true, features = [ @@ -107,10 +104,11 @@ tokio-websockets = { version = "0.13.2", git = "https://github.com/EasyTier/toki "fastrand", "ring", ] } +forwarded-header-value = { version = "0.1.1", optional = true } http = { version = "1", default-features = false, features = [ "std", ], optional = true } -forwarded-header-value = { version = "0.1.1", optional = true } +rcgen = { version = "0.12.1", optional = true } tokio-rustls = { version = "0.26", default-features = false, optional = true } # for tap device @@ -138,16 +136,10 @@ once_cell = "1.18.0" # for rpc prost = "0.14.3" -prost-reflect = { version = "0.16.4", default-features = false, features = ["derive", "serde"] } -prost-wkt-types = "0.7.1" -pbjson = "0.9.0" - anyhow = "1.0" -ariadne = "0.5" url = { version = "2.5", features = ["serde"] } percent-encoding = "2.3.1" -idna = "1.0" # for tun packet byteorder = "1.5.0" @@ -155,11 +147,8 @@ byteorder = "1.5.0" # for proxy cidr = { version = "0.3.1", features = ["serde"] } socket2 = { version = "0.5.10", features = ["all"] } -prefix-trie = { version = "0.7.0", features = ["cidr"] } # for hole punching -stun_codec = "0.3.4" -bytecodec = "0.4.15" rand = "0.8.5" serde = { version = "1.0", features = ["derive"] } @@ -180,20 +169,11 @@ async-recursion = "1.0.5" network-interface = "2.0" -# for ospf route -petgraph = "0.8.1" -ordered_hash_map = "0.5.0" - # for wireguard boringtun = { package = "boringtun-easytier", version = "0.6.1", optional = true } # for encryption ring = { version = "0.17", optional = true } -bitflags = "2.5" -aes-gcm = { version = "0.10.3", optional = true } -openssl = { version = "0.10", optional = true, features = ["vendored"] } -snow = "0.10.0" -x25519-dalek = { version = "2.0", features = ["static_secrets"] } # for cli tabled = "0.16" @@ -208,45 +188,20 @@ mimalloc = { version = "*", optional = true } # mips atomic-shim = "0.2.0" -smoltcp = { git = "https://github.com/smoltcp-rs/smoltcp.git", rev = "0a926767a68bc88d5512afefa7529c5ecdade4ea", optional = true, default-features = false, features = [ - "std", - "medium-ip", - "proto-ipv4", - "proto-ipv6", - "proto-ipv4-fragmentation", - "fragmentation-buffer-size-8192", - "assembler-max-segment-count-16", - "reassembly-buffer-size-8192", - "reassembly-buffer-count-16", - "socket-tcp", - "socket-udp", - # "socket-tcp-cubic", - "async", -] } parking_lot = { version = "0.12.0" } -wildmatch = "2.3.4" rust-i18n = "3" sys-locale = "0.3" -ringbuf = "0.4.5" -async-ringbuf = "0.3.1" service-manager = { git = "https://github.com/EasyTier/service-manager-rs.git", branch = "main" } -zstd = { version = "0.13", optional = true } - kcp-sys = { git = "https://github.com/EasyTier/kcp-sys", rev = "d7427c22d764deb1860a7d37acc446ed5033464c", optional = true } -# for http connector -http_req = { git = "https://github.com/EasyTier/http_req.git", default-features = false, features = [ - "rust-tls", -] } - # for dns connector -hickory-resolver = "0.25.2" -hickory-proto = "0.25.2" +hickory-resolver = { version = "0.25.2", optional = true } +hickory-proto = { version = "0.25.2", optional = true } # for magic dns hickory-client = { version = "0.25.2", optional = true } @@ -257,30 +212,18 @@ hickory-server = { version = "0.25.2", features = [ bon = "3.9.1" derive_builder = "0.20.2" humantime-serde = "1.1.1" -multimap = "0.10.1" -version-compare = "0.2.0" -hmac = "0.12.1" -sha2 = "0.10.8" shellexpand = "3.1.1" # for fake tcp flume = { version = "0.12", optional = true } -igd-next = { version = "0.17.0", features = ["aio_tokio"] } -natpmp = "0.5.0" +igd-next = { version = "0.17.0", features = ["aio_tokio"], optional = true } +natpmp = { version = "0.5.0", optional = true } [target.'cfg(any(target_os = "linux", target_os = "macos", target_os = "windows", target_os = "freebsd"))'.dependencies] machine-uid = "0.5.3" [target.'cfg(any(target_os = "linux"))'.dependencies] -netlink-sys = "0.8.7" -netlink-packet-route = "0.21.0" -netlink-packet-core = { version = "0.7.0" } -netlink-packet-utils = "0.5.2" -# for magic dns -resolv-conf = "0.7.3" -dbus = { version = "0.9.7", features = ["vendored"] } -which = "7.0.3" - +netlink-sys = { version = "0.8.7", optional = true } [target.'cfg(all(windows, any(target_arch = "x86_64", target_arch = "x86")))'.dependencies] windivert = { git = "https://github.com/EasyTier/windivert-rust.git", rev = "adcc56d1550f7b5377ec2b3429f413ee24a77375", features = [ "static", @@ -326,31 +269,27 @@ jemalloc-sys = { package = "tikv-jemalloc-sys", version = "0.6.0", features = [ [build-dependencies] cfg_aliases = "0.2.1" -indoc = "2.0" globwalk = "0.8.1" regex = "1" -prost-build = "0.14.3" -prost-reflect-build = "0.16.0" -pbjson-build = "0.9.0" -proc-macro2 = "1" -quote = "1" thunk-rs = { git = "https://github.com/easytier/thunk.git", default-features = false, features = [ "win7", ] } -[target.'cfg(windows)'.build-dependencies] -reqwest = { version = "0.12.12", features = ["blocking"] } -zip = "4.0.0" - [dev-dependencies] criterion = "0.5.1" +easytier-core = { path = "../easytier-core", version = "2.6.4", default-features = false, features = [ + "test-utils", +] } serial_test = "3.0.0" rstest = "0.25.0" futures-util = "0.3.31" maplit = "1.0.2" tempfile = "3.22.0" ctor = "0.8.0" +stun_codec = "0.3.4" +bytecodec = "0.4.15" +x25519-dalek = { version = "2.0", features = ["static_secrets"] } [target.'cfg(target_os = "linux")'.dev-dependencies] defguard_wireguard_rs = "0.4.2" @@ -369,12 +308,19 @@ default = [ "faketcp", "magic-dns", "zstd", + "upnp", + "icmp-proxy", + "management", + "endpoint-discovery", + "extended-services", + "linux-netlink", + "tcp-hole-punch", ] full = [ "websocket", "wireguard", "aes-gcm", - "openssl-crypto", # need openssl-dev libs + "openssl-crypto", "smoltcp", "tun", "socks5", @@ -383,23 +329,33 @@ full = [ "faketcp", "magic-dns", "zstd", + "upnp", + "icmp-proxy", + "management", + "endpoint-discovery", + "extended-services", + "tcp-hole-punch", ] -wireguard = ["dep:boringtun", "dep:ring"] -quic = ["dep:quinn", "dep:quinn-proto", "dep:seahash", "dep:rustls", "dep:rcgen"] -kcp = ["dep:kcp-sys"] +wireguard = ["vpn-portal", "dep:boringtun", "dep:ring", "easytier-core/aes-gcm", "easytier-core/chacha20", "easytier-proto/wireguard"] +quic = ["wrapped-transport", "easytier-core/proxy-packet", "dep:quinn", "dep:quinn-proto", "dep:seahash", "dep:rustls", "easytier-proto/quic"] +kcp = ["wrapped-transport", "easytier-core/proxy-packet", "dep:kcp-sys"] mimalloc = ["dep:mimalloc"] -aes-gcm = ["dep:aes-gcm"] -openssl-crypto = ["dep:openssl"] -tun = ["dep:tun"] +aes-gcm = ["easytier-core/aes-gcm"] +openssl-crypto = ["easytier-core/aes-gcm", "easytier-core/chacha20"] +tun = ["dep:tun", "linux-netlink"] +linux-netlink = ["dep:netlink-sys"] +proxy-cidr-monitor = ["easytier-core/proxy-cidr-monitor"] websocket = [ "dep:tokio-websockets", - "dep:http", "dep:forwarded-header-value", + "dep:http", + "dep:rcgen", "dep:tokio-rustls", "dep:rustls", - "dep:rcgen", + "easytier-proto/websocket", ] -smoltcp = ["dep:smoltcp"] +smoltcp = ["easytier-core/proxy-smoltcp-stack"] +icmp-proxy = ["easytier-core/proxy-packet"] socks5 = ["smoltcp"] ffi-dataplane = ["socks5"] jemalloc = ["dep:jemallocator", "dep:jemalloc-sys"] @@ -410,10 +366,45 @@ jemalloc-prof = [ "jemalloc-sys/profiling", "jemalloc-sys/stats", ] -tracing = ["tokio/tracing", "dep:console-subscriber"] -magic-dns = ["dep:hickory-client", "dep:hickory-server"] -faketcp = ["dep:flume"] -zstd = ["dep:zstd"] +tracing = ["tokio/tracing", "dep:console-subscriber", "dep:tracing-subscriber"] +tracing-log = ["tracing/log", "easytier-core/tracing-log"] +logging = [] +dns-resolver = ["dep:hickory-proto", "dep:hickory-resolver"] +magic-dns = [ + "dns-resolver", + "dep:hickory-client", + "dep:hickory-server", + "easytier-core/proxy-packet", + "easytier-proto/magic-dns", +] +faketcp = ["dep:flume", "easytier-proto/faketcp"] +zstd = ["easytier-core/zstd", "easytier-proto/zstd"] +upnp = ["dep:igd-next", "dep:natpmp"] +endpoint-discovery = [ + "dns-resolver", + "easytier-core/endpoint-discovery", +] +dhcp-ipv4 = ["easytier-core/dhcp-ipv4"] +public-ipv6-provider = ["linux-netlink", "easytier-core/public-ipv6-provider"] +vpn-portal = ["easytier-core/vpn-portal"] +wrapped-transport = ["easytier-core/wrapped-transport"] +extended-services = [ + "dhcp-ipv4", + "public-ipv6-provider", + "vpn-portal", + "wrapped-transport", + "proxy-cidr-monitor", +] +management = [ + "logging", + "management-rpc", + "easytier-core/management", + "easytier-proto/json-rpc", + "easytier-proto/utils", + "tokio/full", +] +management-rpc = ["easytier-core/management-rpc"] +tcp-hole-punch = ["easytier-core/tcp-hole-punch"] # Deprecated: hotpath profiling has been removed. These feature aliases are # retained as no-ops so existing build scripts using `--features hotpath*` # continue to work without pulling in any dependencies. diff --git a/easytier/benches/README.md b/easytier/benches/README.md index a3aadbfd..34b4dee1 100644 --- a/easytier/benches/README.md +++ b/easytier/benches/README.md @@ -2,10 +2,9 @@ Criterion benchmarks for EasyTier hot paths. -| Bench | What it measures | -| --------------------------- | -------------------------------------------------------------------------------- | -| `tx_throughput` | End-to-end TX injection path through `peer_manager::send_msg_by_ip` | -| `packet_bytes_extraction` | `ZCPacket::payload_bytes` / `tunnel_payload_bytes` extraction (advance hot path) | +| Bench | What it measures | +| ------------------------- | ------------------------------------------------------------------------------- | +| `packet_bytes_extraction` | `ZCPacket::payload_bytes` / `tunnel_payload_bytes` extraction (advance hot path) | ## Packet Bytes Extraction @@ -38,125 +37,3 @@ cargo bench --bench packet_bytes_extraction -- --quiet | `PACKET_BYTES_MEASUREMENT_SECS` | `10` | Criterion `measurement_time` | | `PACKET_BYTES_WARMUP_SECS` | `3` | Criterion `warm_up_time` | | `PACKET_BYTES_SAMPLE_SIZE` | `10` | Criterion `sample_size` (min 10) | - ---- - -## TX Throughput Benchmark - -Criterion benchmark for EasyTier's TX injection path (`peer_manager::send_msg_by_ip`). - -## What it measures - -The benchmark sets up two EasyTier instances (`hot-a` / `hot-b`) and drives -packets from `hot-a` to `hot-b` via `peer_manager.send_msg_by_ip`. This is the -same entry point `easytier-core` uses for daily forwarded traffic, so the -numbers reflect the real TX hot path: NIC pipeline → route lookup → -compress/encrypt → peer connection → tunnel send. - -Two variants are reported per tunnel kind: - -| Bench | What it measures | -| --------------------------------- | ------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------- | -| `tx_throughput/` | Serial baseline. One send in flight at a time. Reports per-packet CPU cost (TX injection latency). | -| `tx_throughput/-saturate` | Spawns `TX_THROUGHPUT_INFLIGHT` tokio tasks that independently pump `send_msg_by_ip`. Reports the aggregate throughput ceiling the peer manager + tunnel can sustain across worker threads. | - -> **Out of scope (by design):** TUN read/write (`no_tun = true`), compression -> (default `None`), reverse/RX-side measurement, multi-peer fanout. Add -> separate benchmarks if you need those. - -## Quick start - -### ring tunnel (no root, fastest) - -```bash -cargo bench --bench tx_throughput -``` - -Smoke run (faster iteration): - -```bash -TX_THROUGHPUT_MEASUREMENT_SECS=2 \ -TX_THROUGHPUT_WARMUP_SECS=1 \ -TX_THROUGHPUT_SAMPLE_SIZE=10 \ -cargo bench --bench tx_throughput -- --quiet -``` - -### tcp / udp tunnels (requires Docker + root) - -The benchmark creates a Docker network and registers each container's netns -under `/var/run/netns`, which requires root. Run the whole command under -`sudo`: - -```bash -sudo TX_THROUGHPUT_TUNNEL=tcp \ - TX_THROUGHPUT_MEASUREMENT_SECS=5 \ - TX_THROUGHPUT_WARMUP_SECS=2 \ - TX_THROUGHPUT_INFLIGHT=64 \ - cargo bench --bench tx_throughput -- --quiet - -sudo TX_THROUGHPUT_TUNNEL=udp cargo bench --bench tx_throughput -- --quiet -``` - -> If `sudo` cannot find `cargo`, use `sudo -E` or the absolute path -> (`$(which cargo)`). - -## Environment variables - -| Variable | Default | Notes | -| -------------------------------- | --------------------- | -------------------------------------- | -| `TX_THROUGHPUT_TUNNEL` | `ring` | `ring` / `tcp` / `udp` | -| `TX_THROUGHPUT_PKT_SIZE` | `1400` | IP total length in bytes | -| `TX_THROUGHPUT_WORKER_THREADS` | `4` | tokio worker threads | -| `TX_THROUGHPUT_INFLIGHT` | `64` | saturate-mode concurrency (task count) | -| `TX_THROUGHPUT_TUNNEL_PORT` | `35521` | tcp/udp listen port | -| `TX_THROUGHPUT_MEASUREMENT_SECS` | `10` | Criterion `measurement_time` | -| `TX_THROUGHPUT_WARMUP_SECS` | `3` | Criterion `warm_up_time` | -| `TX_THROUGHPUT_SAMPLE_SIZE` | `10` | Criterion `sample_size` (min 10) | -| `TX_THROUGHPUT_DOCKER_IMAGE` | `busybox:latest` | tcp/udp only | -| `TX_THROUGHPUT_DOCKER_NET` | `easytier-bench-` | auto-generated unique name | -| `TX_THROUGHPUT_DOCKER_SUBNET` | `172.31.250.0/24` | | -| `TX_THROUGHPUT_DOCKER_IP_A` | `172.31.250.2` | | -| `TX_THROUGHPUT_DOCKER_IP_B` | `172.31.250.3` | | - -## Parameter sweeps - -```bash -# Packet size -for sz in 64 256 1400 9000; do - TX_THROUGHPUT_PKT_SIZE=$sz cargo bench --bench tx_throughput -- --quick -done - -# Inflight depth (self-check: depth=1 should match serial baseline) -for d in 1 4 16 64 256; do - TX_THROUGHPUT_INFLIGHT=$d cargo bench --bench tx_throughput -- --quick -done - -# Worker threads -for w in 1 2 4 8; do - TX_THROUGHPUT_WORKER_THREADS=$w cargo bench --bench tx_throughput -- --quick -done -``` - -## Interpreting results - -- **``** reports per-packet latency. Lower is better. Throughput - column here is "what one in-flight sender sustains". -- **`-saturate`** reports aggregate throughput across - `TX_THROUGHPUT_INFLIGHT` concurrent senders. If this matches the serial - baseline, the TX path is bottlenecked on an internal serialization point - (lock, single-threaded queue, etc.) rather than CPU or link bandwidth. - -### Known finding (ring, single peer) - -On the ring tunnel with a single destination peer, saturate does **not** beat -serial (observed ~277 MiB/s saturate vs ~288 MiB/s serial on a 4-worker -runtime). This points to a serialization point inside the peer-connection TX -path. Tunnels with real I/O await points (tcp/udp via Docker) are expected to -show a saturate > serial gap; verify with the sudo commands above. - -## Output artifacts - -Criterion writes HTML reports + SVG plots under -`easytier/target/criterion/`. Open `tx_throughput//report/index.html` -or `.../-saturate/report/index.html` in a browser to inspect -distributions and regressions across runs. diff --git a/easytier/benches/packet_bytes_extraction.rs b/easytier/benches/packet_bytes_extraction.rs index 1ff2ecbd..a7c5bc31 100644 --- a/easytier/benches/packet_bytes_extraction.rs +++ b/easytier/benches/packet_bytes_extraction.rs @@ -3,7 +3,7 @@ use std::time::Duration; use criterion::{BatchSize, Criterion, Throughput, criterion_group, criterion_main}; -use easytier::tunnel::packet_def::ZCPacket; +use easytier_core::packet::ZCPacket; const PAYLOAD_SIZES: &[usize] = &[1280, 4096]; diff --git a/easytier/benches/tx_throughput.rs b/easytier/benches/tx_throughput.rs deleted file mode 100644 index 04815b34..00000000 --- a/easytier/benches/tx_throughput.rs +++ /dev/null @@ -1,472 +0,0 @@ -use std::{ - net::IpAddr, - path::PathBuf, - process::{Command, Stdio}, - str::FromStr, - sync::Arc, - sync::atomic::{AtomicU64, Ordering}, - time::{Duration, Instant, SystemTime, UNIX_EPOCH}, -}; - -use bytes::BytesMut; -use criterion::{Criterion, Throughput, criterion_group, criterion_main}; - -use easytier::{ - common::config::{ConfigLoader, TomlConfigLoader}, - instance::instance::Instance, - tunnel::{ - packet_def::ZCPacket, ring::RingTunnelConnector, tcp::TcpTunnelConnector, - udp::UdpTunnelConnector, - }, -}; - -const VIRTUAL_IP_A: &str = "10.144.144.1"; -const VIRTUAL_IP_B: &str = "10.144.144.2"; -const DEFAULT_DOCKER_SUBNET: &str = "172.31.250.0/24"; -const DEFAULT_DOCKER_IP_A: &str = "172.31.250.2"; -const DEFAULT_DOCKER_IP_B: &str = "172.31.250.3"; -const DEFAULT_TUNNEL_PORT: u16 = 35521; - -#[derive(Clone, Copy, Debug)] -enum TunnelKind { - Ring, - Tcp, - Udp, -} - -impl TunnelKind { - fn as_str(self) -> &'static str { - match self { - TunnelKind::Ring => "ring", - TunnelKind::Tcp => "tcp", - TunnelKind::Udp => "udp", - } - } -} - -impl FromStr for TunnelKind { - type Err = String; - - fn from_str(value: &str) -> Result { - match value { - "ring" => Ok(TunnelKind::Ring), - "tcp" => Ok(TunnelKind::Tcp), - "udp" => Ok(TunnelKind::Udp), - other => Err(format!( - "unsupported TX_THROUGHPUT_TUNNEL={other:?}; expected ring, tcp, or udp" - )), - } - } -} - -struct BenchTopology { - _docker: Option, - inst_a: Instance, - _inst_b: Instance, - dst: IpAddr, - packet: ZCPacket, -} - -struct DockerNetns { - network: String, - container_a: String, - container_b: String, - netns_a: String, - netns_b: String, - ip_a: String, - netns_a_path: PathBuf, - netns_b_path: PathBuf, -} - -impl DockerNetns { - fn create() -> Self { - let id = unique_id(); - let image = env_string("TX_THROUGHPUT_DOCKER_IMAGE", "busybox:latest"); - let network = env_string("TX_THROUGHPUT_DOCKER_NET", &format!("easytier-bench-{id}")); - let subnet = env_string("TX_THROUGHPUT_DOCKER_SUBNET", DEFAULT_DOCKER_SUBNET); - let ip_a = env_string("TX_THROUGHPUT_DOCKER_IP_A", DEFAULT_DOCKER_IP_A); - let ip_b = env_string("TX_THROUGHPUT_DOCKER_IP_B", DEFAULT_DOCKER_IP_B); - let container_a = format!("easytier-bench-a-{id}"); - let container_b = format!("easytier-bench-b-{id}"); - let netns_a = format!("easytier-bench-a-{id}"); - let netns_b = format!("easytier-bench-b-{id}"); - - docker(&[ - "network", "create", "--driver", "bridge", "--subnet", &subnet, &network, - ]); - - let mut docker_netns = Self { - network, - container_a, - container_b, - netns_a, - netns_b, - ip_a: ip_a.clone(), - netns_a_path: PathBuf::new(), - netns_b_path: PathBuf::new(), - }; - - docker_netns.start_container(&docker_netns.container_a, &ip_a, &image); - docker_netns.start_container(&docker_netns.container_b, &ip_b, &image); - - let pid_a = docker(&["inspect", "-f", "{{.State.Pid}}", &docker_netns.container_a]); - let pid_b = docker(&["inspect", "-f", "{{.State.Pid}}", &docker_netns.container_b]); - - docker_netns.netns_a_path = register_netns(&docker_netns.netns_a, &pid_a); - docker_netns.netns_b_path = register_netns(&docker_netns.netns_b, &pid_b); - docker_netns - } - - fn start_container(&self, name: &str, ip: &str, image: &str) { - docker(&[ - "run", - "-d", - "--name", - name, - "--network", - &self.network, - "--ip", - ip, - image, - "sleep", - "3600", - ]); - } -} - -impl Drop for DockerNetns { - fn drop(&mut self) { - let _ = std::fs::remove_file(&self.netns_a_path); - let _ = std::fs::remove_file(&self.netns_b_path); - docker_ignore(&["rm", "-f", &self.container_a, &self.container_b]); - docker_ignore(&["network", "rm", &self.network]); - } -} - -fn bench_tx_throughput(c: &mut Criterion) { - let tunnel = env_string("TX_THROUGHPUT_TUNNEL", "ring") - .parse::() - .unwrap_or_else(|err| panic!("{err}")); - let packet_size = env_parse("TX_THROUGHPUT_PKT_SIZE", 1400usize); - const MIN_PKT_SIZE: usize = 28; // IPv4 (20) + UDP (8) header - assert!( - packet_size >= MIN_PKT_SIZE, - "TX_THROUGHPUT_PKT_SIZE={packet_size} is smaller than the minimum {MIN_PKT_SIZE} (IPv4+UDP headers)" - ); - let worker_threads = env_parse("TX_THROUGHPUT_WORKER_THREADS", 4usize); - let inflight_depth = env_parse("TX_THROUGHPUT_INFLIGHT", 64usize).max(1); - let runtime = tokio::runtime::Builder::new_multi_thread() - .worker_threads(worker_threads) - .enable_all() - .build() - .expect("create tokio runtime"); - - let topology = runtime.block_on(setup_topology(tunnel, packet_size)); - let peer_manager = topology.inst_a.get_peer_manager(); - let packet = topology.packet.clone(); - let dst = topology.dst; - - eprintln!( - "tx_throughput: tunnel={} inflight={} workers={} pkt_size={}", - tunnel.as_str(), - inflight_depth.max(1), - worker_threads, - packet_size - ); - - let mut group = c.benchmark_group("tx_throughput"); - group.throughput(Throughput::Bytes(packet_size as u64)); - - // Serial baseline: one packet in flight at a time. - // Measures per-packet CPU cost (TX injection latency). - group.bench_function(tunnel.as_str(), |b| { - b.iter_custom(|iterations| { - let pm = peer_manager.clone(); - let pkt = packet.clone(); - runtime.block_on(async move { - let start = Instant::now(); - for _ in 0..iterations { - pm.send_msg_by_ip(pkt.clone(), dst, false) - .await - .expect("send packet by EasyTier IP"); - } - start.elapsed() - }) - }); - }); - - // Saturate: spawn TX_THROUGHPUT_INFLIGHT worker tasks, each independently - // pumping send_msg_by_ip. Work is distributed across tokio worker threads, - // exposing the peer manager + tunnel's true aggregate throughput ceiling. - // With TX_THROUGHPUT_INFLIGHT=1 it degrades to the serial baseline. - group.bench_function(format!("{}-saturate", tunnel.as_str()), |b| { - b.iter_custom(|iterations| { - let pm = peer_manager.clone(); - let pkt = packet.clone(); - let concurrency = inflight_depth.min(iterations as usize).max(1); - runtime.block_on(async move { - let counter = Arc::new(AtomicU64::new(iterations)); - let start = Instant::now(); - let mut handles = Vec::with_capacity(concurrency); - for _ in 0..concurrency { - let pm = pm.clone(); - let pkt = pkt.clone(); - let counter = counter.clone(); - handles.push(tokio::spawn(async move { - loop { - if counter - .fetch_update(Ordering::AcqRel, Ordering::Acquire, |cur| { - if cur > 0 { Some(cur - 1) } else { None } - }) - .is_err() - { - return; - } - pm.send_msg_by_ip(pkt.clone(), dst, false) - .await - .expect("send packet by EasyTier IP"); - } - })); - } - for h in handles { - h.await.expect("saturate worker task panicked"); - } - start.elapsed() - }) - }); - }); - - group.finish(); - - runtime.block_on(async move { - drop(topology); - }); -} - -async fn setup_topology(tunnel: TunnelKind, packet_size: usize) -> BenchTopology { - let tunnel_port = env_parse("TX_THROUGHPUT_TUNNEL_PORT", DEFAULT_TUNNEL_PORT); - let docker = match tunnel { - TunnelKind::Ring => None, - TunnelKind::Tcp | TunnelKind::Udp => Some(DockerNetns::create()), - }; - - let (netns_a, netns_b) = match &docker { - Some(docker) => (Some(docker.netns_a.clone()), Some(docker.netns_b.clone())), - None => (None, None), - }; - let listeners_a = match tunnel { - TunnelKind::Ring => Vec::new(), - TunnelKind::Tcp | TunnelKind::Udp => vec![ - format!("{}://0.0.0.0:{}", tunnel.as_str(), tunnel_port) - .parse() - .unwrap(), - ], - }; - - let mut inst_a = Instance::new(no_tun_config("hot-a", VIRTUAL_IP_A, netns_a, listeners_a)); - let mut inst_b = Instance::new(no_tun_config("hot-b", VIRTUAL_IP_B, netns_b, Vec::new())); - - inst_a.run().await.expect("inst_a run"); - inst_b.run().await.expect("inst_b run"); - - match tunnel { - TunnelKind::Ring => inst_b - .get_conn_manager() - .add_connector(RingTunnelConnector::new( - format!("ring://{}", inst_a.id()).parse().unwrap(), - )), - TunnelKind::Tcp => inst_b - .get_conn_manager() - .add_connector(TcpTunnelConnector::new( - format!( - "tcp://{}:{}", - docker.as_ref().expect("tcp benchmark needs Docker").ip_a, - tunnel_port - ) - .parse() - .unwrap(), - )), - TunnelKind::Udp => inst_b - .get_conn_manager() - .add_connector(UdpTunnelConnector::new( - format!( - "udp://{}:{}", - docker.as_ref().expect("udp benchmark needs Docker").ip_a, - tunnel_port - ) - .parse() - .unwrap(), - )), - } - - wait_for_routes(&inst_a, &inst_b).await; - - BenchTopology { - _docker: docker, - inst_a, - _inst_b: inst_b, - dst: VIRTUAL_IP_B.parse().unwrap(), - packet: make_data_packet(VIRTUAL_IP_A, VIRTUAL_IP_B, packet_size), - } -} - -async fn wait_for_routes(inst_a: &Instance, inst_b: &Instance) { - tokio::time::timeout(Duration::from_secs(15), async { - loop { - let routes_a = inst_a.get_peer_manager().list_routes().await; - let routes_b = inst_b.get_peer_manager().list_routes().await; - if !routes_a.is_empty() && !routes_b.is_empty() { - return; - } - tokio::time::sleep(Duration::from_millis(500)).await; - } - }) - .await - .expect("EasyTier routes did not converge within 15s"); -} - -fn make_data_packet(src: &str, dst: &str, total_size: usize) -> ZCPacket { - use std::net::Ipv4Addr; - - let hdr_len = 28; - let payload_len = total_size.saturating_sub(hdr_len); - let ip_total_len = (hdr_len + payload_len) as u16; - let mut buf = BytesMut::with_capacity(total_size); - - buf.extend_from_slice(&[ - 0x45, - 0x00, - (ip_total_len >> 8) as u8, - (ip_total_len & 0xff) as u8, - 0x00, - 0x00, - 0x40, - 0x00, - 0x40, - 0x11, - 0x00, - 0x00, - ]); - let src: Ipv4Addr = src.parse().unwrap(); - buf.extend_from_slice(&src.octets()); - let dst: Ipv4Addr = dst.parse().unwrap(); - buf.extend_from_slice(&dst.octets()); - - let udp_len = (8 + payload_len) as u16; - buf.extend_from_slice(&[ - 0x30, - 0x39, - 0xd4, - 0x31, - (udp_len >> 8) as u8, - (udp_len & 0xff) as u8, - 0x00, - 0x00, - ]); - - buf.resize(total_size, 0xaa); - ZCPacket::new_with_payload(&buf) -} - -fn no_tun_config( - name: &str, - ipv4: &str, - netns: Option, - listeners: Vec, -) -> TomlConfigLoader { - let config = TomlConfigLoader::default(); - config.set_inst_name(name.to_owned()); - config.set_netns(netns); - config.set_ipv4(Some(ipv4.parse().unwrap())); - config.set_listeners(listeners); - let mut flags = config.get_flags(); - flags.no_tun = true; - config.set_flags(flags); - config -} - -fn register_netns(name: &str, pid: &str) -> PathBuf { - #[cfg(target_os = "linux")] - { - let dir = PathBuf::from("/var/run/netns"); - std::fs::create_dir_all(&dir).expect("create /var/run/netns"); - let path = dir.join(name); - let _ = std::fs::remove_file(&path); - std::os::unix::fs::symlink(format!("/proc/{pid}/ns/net"), &path) - .expect("link Docker netns into /var/run/netns"); - path - } - - #[cfg(not(target_os = "linux"))] - { - let _ = (name, pid); - panic!("Docker netns benchmark requires Linux"); - } -} - -fn docker(args: &[&str]) -> String { - let output = Command::new("docker") - .args(args) - .output() - .unwrap_or_else(|err| panic!("failed to run docker {args:?}: {err}")); - if !output.status.success() { - panic!( - "docker {:?} failed with status {:?}: {}", - args, - output.status.code(), - String::from_utf8_lossy(&output.stderr) - ); - } - String::from_utf8_lossy(&output.stdout).trim().to_owned() -} - -fn docker_ignore(args: &[&str]) { - let _ = Command::new("docker") - .args(args) - .stdout(Stdio::null()) - .stderr(Stdio::null()) - .status(); -} - -fn env_string(name: &str, default: &str) -> String { - std::env::var(name).unwrap_or_else(|_| default.to_owned()) -} - -fn env_parse(name: &str, default: T) -> T -where - T: FromStr, - T::Err: std::fmt::Display, -{ - match std::env::var(name) { - Ok(value) => value - .parse() - .unwrap_or_else(|err| panic!("invalid {name}={value:?}: {err}")), - Err(_) => default, - } -} - -fn criterion_config() -> Criterion { - let measurement_secs = env_parse("TX_THROUGHPUT_MEASUREMENT_SECS", 10u64); - let warmup_secs = env_parse("TX_THROUGHPUT_WARMUP_SECS", 3u64); - let sample_size = env_parse("TX_THROUGHPUT_SAMPLE_SIZE", 10usize).max(10); - - Criterion::default() - .measurement_time(Duration::from_secs(measurement_secs)) - .warm_up_time(Duration::from_secs(warmup_secs)) - .sample_size(sample_size) -} - -fn unique_id() -> String { - let nanos = SystemTime::now() - .duration_since(UNIX_EPOCH) - .expect("system clock before UNIX epoch") - .as_nanos(); - format!("{}-{nanos}", std::process::id()) -} - -criterion_group! { - name = benches; - config = criterion_config(); - targets = bench_tx_throughput -} -criterion_main!(benches); diff --git a/easytier/build/main.rs b/easytier/build/main.rs index 49cbdbae..7cff6847 100644 --- a/easytier/build/main.rs +++ b/easytier/build/main.rs @@ -1,75 +1,11 @@ -mod rpc; - -use crate::rpc::ServiceGenerator; use cfg_aliases::cfg_aliases; -#[cfg(target_os = "windows")] -use std::io::Cursor; -use std::{env, path::PathBuf}; +use std::env; #[cfg(target_os = "windows")] struct WindowsBuild {} #[cfg(target_os = "windows")] impl WindowsBuild { - fn check_protoc_exist() -> Option { - let path = env::var_os("PROTOC").map(PathBuf::from); - if path.is_some() && path.as_ref().unwrap().exists() { - return path; - } - - let path = env::var_os("PATH").unwrap_or_default(); - for p in env::split_paths(&path) { - let p = p.join("protoc.exe"); - if p.exists() && p.is_file() { - return Some(p); - } - } - - None - } - - fn get_cargo_target_dir() -> Result> { - let out_dir = std::path::PathBuf::from(std::env::var("OUT_DIR")?); - let profile = std::env::var("PROFILE")?; - let mut target_dir = None; - let mut sub_path = out_dir.as_path(); - while let Some(parent) = sub_path.parent() { - if parent.ends_with(&profile) { - target_dir = Some(parent); - break; - } - sub_path = parent; - } - let target_dir = target_dir.ok_or("not found")?; - Ok(target_dir.to_path_buf()) - } - - fn download_protoc() -> PathBuf { - println!("cargo:info=use exist protoc: {:?}", "k"); - let out_dir = Self::get_cargo_target_dir().unwrap().join("protobuf"); - let fname = out_dir.join("bin/protoc.exe"); - if fname.exists() { - println!("cargo:info=use exist protoc: {:?}", fname); - return fname; - } - - println!("cargo:info=need download protoc, please wait..."); - - let url = "https://github.com/protocolbuffers/protobuf/releases/download/v26.0-rc1/protoc-26.0-rc-1-win64.zip"; - let response = reqwest::blocking::get(url).unwrap(); - println!("{:?}", response); - let mut content = response - .bytes() - .map(|v| v.to_vec()) - .map(Cursor::new) - .map(zip::ZipArchive::new) - .unwrap() - .unwrap(); - content.extract(out_dir).unwrap(); - - fname - } - pub fn check_for_win() { // add third_party dir to link search path let target = std::env::var("TARGET").unwrap_or_default(); @@ -81,16 +17,6 @@ impl WindowsBuild { } else if target.contains("aarch64") { println!("cargo:rustc-link-search=native=easytier/third_party/arm64/"); } - - let protoc_path = if let Some(o) = Self::check_protoc_exist() { - println!("cargo:info=use os exist protoc: {:?}", o); - o - } else { - Self::download_protoc() - }; - unsafe { - std::env::set_var("PROTOC", protoc_path); - } } } @@ -155,50 +81,6 @@ fn main() -> Result<(), Box> { #[cfg(target_os = "windows")] WindowsBuild::check_for_win(); - let proto_files_reflect = ["src/proto/peer_rpc.proto", "src/proto/common.proto"]; - - let proto_files = [ - "src/proto/error.proto", - "src/proto/tests.proto", - "src/proto/api_instance.proto", - "src/proto/api_logger.proto", - "src/proto/api_config.proto", - "src/proto/api_manage.proto", - "src/proto/web.proto", - "src/proto/magic_dns.proto", - "src/proto/acl.proto", - ]; - - for proto_file in proto_files.iter().chain(proto_files_reflect.iter()) { - println!("cargo:rerun-if-changed={proto_file}"); - } - - let out = PathBuf::from(env::var("OUT_DIR")?); - let descriptor = out.join("descriptors.bin"); - - let mut config = prost_build::Config::new(); - config - .extern_path(".google.protobuf.Any", "::prost_wkt_types::Any") - .extern_path(".google.protobuf.Timestamp", "::prost_wkt_types::Timestamp") - .extern_path(".google.protobuf.Value", "::prost_wkt_types::Value") - .file_descriptor_set_path(&descriptor) - .service_generator(Box::new(ServiceGenerator::default())) - .btree_map(["."]) - .skip_debug([".common.Ipv4Addr", ".common.Ipv6Addr", ".common.UUID"]); - - config.compile_protos(&proto_files, &["src/proto/"])?; - - prost_reflect_build::Builder::new() - .file_descriptor_set_bytes("crate::proto::DESCRIPTOR_POOL_BYTES") - .compile_protos_with_config(config, &proto_files_reflect, &["src/proto/"])?; - - let descriptor = std::fs::read(descriptor)?; - pbjson_build::Builder::new() - .register_descriptors(&descriptor)? - .preserve_proto_field_names() - .btree_map(["."]) - .build(&["."])?; - check_locale(); Ok(()) } diff --git a/easytier/src/common/config.rs b/easytier/src/common/config.rs index 65005be4..7c45c631 100644 --- a/easytier/src/common/config.rs +++ b/easytier/src/common/config.rs @@ -1,1274 +1,65 @@ -use std::{ - hash::Hasher, - net::{IpAddr, SocketAddr}, - path::PathBuf, - sync::{Arc, Mutex}, -}; +//! Native Adapters around the core-owned TOML configuration model. -use anyhow::Context; -use ariadne::{CharSet, Config as AriadneConfig, IndexType, Label, Report, ReportKind, Source}; -use base64::{Engine as _, prelude::BASE64_STANDARD}; -use clap::ValueEnum; -use clap::builder::PossibleValue; -use prost_reflect::{DynamicMessage, ReflectMessage, SerializeOptions}; -use serde::{Deserialize, Serialize}; -use strum::{Display, EnumString, VariantArray}; +use std::path::PathBuf; + +use anyhow::Context as _; +use strum::VariantArray as _; +#[cfg(feature = "management")] use tokio::io::AsyncReadExt as _; -use crate::{ - common::stun::StunInfoCollector, - instance::dns_server::DEFAULT_ET_DNS_ZONE, - proto::{ - acl::Acl, - api::manage::ConfigSource as RpcConfigSource, - common::{CompressionAlgoPb, PortForwardConfigPb, SecureModeConfig, SocketType}, - }, - tunnel::{IpScheme, TunnelScheme, generate_digest_from_str}, +use easytier_core::config::MappedListenerPolicy; +#[cfg(feature = "management")] +pub use easytier_core::config::api_input::{ + NetworkConfig, NetworkConfigExt, NetworkingMethod, add_proxy_network_to_config, }; +pub use easytier_core::config::toml::*; -use super::env_parser; +#[cfg(feature = "management")] +use crate::common::env_parser; +use crate::tunnel::IpScheme; -pub type Flags = crate::proto::common::FlagsInConfig; - -pub fn gen_default_flags() -> Flags { - #[allow(deprecated)] - Flags { - default_protocol: "tcp".to_string(), - dev_name: "".to_string(), - enable_encryption: true, - enable_ipv6: true, - mtu: 1380, - latency_first: false, - enable_exit_node: false, - proxy_forward_by_system: false, - no_tun: false, - use_smoltcp: false, - relay_network_whitelist: "*".to_string(), - disable_p2p: false, - p2p_only: false, - lazy_p2p: false, - relay_all_peer_rpc: false, - disable_tcp_hole_punching: false, - disable_udp_hole_punching: false, - multi_thread: true, - data_compress_algo: CompressionAlgoPb::None.into(), - bind_device: true, - enable_kcp_proxy: false, - disable_kcp_input: false, - disable_relay_kcp: false, - enable_relay_foreign_network_kcp: false, - accept_dns: false, - private_mode: false, - enable_quic_proxy: false, - disable_quic_input: false, - disable_relay_quic: false, - enable_relay_foreign_network_quic: false, - foreign_relay_bps_limit: u64::MAX, - multi_thread_count: 2, - encryption_algorithm: EncryptionAlgorithm::default().to_string(), - disable_sym_hole_punching: false, - tld_dns_zone: DEFAULT_ET_DNS_ZONE.to_string(), - - quic_listen_port: u32::MAX, - need_p2p: false, - instance_recv_bps_limit: u64::MAX, - disable_upnp: false, - disable_relay_data: false, - enable_udp_broadcast_relay: false, - socket_mark: None, - } -} - -fn flags_to_dynamic_message(flags: &Flags) -> DynamicMessage { - let mut message = DynamicMessage::new(flags.descriptor()); - message - .transcode_from(flags) - .expect("FlagsInConfig should transcode to DynamicMessage"); - message -} - -fn flags_to_full_json_map(flags: &DynamicMessage) -> serde_json::Map { - let options = SerializeOptions::new() - .use_proto_field_name(true) - .skip_default_fields(false); - - match flags - .serialize_with_options(serde_json::value::Serializer, &options) - .expect("FlagsInConfig should serialize to JSON") - { - serde_json::Value::Object(map) => map, - _ => unreachable!("FlagsInConfig should serialize to a JSON object"), - } -} - -fn flags_diff_from_default(flags: &Flags) -> serde_json::Map { - let default_flags = gen_default_flags(); - let default_message = flags_to_dynamic_message(&default_flags); - let current_message = flags_to_dynamic_message(flags); - let default_map = flags_to_full_json_map(&default_message); - let current_map = flags_to_full_json_map(¤t_message); - - current_message - .descriptor() - .fields() - .filter_map(|field| { - let key = field.name(); - let value_changed = default_map.get(key) != current_map.get(key); - let presence_changed = - default_message.has_field(&field) != current_message.has_field(&field); - if value_changed || presence_changed { - current_map - .get(key) - .map(|value| (key.to_string(), value.clone())) - } else { - None - } - }) - .collect() -} - -fn mapped_listener_allows_implicit_port(url: &url::Url) -> bool { - TunnelScheme::try_from(url) - .ok() - .and_then(|scheme| IpScheme::try_from(scheme).ok()) - .is_some() -} - -pub fn validate_mapped_listener_url(url: &url::Url) -> Result<(), anyhow::Error> { - if url.port().is_none() && !mapped_listener_allows_implicit_port(url) { - anyhow::bail!("mapped listener port is missing: {}", url); - } - - Ok(()) -} +#[cfg(feature = "management")] +pub use easytier_core::management::{config_source_from_rpc, config_source_to_rpc}; pub fn parse_mapped_listener_urls( mapped_listeners: &[String], ) -> Result, anyhow::Error> { - mapped_listeners - .iter() - .map(|s| { - let url: url::Url = s - .parse() - .with_context(|| format!("mapped listener is not a valid url: {}", s))?; - validate_mapped_listener_url(&url)?; - Ok(url) - }) - .collect() + MappedListenerPolicy::new(IpScheme::VARIANTS.iter().map(ToString::to_string)) + .parse_urls(mapped_listeners) } -#[derive(Debug, Clone, PartialEq, Eq, Display, EnumString, VariantArray)] -#[strum(ascii_case_insensitive)] -pub enum EncryptionAlgorithm { - #[strum(serialize = "xor")] - Xor, - - #[cfg(any(feature = "aes-gcm", feature = "wireguard", feature = "openssl-crypto"))] - #[strum(serialize = "aes-gcm")] - AesGcm, - #[cfg(any(feature = "aes-gcm", feature = "wireguard", feature = "openssl-crypto"))] - #[strum(serialize = "aes-256-gcm")] - Aes256Gcm, - #[cfg(any(feature = "wireguard", feature = "openssl-crypto"))] - #[strum(serialize = "chacha20")] - ChaCha20, +pub fn parse_encryption_algorithm(value: &str) -> Result { + value + .parse() + .map_err(|_| format!("'{value}' is not a valid encryption algorithm")) } -impl ValueEnum for EncryptionAlgorithm { - fn value_variants<'a>() -> &'a [Self] { - Self::VARIANTS - } - - fn from_str(input: &str, _ignore_case: bool) -> Result { - input - .parse() - .map_err(|_| format!("'{}' is not a valid encryption algorithm", input)) - } - - fn to_possible_value(&self) -> Option { - Some(PossibleValue::new(self.to_string())) - } +pub fn load_toml_config_from_path(path: &PathBuf) -> Result { + let config = std::fs::read_to_string(path) + .with_context(|| format!("failed to read config file: {}", path.display()))?; + TomlConfigLoader::new_from_str_with_source(&path.display().to_string(), &config) } -#[allow(clippy::derivable_impls)] -impl Default for EncryptionAlgorithm { - fn default() -> Self { - cfg_select! { - any(feature = "aes-gcm", feature = "wireguard", feature = "openssl-crypto") => EncryptionAlgorithm::AesGcm, - _ => { - crate::common::log::warn!("no AEAD encryption algorithm is available, using INSECURE XOR"); - EncryptionAlgorithm::Xor - } - } - } -} - -#[auto_impl::auto_impl(Box, &)] -pub trait ConfigLoader: Send + Sync { - fn get_id(&self) -> uuid::Uuid; - fn set_id(&self, id: uuid::Uuid); - - fn get_hostname(&self) -> String; - fn set_hostname(&self, name: Option); - - fn get_inst_name(&self) -> String; - fn set_inst_name(&self, name: String); - - fn get_netns(&self) -> Option; - fn set_netns(&self, ns: Option); - - fn get_ipv4(&self) -> Option; - fn set_ipv4(&self, addr: Option); - - fn get_ipv6(&self) -> Option; - fn set_ipv6(&self, addr: Option); - - fn get_ipv6_public_addr_provider(&self) -> bool; - fn set_ipv6_public_addr_provider(&self, enabled: bool); - - fn get_ipv6_public_addr_auto(&self) -> bool; - fn set_ipv6_public_addr_auto(&self, enabled: bool); - - fn get_ipv6_public_addr_prefix(&self) -> Option; - fn set_ipv6_public_addr_prefix(&self, prefix: Option); - - fn get_dhcp(&self) -> bool; - fn set_dhcp(&self, dhcp: bool); - - fn add_proxy_cidr( - &self, - cidr: cidr::Ipv4Cidr, - mapped_cidr: Option, - ) -> Result<(), anyhow::Error>; - fn remove_proxy_cidr(&self, cidr: cidr::Ipv4Cidr); - fn clear_proxy_cidrs(&self); - fn get_proxy_cidrs(&self) -> Vec; - - fn get_network_identity(&self) -> NetworkIdentity; - fn set_network_identity(&self, identity: NetworkIdentity); - - fn get_listener_uris(&self) -> Vec; - - fn get_peers(&self) -> Vec; - fn set_peers(&self, peers: Vec); - - fn get_listeners(&self) -> Option>; - fn set_listeners(&self, listeners: Vec); - - fn get_mapped_listeners(&self) -> Vec; - fn set_mapped_listeners(&self, listeners: Option>); - - fn get_vpn_portal_config(&self) -> Option; - fn set_vpn_portal_config(&self, config: VpnPortalConfig); - - fn get_flags(&self) -> Flags; - fn set_flags(&self, flags: Flags); - - fn get_exit_nodes(&self) -> Vec; - fn set_exit_nodes(&self, nodes: Vec); - - fn get_routes(&self) -> Option>; - fn set_routes(&self, routes: Option>); - - fn get_socks5_portal(&self) -> Option; - fn set_socks5_portal(&self, addr: Option); - - fn get_port_forwards(&self) -> Vec; - fn set_port_forwards(&self, forwards: Vec); - - fn get_acl(&self) -> Option; - fn set_acl(&self, acl: Option); - - fn get_tcp_whitelist(&self) -> Vec; - fn set_tcp_whitelist(&self, whitelist: Vec); - - fn get_udp_whitelist(&self) -> Vec; - fn set_udp_whitelist(&self, whitelist: Vec); - - fn get_stun_servers(&self) -> Option>; - fn set_stun_servers(&self, servers: Option>); - - fn get_stun_servers_v6(&self) -> Option>; - fn set_stun_servers_v6(&self, servers: Option>); - - fn get_secure_mode(&self) -> Option; - fn set_secure_mode(&self, secure_mode: Option); - - fn get_credential_file(&self) -> Option { - None - } - fn set_credential_file(&self, _path: Option) {} - - fn get_network_config_source(&self) -> ConfigSource { - ConfigSource::User - } - fn set_network_config_source(&self, _source: Option) {} - - fn dump(&self) -> String; -} - -pub trait LoggingConfigLoader { - fn get_file_logger_config(&self) -> FileLoggerConfig; - - fn get_console_logger_config(&self) -> ConsoleLoggerConfig; -} - -pub type NetworkSecretDigest = [u8; 32]; - -#[derive(Debug, Clone, Deserialize, Serialize)] -pub struct NetworkIdentity { - pub network_name: String, - pub network_secret: Option, - #[serde(skip)] - pub network_secret_digest: Option, -} - -#[derive(Debug, Clone, Copy, Deserialize, Serialize, PartialEq, Eq, Default)] -#[serde(rename_all = "snake_case")] -pub enum ConfigSource { - #[default] - User, - Web, -} - -impl ConfigSource { - pub fn as_str(self) -> &'static str { - match self { - Self::User => "user", - Self::Web => "web", - } - } - - pub fn from_rpc(source: i32) -> Option { - match RpcConfigSource::try_from(source).ok() { - Some(RpcConfigSource::Web) => Some(Self::Web), - Some(RpcConfigSource::User) => Some(Self::User), - _ => None, - } - } - - pub fn to_rpc(self) -> i32 { - match self { - Self::User => RpcConfigSource::User as i32, - Self::Web => RpcConfigSource::Web as i32, - } - } -} - -impl std::str::FromStr for ConfigSource { - type Err = String; - - fn from_str(s: &str) -> Result { - match s { - "user" => Ok(Self::User), - "web" => Ok(Self::Web), - other => Err(format!("unknown network config source: {other}")), - } - } -} - -#[derive(Debug, Clone, Deserialize, Serialize, PartialEq)] -struct ConfigSourceConfig { - source: ConfigSource, -} - -#[derive(Eq, PartialEq, Hash)] -struct NetworkIdentityWithOnlyDigest { - network_name: String, - network_secret_digest: Option, -} - -impl From for NetworkIdentityWithOnlyDigest { - fn from(identity: NetworkIdentity) -> Self { - if identity.network_secret_digest.is_some() { - Self { - network_name: identity.network_name, - network_secret_digest: identity.network_secret_digest, - } - } else if identity.network_secret.is_some() { - let mut network_secret_digest = [0u8; 32]; - generate_digest_from_str( - &identity.network_name, - identity.network_secret.as_ref().unwrap(), - &mut network_secret_digest, - ); - Self { - network_name: identity.network_name, - network_secret_digest: Some(network_secret_digest), - } - } else { - Self { - network_name: identity.network_name, - network_secret_digest: None, - } - } - } -} - -impl PartialEq for NetworkIdentity { - fn eq(&self, other: &Self) -> bool { - let self_with_digest = NetworkIdentityWithOnlyDigest::from(self.clone()); - let other_with_digest = NetworkIdentityWithOnlyDigest::from(other.clone()); - self_with_digest == other_with_digest - } -} - -impl Eq for NetworkIdentity {} - -impl std::hash::Hash for NetworkIdentity { - fn hash(&self, state: &mut H) { - let self_with_digest = NetworkIdentityWithOnlyDigest::from(self.clone()); - self_with_digest.hash(state); - } -} - -impl NetworkIdentity { - pub fn new(network_name: String, network_secret: String) -> Self { - let mut network_secret_digest = [0u8; 32]; - generate_digest_from_str(&network_name, &network_secret, &mut network_secret_digest); - - NetworkIdentity { - network_name, - network_secret: Some(network_secret), - network_secret_digest: Some(network_secret_digest), - } - } - - /// Create a NetworkIdentity for a credential node (no network_secret). - /// The node identifies by network_name only and authenticates via credential keypair. - pub fn new_credential(network_name: String) -> Self { - NetworkIdentity { - network_name, - network_secret: None, - network_secret_digest: None, - } - } -} - -impl Default for NetworkIdentity { - fn default() -> Self { - Self::new("default".to_string(), "".to_string()) - } -} - -#[derive(Debug, Clone, Deserialize, Serialize, PartialEq)] -pub struct PeerConfig { - pub uri: url::Url, - pub peer_public_key: Option, -} - -#[derive(Debug, Clone, Deserialize, Serialize, PartialEq)] -pub struct ProxyNetworkConfig { - pub cidr: cidr::Ipv4Cidr, // the CIDR of the proxy network - pub mapped_cidr: Option, // allow remap the proxy CIDR to another CIDR - pub allow: Option>, -} - -#[derive(Debug, Clone, Deserialize, Serialize, PartialEq, Default)] -pub struct FileLoggerConfig { - pub level: Option, - pub file: Option, - pub dir: Option, - pub size_mb: Option, - pub count: Option, -} - -#[derive(Debug, Clone, Deserialize, Serialize, PartialEq, Default)] -pub struct ConsoleLoggerConfig { - pub level: Option, -} - -#[derive(Debug, Clone, Deserialize, Serialize, PartialEq, derive_builder::Builder)] -pub struct LoggingConfig { - #[builder(setter(into, strip_option), default = None)] - pub file_logger: Option, - #[builder(setter(into, strip_option), default = None)] - pub console_logger: Option, -} - -impl LoggingConfigLoader for &LoggingConfig { - fn get_file_logger_config(&self) -> FileLoggerConfig { - self.file_logger.clone().unwrap_or_default() - } - - fn get_console_logger_config(&self) -> ConsoleLoggerConfig { - self.console_logger.clone().unwrap_or_default() - } -} - -#[derive(Debug, Clone, Deserialize, Serialize, PartialEq)] -pub struct VpnPortalConfig { - pub client_cidr: cidr::Ipv4Cidr, - pub wireguard_listen: SocketAddr, -} - -#[derive(Debug, Clone, Deserialize, Serialize, PartialEq, Eq, Hash)] -pub struct PortForwardConfig { - pub bind_addr: SocketAddr, - pub dst_addr: SocketAddr, - pub proto: String, -} - -impl From for PortForwardConfig { - fn from(config: PortForwardConfigPb) -> Self { - PortForwardConfig { - bind_addr: config.bind_addr.unwrap_or_default().into(), - dst_addr: config.dst_addr.unwrap_or_default().into(), - proto: match SocketType::try_from(config.socket_type) { - Ok(SocketType::Tcp) => "tcp".to_string(), - Ok(SocketType::Udp) => "udp".to_string(), - _ => "tcp".to_string(), - }, - } - } -} - -impl From for PortForwardConfigPb { - fn from(val: PortForwardConfig) -> Self { - PortForwardConfigPb { - bind_addr: Some(val.bind_addr.into()), - dst_addr: Some(val.dst_addr.into()), - socket_type: match val.proto.to_lowercase().as_str() { - "tcp" => SocketType::Tcp as i32, - "udp" => SocketType::Udp as i32, - _ => SocketType::Tcp as i32, - }, - } - } -} - -pub fn process_secure_mode_cfg(mut user_cfg: SecureModeConfig) -> anyhow::Result { - if !user_cfg.enabled { - return Ok(user_cfg); - } - - let private_key = if user_cfg.local_private_key.is_none() { - // if no private key, generate random one - let private = x25519_dalek::StaticSecret::random_from_rng(rand::rngs::OsRng); - user_cfg.local_private_key = Some(BASE64_STANDARD.encode(private.clone().as_bytes())); - private - } else { - // check if private key is valid - user_cfg.private_key()? - }; - - let public = x25519_dalek::PublicKey::from(&private_key); - - match user_cfg.local_public_key { - None => { - user_cfg.local_public_key = Some(BASE64_STANDARD.encode(public.as_bytes())); - } - Some(ref user_pub) => { - let public = user_cfg.public_key()?; - if *user_pub != BASE64_STANDARD.encode(public.as_bytes()) { - return Err(anyhow::anyhow!( - "local public key {} does not match generated public key {}", - user_pub, - BASE64_STANDARD.encode(public.as_bytes()) - )); - } - } - } - - Ok(user_cfg) -} - -#[derive(Debug, Clone, PartialEq, Deserialize, Serialize)] -struct Config { - netns: Option, - hostname: Option, - instance_name: Option, - instance_id: Option, - ipv4: Option, - ipv6: Option, - ipv6_public_addr_provider: Option, - ipv6_public_addr_auto: Option, - ipv6_public_addr_prefix: Option, - dhcp: Option, - network_identity: Option, - listeners: Option>, - mapped_listeners: Option>, - exit_nodes: Option>, - - peer: Option>, - proxy_network: Option>, - - vpn_portal_config: Option, - - routes: Option>, - - socks5_proxy: Option, - - port_forward: Option>, - - secure_mode: Option, - - flags: Option>, - - #[serde(skip)] - flags_struct: Option, - - acl: Option, - - tcp_whitelist: Option>, - udp_whitelist: Option>, - stun_servers: Option>, - stun_servers_v6: Option>, - - credential_file: Option, - source: Option, -} - -fn format_toml_parse_error(source_name: &str, config_str: &str, error: &toml::de::Error) -> String { - let message = format!("failed to parse config TOML from {source_name}"); - - let Some(span) = error.span() else { - return format!("{message}\ndetail: {error}"); - }; - - let mut output = Vec::new(); - let report = Report::build(ReportKind::Error, (source_name, span.clone())) - .with_config( - AriadneConfig::default() - .with_color(false) - .with_char_set(CharSet::Ascii) - .with_index_type(IndexType::Byte), - ) - .with_message(&message) - .with_label(Label::new((source_name, span)).with_message(error.message())) - .finish(); - - if report - .write((source_name, Source::from(config_str)), &mut output) - .is_ok() - { - String::from_utf8_lossy(&output).into_owned() - } else { - format!("{message}\ndetail: {error}") - } -} - -#[derive(Debug, Clone)] -pub struct TomlConfigLoader { - config: Arc>, -} - -impl Default for TomlConfigLoader { - fn default() -> Self { - TomlConfigLoader::new_from_str("").unwrap() - } -} - -impl TomlConfigLoader { - fn normalize_config_source(config: &mut Config) { - if matches!( - config.source.as_ref().map(|source| source.source), - Some(ConfigSource::User) - ) { - config.source = None; - } - } - - pub fn new_from_str(config_str: &str) -> Result { - Self::new_from_str_with_source("inline config", config_str) - } - - pub fn new(config_path: &PathBuf) -> Result { - let config_str = std::fs::read_to_string(config_path) - .with_context(|| format!("failed to read config file: {}", config_path.display()))?; - - let source_name = config_path.display().to_string(); - Self::new_from_str_with_source(&source_name, &config_str) - } - - pub(crate) fn new_from_str_with_source( - source_name: &str, - config_str: &str, - ) -> Result { - let mut config = toml::de::from_str::(config_str).map_err(|err| { - let message = format_toml_parse_error(source_name, config_str, &err); - anyhow::Error::new(err).context(message) - })?; - - Self::normalize_config_source(&mut config); - - Self::new_from_config(config).map_err(|err| { - let message = format!("failed to load config from {source_name}: {err}"); - err.context(message) - }) - } - - fn new_from_config(mut config: Config) -> Result { - config.flags_struct = Some( - Self::gen_flags(config.flags.clone().unwrap_or_default()) - .context("failed to parse flags")?, - ); - let has_network_identity = config.network_identity.is_some(); - - let config = TomlConfigLoader { - config: Arc::new(Mutex::new(config)), - }; - - let old_ns = config.get_network_identity(); - - // Detect credential mode: secure_mode enabled + no network_secret in TOML - let is_credential = has_network_identity - && config - .get_secure_mode() - .map(|sm| sm.enabled) - .unwrap_or(false) - && old_ns - .network_secret - .as_deref() - .is_none_or(|s| s.is_empty()); - - if is_credential { - config.set_network_identity(NetworkIdentity::new_credential(old_ns.network_name)); - } else { - config.set_network_identity(NetworkIdentity::new( - old_ns.network_name, - old_ns.network_secret.unwrap_or_default(), - )); - } - - Ok(config) - } - - fn gen_flags( - flags_hashmap: serde_json::Map, - ) -> serde_json::Result { - let mut merged_hashmap = match serde_json::to_value(gen_default_flags()) { - Ok(serde_json::Value::Object(map)) => map, - _ => serde_json::Map::new(), - }; - merged_hashmap.extend(flags_hashmap); - serde_json::from_value(serde_json::Value::Object(merged_hashmap)) - } -} - -impl ConfigLoader for TomlConfigLoader { - fn get_inst_name(&self) -> String { - self.config - .lock() - .unwrap() - .instance_name - .clone() - .unwrap_or("default".to_string()) - } - - fn set_inst_name(&self, name: String) { - self.config.lock().unwrap().instance_name = Some(name); - } - - fn get_hostname(&self) -> String { - let hostname = self.config.lock().unwrap().hostname.clone(); - - match hostname { - Some(hostname) => { - let hostname = hostname - .chars() - .filter(|c| !c.is_control()) - .take(32) - .collect::(); - - if !hostname.is_empty() { - self.set_hostname(Some(hostname.clone())); - hostname - } else { - self.set_hostname(None); - gethostname::gethostname().to_string_lossy().to_string() - } - } - None => gethostname::gethostname().to_string_lossy().to_string(), - } - } - - fn set_hostname(&self, name: Option) { - self.config.lock().unwrap().hostname = name; - } - - fn get_netns(&self) -> Option { - self.config.lock().unwrap().netns.clone() - } - - fn set_netns(&self, ns: Option) { - self.config.lock().unwrap().netns = ns; - } - - fn get_ipv4(&self) -> Option { - let locked_config = self.config.lock().unwrap(); - locked_config - .ipv4 - .as_ref() - .and_then(|s| s.parse().ok()) - .map(|c: cidr::Ipv4Inet| { - if c.network_length() == 32 { - cidr::Ipv4Inet::new(c.address(), 24).unwrap() - } else { - c - } - }) - } - - fn set_ipv4(&self, addr: Option) { - self.config.lock().unwrap().ipv4 = addr.map(|addr| addr.to_string()); - } - - fn get_ipv6(&self) -> Option { - let locked_config = self.config.lock().unwrap(); - locked_config.ipv6.as_ref().and_then(|s| s.parse().ok()) - } - - fn set_ipv6(&self, addr: Option) { - self.config.lock().unwrap().ipv6 = addr.map(|addr| addr.to_string()); - } - - fn get_ipv6_public_addr_provider(&self) -> bool { - self.config - .lock() - .unwrap() - .ipv6_public_addr_provider - .unwrap_or_default() - } - - fn set_ipv6_public_addr_provider(&self, enabled: bool) { - self.config.lock().unwrap().ipv6_public_addr_provider = Some(enabled); - } - - fn get_ipv6_public_addr_auto(&self) -> bool { - self.config - .lock() - .unwrap() - .ipv6_public_addr_auto - .unwrap_or_default() - } - - fn set_ipv6_public_addr_auto(&self, enabled: bool) { - self.config.lock().unwrap().ipv6_public_addr_auto = Some(enabled); - } - - fn get_ipv6_public_addr_prefix(&self) -> Option { - let locked_config = self.config.lock().unwrap(); - locked_config - .ipv6_public_addr_prefix - .as_ref() - .and_then(|s| s.parse().ok()) - } - - fn set_ipv6_public_addr_prefix(&self, prefix: Option) { - self.config.lock().unwrap().ipv6_public_addr_prefix = - prefix.map(|prefix| prefix.to_string()); - } - - fn get_dhcp(&self) -> bool { - self.config.lock().unwrap().dhcp.unwrap_or_default() - } - - fn set_dhcp(&self, dhcp: bool) { - self.config.lock().unwrap().dhcp = Some(dhcp); - } - - fn add_proxy_cidr( - &self, - cidr: cidr::Ipv4Cidr, - mapped_cidr: Option, - ) -> Result<(), anyhow::Error> { - let mut locked_config = self.config.lock().unwrap(); - if locked_config.proxy_network.is_none() { - locked_config.proxy_network = Some(vec![]); - } - if let Some(mapped_cidr) = mapped_cidr.as_ref() - && cidr.network_length() != mapped_cidr.network_length() - { - return Err(anyhow::anyhow!( - "Mapped CIDR must have the same network length as the original CIDR: {} != {}", - cidr.network_length(), - mapped_cidr.network_length() - )); - } - // insert if no duplicate - if !locked_config - .proxy_network - .as_ref() - .unwrap() - .iter() - .any(|c| c.cidr == cidr && c.mapped_cidr == mapped_cidr) - { - locked_config - .proxy_network - .as_mut() - .unwrap() - .push(ProxyNetworkConfig { - cidr, - mapped_cidr, - allow: None, - }); - } - Ok(()) - } - - fn remove_proxy_cidr(&self, cidr: cidr::Ipv4Cidr) { - let mut locked_config = self.config.lock().unwrap(); - if let Some(proxy_cidrs) = &mut locked_config.proxy_network { - proxy_cidrs.retain(|c| c.cidr != cidr); - } - } - - fn clear_proxy_cidrs(&self) { - let mut locked_config = self.config.lock().unwrap(); - locked_config.proxy_network = None; - } - - fn get_proxy_cidrs(&self) -> Vec { - self.config - .lock() - .unwrap() - .proxy_network - .as_ref() - .cloned() - .unwrap_or_default() - } - - fn get_id(&self) -> uuid::Uuid { - let mut locked_config = self.config.lock().unwrap(); - match locked_config.instance_id { - Some(id) => id, - None => { - let id = uuid::Uuid::new_v4(); - locked_config.instance_id = Some(id); - id - } - } - } - - fn set_id(&self, id: uuid::Uuid) { - self.config.lock().unwrap().instance_id = Some(id); - } - - fn get_network_identity(&self) -> NetworkIdentity { - self.config - .lock() - .unwrap() - .network_identity - .clone() - .unwrap_or_default() - } - - fn set_network_identity(&self, identity: NetworkIdentity) { - self.config.lock().unwrap().network_identity = Some(identity); - } - - fn get_listener_uris(&self) -> Vec { - self.config - .lock() - .unwrap() - .listeners - .clone() - .unwrap_or_default() - } - - fn get_peers(&self) -> Vec { - self.config.lock().unwrap().peer.clone().unwrap_or_default() - } - - fn set_peers(&self, peers: Vec) { - self.config.lock().unwrap().peer = Some(peers); - } - - fn get_listeners(&self) -> Option> { - self.config.lock().unwrap().listeners.clone() - } - - fn set_listeners(&self, listeners: Vec) { - self.config.lock().unwrap().listeners = Some(listeners); - } - - fn get_mapped_listeners(&self) -> Vec { - self.config - .lock() - .unwrap() - .mapped_listeners - .clone() - .unwrap_or_default() - } - - fn set_mapped_listeners(&self, listeners: Option>) { - self.config.lock().unwrap().mapped_listeners = listeners; - } - - fn get_vpn_portal_config(&self) -> Option { - self.config.lock().unwrap().vpn_portal_config.clone() - } - fn set_vpn_portal_config(&self, config: VpnPortalConfig) { - self.config.lock().unwrap().vpn_portal_config = Some(config); - } - - fn get_flags(&self) -> Flags { - self.config - .lock() - .unwrap() - .flags_struct - .clone() - .unwrap_or_default() - } - - fn set_flags(&self, flags: Flags) { - self.config.lock().unwrap().flags_struct = Some(flags); - } - - fn get_exit_nodes(&self) -> Vec { - self.config - .lock() - .unwrap() - .exit_nodes - .clone() - .unwrap_or_default() - } - - fn set_exit_nodes(&self, nodes: Vec) { - self.config.lock().unwrap().exit_nodes = Some(nodes); - } - - fn get_routes(&self) -> Option> { - self.config.lock().unwrap().routes.clone() - } - - fn set_routes(&self, routes: Option>) { - self.config.lock().unwrap().routes = routes; - } - - fn get_socks5_portal(&self) -> Option { - self.config.lock().unwrap().socks5_proxy.clone() - } - - fn set_socks5_portal(&self, addr: Option) { - self.config.lock().unwrap().socks5_proxy = addr; - } - - fn get_port_forwards(&self) -> Vec { - self.config - .lock() - .unwrap() - .port_forward - .clone() - .unwrap_or_default() - } - - fn set_port_forwards(&self, forwards: Vec) { - self.config.lock().unwrap().port_forward = Some(forwards); - } - - fn get_acl(&self) -> Option { - self.config.lock().unwrap().acl.clone() - } - - fn set_acl(&self, acl: Option) { - self.config.lock().unwrap().acl = acl; - } - - fn get_tcp_whitelist(&self) -> Vec { - self.config - .lock() - .unwrap() - .tcp_whitelist - .clone() - .unwrap_or_default() - } - - fn set_tcp_whitelist(&self, whitelist: Vec) { - self.config.lock().unwrap().tcp_whitelist = Some(whitelist); - } - - fn get_udp_whitelist(&self) -> Vec { - self.config - .lock() - .unwrap() - .udp_whitelist - .clone() - .unwrap_or_default() - } - - fn set_udp_whitelist(&self, whitelist: Vec) { - self.config.lock().unwrap().udp_whitelist = Some(whitelist); - } - - fn get_stun_servers(&self) -> Option> { - self.config.lock().unwrap().stun_servers.clone() - } - - fn set_stun_servers(&self, servers: Option>) { - self.config.lock().unwrap().stun_servers = servers; - } - - fn get_stun_servers_v6(&self) -> Option> { - self.config.lock().unwrap().stun_servers_v6.clone() - } - - fn set_stun_servers_v6(&self, servers: Option>) { - self.config.lock().unwrap().stun_servers_v6 = servers; - } - - fn get_secure_mode(&self) -> Option { - self.config.lock().unwrap().secure_mode.clone() - } - - fn set_secure_mode(&self, secure_mode: Option) { - self.config.lock().unwrap().secure_mode = secure_mode; - } - - fn get_credential_file(&self) -> Option { - self.config.lock().unwrap().credential_file.clone() - } - - fn set_credential_file(&self, path: Option) { - self.config.lock().unwrap().credential_file = path; - } - - fn get_network_config_source(&self) -> ConfigSource { - self.config - .lock() - .unwrap() - .source - .as_ref() - .map(|source| source.source) - .unwrap_or(ConfigSource::User) - } - - fn set_network_config_source(&self, source: Option) { - self.config.lock().unwrap().source = source.and_then(|source| match source { - ConfigSource::User => None, - other => Some(ConfigSourceConfig { source: other }), - }); - } - - fn dump(&self) -> String { - let mut config = self.config.lock().unwrap().clone(); - Self::normalize_config_source(&mut config); - config.flags = Some(flags_diff_from_default(&self.get_flags())); - if config.stun_servers == Some(StunInfoCollector::get_default_servers()) { - config.stun_servers = None; - } - if config.stun_servers_v6 == Some(StunInfoCollector::get_default_servers_v6()) { - config.stun_servers_v6 = None; - } - toml::to_string_pretty(&config).unwrap() - } -} - -#[derive(Clone, Copy, Default)] -pub struct ConfigFilePermission(u8); -impl ConfigFilePermission { - pub const READ_ONLY: u8 = 1 << 0; - pub const NO_DELETE: u8 = 1 << 1; - - pub fn with_flag(self, flag: u8) -> Self { - Self(self.0 | flag) - } - pub fn remove_flag(self, flag: u8) -> Self { - Self(self.0 & !flag) - } - pub fn has_flag(&self, flag: u8) -> bool { - (self.0 & flag) != 0 - } -} -impl From for ConfigFilePermission { - fn from(value: u8) -> Self { - ConfigFilePermission(value) - } -} -impl From for ConfigFilePermission { - fn from(value: u32) -> Self { - ConfigFilePermission(value as u8) - } -} -impl From for u8 { - fn from(value: ConfigFilePermission) -> Self { - value.0 - } -} -impl From for u32 { - fn from(value: ConfigFilePermission) -> Self { - value.0 as u32 - } -} -impl std::fmt::Debug for ConfigFilePermission { - fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { - let mut flags = vec![]; - if self.has_flag(ConfigFilePermission::READ_ONLY) { - flags.push("READ_ONLY"); - } else { - flags.push("EDITABLE"); - } - if self.has_flag(ConfigFilePermission::NO_DELETE) { - flags.push("NO_DELETE"); - } else { - flags.push("DELETABLE"); - } - write!(f, "{}", flags.join("|")) - } -} - -#[derive(Debug, Clone)] -pub struct ConfigFileControl { - pub path: Option, - pub permission: ConfigFilePermission, -} - -impl ConfigFileControl { - pub const STATIC_CONFIG: ConfigFileControl = Self { - path: None, - permission: ConfigFilePermission( - ConfigFilePermission::READ_ONLY | ConfigFilePermission::NO_DELETE, - ), - }; - - pub fn new(path: Option, permission: ConfigFilePermission) -> Self { - ConfigFileControl { path, permission } - } - - pub async fn from_path(path: PathBuf) -> Self { - let read_only = if let Ok(metadata) = tokio::fs::metadata(&path).await { - metadata.permissions().readonly() - } else { - true - }; - Self::new( - Some(path), - if read_only { - ConfigFilePermission(ConfigFilePermission::READ_ONLY) - } else { - ConfigFilePermission(0) - }, - ) - } - - pub fn is_read_only(&self) -> bool { - self.permission.has_flag(ConfigFilePermission::READ_ONLY) - } - pub fn set_read_only(&mut self, read_only: bool) { +#[cfg(feature = "management-rpc")] +pub use easytier_core::management::{ConfigFileControl, ConfigFilePermission}; + +#[cfg(feature = "management")] +pub async fn config_file_control_from_path(path: PathBuf) -> ConfigFileControl { + let read_only = tokio::fs::metadata(&path) + .await + .map(|metadata| metadata.permissions().readonly()) + .unwrap_or(true); + ConfigFileControl::new( + Some(path), if read_only { - self.permission = self.permission.with_flag(ConfigFilePermission::READ_ONLY); + ConfigFilePermission::from(ConfigFilePermission::READ_ONLY) } else { - self.permission = self.permission.remove_flag(ConfigFilePermission::READ_ONLY); - } - } - - pub fn is_no_delete(&self) -> bool { - self.permission.has_flag(ConfigFilePermission::NO_DELETE) - } - pub fn set_no_delete(&mut self, no_delete: bool) { - if no_delete { - self.permission = self.permission.with_flag(ConfigFilePermission::NO_DELETE); - } else { - self.permission = self.permission.remove_flag(ConfigFilePermission::NO_DELETE); - } - } - - pub fn is_deletable(&self) -> bool { - !self.is_no_delete() - } + ConfigFilePermission::default() + }, + ) } +#[cfg(feature = "management")] pub async fn load_config_from_file( config_file: &PathBuf, config_dir: Option<&PathBuf>, @@ -1276,7 +67,7 @@ pub async fn load_config_from_file( ) -> Result<(TomlConfigLoader, ConfigFileControl), anyhow::Error> { if config_file.as_os_str() == "-" { let mut stdin = String::new(); - _ = tokio::io::stdin() + tokio::io::stdin() .read_to_string(&mut stdin) .await .context("failed to read config from stdin")?; @@ -1287,53 +78,32 @@ pub async fn load_config_from_file( let config_str = tokio::fs::read_to_string(config_file) .await .with_context(|| format!("failed to read config file: {}", config_file.display()))?; - let (expanded_config_str, uses_env_vars) = if disable_env_parsing { - (config_str.clone(), false) + (config_str, false) } else { env_parser::expand_env_vars(&config_str) }; if disable_env_parsing { - tracing::info!( - "Environment variable parsing is disabled for config file: {:?}", - config_file - ); - } - - if uses_env_vars { - tracing::info!( - "Environment variables detected and expanded in config file: {:?}", - config_file - ); + tracing::info!(?config_file, "environment variable parsing is disabled"); + } else if uses_env_vars { + tracing::info!(?config_file, "environment variables detected and expanded"); } let source_name = config_file.display().to_string(); let config = TomlConfigLoader::new_from_str_with_source(&source_name, &expanded_config_str)?; - - let mut control = ConfigFileControl::from_path(config_file.clone()).await; + let mut control = config_file_control_from_path(config_file.clone()).await; if uses_env_vars { control.set_read_only(true); control.set_no_delete(true); - tracing::info!( - "Config file {:?} uses environment variables, marked as READ_ONLY and NO_DELETE", - config_file - ); } else if control.is_read_only() { control.set_no_delete(true); } else if let Some(config_dir) = config_dir { - if let Some(config_file_dir) = config_file.parent() { - // if the config file is in the config dir and named as the instance id, it can be saved remotely - if config_file_dir == config_dir - && config_file.file_stem() == Some(config.get_id().to_string().as_ref()) - && config_file.extension() == Some(std::ffi::OsStr::new("toml")) - { - control.set_no_delete(false); - } else { - control.set_no_delete(true); - } - } + let is_managed_file = config_file.parent() == Some(config_dir.as_path()) + && config_file.file_stem() == Some(config.get_id().to_string().as_ref()) + && config_file.extension() == Some(std::ffi::OsStr::new("toml")); + control.set_no_delete(!is_managed_file); } else { control.set_no_delete(true); } @@ -1342,495 +112,48 @@ pub async fn load_config_from_file( } #[cfg(test)] -pub mod tests { - use super::*; - use crate::tests::{remove_env_var, set_env_var}; - use std::io::Write; - use std::path::PathBuf; +mod tests { + use std::io::Write as _; + use tempfile::NamedTempFile; - #[test] - fn invalid_toml_error_includes_location_and_source_line() { - let error = TomlConfigLoader::new_from_str("dhcp = \"yes\"").unwrap_err(); - let display = error.to_string(); - - assert!(display.contains("failed to parse config TOML")); - assert!(display.contains("inline config")); - assert!(display.contains("dhcp = \"yes\"")); - assert!(display.contains("^")); - assert!(display.contains("invalid type: string")); - assert!(!display.contains("")); - assert!( - error - .chain() - .any(|err| err.downcast_ref::().is_some()) - ); - } + use super::*; #[test] - fn invalid_file_toml_error_includes_config_source() { - let mut config_file = NamedTempFile::new().unwrap(); - writeln!(config_file, "dhcp = \"yes\"").unwrap(); + fn path_adapter_preserves_file_name_in_parse_error() { + let mut file = NamedTempFile::new().unwrap(); + writeln!(file, "dhcp = \"yes\"").unwrap(); - let error = TomlConfigLoader::new(&config_file.path().to_path_buf()).unwrap_err(); - let error = error.to_string(); - - assert!(error.contains(config_file.path().to_string_lossy().as_ref())); - assert!(error.contains("failed to parse config TOML")); - assert!(error.contains("dhcp = \"yes\"")); - assert!(error.contains("^")); - assert!(error.contains("invalid type: string")); - assert!(!error.contains("")); - } - - #[test] - fn invalid_stdin_toml_error_includes_config_source_in_display() { - let error = TomlConfigLoader::new_from_str_with_source("stdin", "dhcp = \"yes\"") + let error = load_toml_config_from_path(&file.path().to_path_buf()) .unwrap_err() .to_string(); - assert!(error.contains("stdin")); - assert!(error.contains("failed to parse config TOML")); + assert!(error.contains(file.path().to_string_lossy().as_ref())); assert!(error.contains("dhcp = \"yes\"")); - assert!(error.contains("^")); - assert!(error.contains("invalid type: string")); - assert!(!error.contains("")); - } - - #[test] - fn invalid_toml_error_handles_non_ascii_before_error() { - let error = TomlConfigLoader::new_from_str("hostname = \"节点\"\ndhcp = \"yes\"") - .unwrap_err() - .to_string(); - - assert!(error.contains("dhcp = \"yes\"")); - assert!(error.contains("^")); - assert!(error.contains("invalid type: string")); - } - - #[test] - fn invalid_toml_error_handles_non_ascii_before_error_on_same_line() { - let error = TomlConfigLoader::new_from_str("hostname = \"节点\" dhcp = \"yes\"") - .unwrap_err() - .to_string(); - - assert!(error.contains("failed to parse config TOML")); - assert!(error.contains("inline config:1:")); - assert!(error.contains("hostname = \"节点\" dhcp = \"yes\"")); - assert!(error.contains("expected newline")); - assert!(!error.contains("")); - } - - #[test] - fn invalid_file_flags_error_includes_config_source_in_display() { - let mut config_file = NamedTempFile::new().unwrap(); - writeln!(config_file, "[flags]").unwrap(); - writeln!(config_file, "socket_mark = \"bad\"").unwrap(); - - let error = TomlConfigLoader::new(&config_file.path().to_path_buf()).unwrap_err(); - - let display = error.to_string(); - assert!(display.contains(config_file.path().to_string_lossy().as_ref())); - assert!(display.contains("failed to load config")); - assert!(display.contains("failed to parse flags")); - - // with_context preserves the cause chain so callers can inspect the root reason. - let chain: Vec = error.chain().map(|e| e.to_string()).collect(); - assert!(chain.iter().any(|m| m.contains("failed to parse flags"))); - } - - #[test] - fn socket_mark_config_file_roundtrip_none_some_and_zero() { - // Omitting the flag leaves socket_mark unset (None) -> SO_MARK untouched. - let cfg = TomlConfigLoader::new_from_str( - r#" -[network_identity] -network_name = "n" -network_secret = "s" -"#, - ) - .unwrap(); - assert_eq!(cfg.get_flags().socket_mark, None); - - // socket_mark = 0 is a legitimate value distinct from "unset". - let cfg = TomlConfigLoader::new_from_str( - r#" -[network_identity] -network_name = "n" -network_secret = "s" - -[flags] -socket_mark = 0 -"#, - ) - .unwrap(); - assert_eq!(cfg.get_flags().socket_mark, Some(0)); - - // A non-zero mark round-trips as Some(v). - let cfg = TomlConfigLoader::new_from_str( - r#" -[network_identity] -network_name = "n" -network_secret = "s" - -[flags] -socket_mark = 66 -"#, - ) - .unwrap(); - assert_eq!(cfg.get_flags().socket_mark, Some(66)); - - // set_flags(None) must serialize back through gen_config without - // resurrecting a value (guards the gen_flags merge against dropping - // the key when the serialized default is null). - cfg.set_flags(Flags { - socket_mark: None, - ..cfg.get_flags() - }); - assert_eq!(cfg.get_flags().socket_mark, None); - } - - #[test] - fn dump_preserves_flags_that_differ_from_easytier_defaults() { - let cfg = TomlConfigLoader::default(); - let mut flags = gen_default_flags(); - flags.dev_name = "et_test".to_string(); - flags.enable_quic_proxy = true; - flags.disable_tcp_hole_punching = true; - flags.disable_sym_hole_punching = true; - flags.multi_thread = false; - flags.bind_device = false; - flags.enable_ipv6 = false; - flags.relay_network_whitelist = "".to_string(); - flags.mtu = 0; - flags.socket_mark = Some(0); - cfg.set_flags(flags); - - let dumped = cfg.dump(); - - assert!(dumped.contains("dev_name = \"et_test\"")); - assert!(dumped.contains("enable_quic_proxy = true")); - assert!(dumped.contains("disable_tcp_hole_punching = true")); - assert!(dumped.contains("disable_sym_hole_punching = true")); - assert!(dumped.contains("multi_thread = false")); - assert!(dumped.contains("bind_device = false")); - assert!(dumped.contains("enable_ipv6 = false")); - assert!(dumped.contains("relay_network_whitelist = \"\"")); - assert!(dumped.contains("mtu = 0")); - assert!(dumped.contains("socket_mark = 0")); - - let reloaded = TomlConfigLoader::new_from_str(&dumped).unwrap(); - let reloaded_flags = reloaded.get_flags(); - assert_eq!(reloaded_flags.dev_name, "et_test"); - assert!(reloaded_flags.enable_quic_proxy); - assert!(reloaded_flags.disable_tcp_hole_punching); - assert!(reloaded_flags.disable_sym_hole_punching); - assert!(!reloaded_flags.multi_thread); - assert!(!reloaded_flags.bind_device); - assert!(!reloaded_flags.enable_ipv6); - assert_eq!(reloaded_flags.relay_network_whitelist, ""); - assert_eq!(reloaded_flags.mtu, 0); - assert_eq!(reloaded_flags.socket_mark, Some(0)); - } - - #[test] - fn test_stun_servers_config() { - let config = TomlConfigLoader::default(); - let stun_servers = config.get_stun_servers(); - assert!(stun_servers.is_none()); - - // Test setting custom stun servers - let custom_servers = vec!["txt:stun.easytier.cn".to_string()]; - config.set_stun_servers(Some(custom_servers.clone())); - - let retrieved_servers = config.get_stun_servers(); - assert_eq!(retrieved_servers.unwrap(), custom_servers); - } - - #[test] - fn test_stun_servers_toml_parsing() { - let config_str = r#" -instance_name = "test" -stun_servers = [ - "stun.l.google.com:19302", - "stun1.l.google.com:19302", - "txt:stun.easytier.cn" -]"#; - - let config = TomlConfigLoader::new_from_str(config_str).unwrap(); - let stun_servers = config.get_stun_servers().unwrap(); - - assert_eq!(stun_servers.len(), 3); - assert_eq!(stun_servers[0], "stun.l.google.com:19302"); - assert_eq!(stun_servers[1], "stun1.l.google.com:19302"); - assert_eq!(stun_servers[2], "txt:stun.easytier.cn"); - } - - #[test] - fn test_network_config_source_toml_roundtrip() { - let config = TomlConfigLoader::default(); - assert_eq!(config.get_network_config_source(), ConfigSource::User); - - config.set_network_config_source(Some(ConfigSource::Web)); - let dumped = config.dump(); - - assert!(dumped.contains("[source]")); - assert!(dumped.contains("source = \"web\"")); - - let loaded = TomlConfigLoader::new_from_str(&dumped).unwrap(); - assert_eq!(loaded.get_network_config_source(), ConfigSource::Web); - } - - #[test] - fn test_toml_credential_mode_omits_network_secret() { - for network_secret in ["", r#"network_secret = """#] { - let config = TomlConfigLoader::new_from_str(&format!( - r#" -[network_identity] -network_name = "credential-network" -{network_secret} - -[secure_mode] -enabled = true -"# - )) - .unwrap(); - - let identity = config.get_network_identity(); - assert_eq!(identity.network_name, "credential-network"); - assert_eq!(identity.network_secret, None); - assert_eq!(identity.network_secret_digest, None); - assert!(!config.dump().contains("network_secret")); - } - } - - #[test] - fn test_toml_secure_mode_without_network_identity_uses_default_secret() { - let config = TomlConfigLoader::new_from_str( - r#" -[secure_mode] -enabled = true -"#, - ) - .unwrap(); - - let identity = config.get_network_identity(); - assert_eq!(identity.network_name, "default"); - assert_eq!(identity.network_secret.as_deref(), Some("")); - assert!(identity.network_secret_digest.is_some()); - } - - #[test] - fn test_parse_mapped_listener_urls_allows_ws_without_port() { - let parsed = parse_mapped_listener_urls(&[ - "ws://example.com".to_string(), - "wss://example.com/path".to_string(), - ]) - .unwrap(); - - assert_eq!(parsed.len(), 2); - assert_eq!(parsed[0].scheme(), "ws"); - assert_eq!(parsed[0].port(), None); - assert_eq!(parsed[1].scheme(), "wss"); - assert_eq!(parsed[1].port(), None); - } - - #[test] - fn test_parse_mapped_listener_urls_allows_tcp_without_port() { - let parsed = parse_mapped_listener_urls(&["tcp://127.0.0.1".to_string()]).unwrap(); - - assert_eq!(parsed.len(), 1); - assert_eq!(parsed[0].scheme(), "tcp"); - assert_eq!(parsed[0].port(), None); - } - - #[test] - fn test_parse_mapped_listener_urls_requires_port_for_non_ip_scheme() { - let err = parse_mapped_listener_urls(&["ring://peer-id".to_string()]).unwrap_err(); - - assert!(err.to_string().contains("mapped listener port is missing")); - } - - #[test] - fn test_acl_toml_rule_uses_defaults_for_omitted_fields() { - use crate::proto::acl::{Action, ChainType, Protocol}; - - let config_str = r#" -[[acl.acl_v1.chains]] -name = "subnet_proxy_protect" -chain_type = 3 -enabled = true -default_action = 2 - -[[acl.acl_v1.chains.rules]] -name = "allow_my_devices" -priority = 1000 -action = 1 -source_ips = ["10.172.192.2/32"] -protocol = 5 -enabled = true -"#; - - let config = TomlConfigLoader::new_from_str(config_str).unwrap(); - let acl = config.get_acl().unwrap(); - let acl_v1 = acl.acl_v1.unwrap(); - let chain = &acl_v1.chains[0]; - let rule = &chain.rules[0]; - - assert_eq!(chain.chain_type, ChainType::Forward as i32); - assert_eq!(chain.default_action, Action::Drop as i32); - assert_eq!(rule.action, Action::Allow as i32); - assert_eq!(rule.protocol, Protocol::Any as i32); - assert_eq!(rule.source_ips, vec!["10.172.192.2/32"]); - assert!(rule.ports.is_empty()); - assert!(rule.source_ports.is_empty()); - assert!(rule.destination_ips.is_empty()); - assert!(rule.source_groups.is_empty()); - assert!(rule.destination_groups.is_empty()); - assert_eq!(rule.rate_limit, 0); - assert_eq!(rule.burst_limit, 0); - assert!(!rule.stateful); - } - - #[test] - fn test_acl_toml_group_can_omit_declares_or_members() { - let declares_only = r#" -[acl.acl_v1.group] - -[[acl.acl_v1.group.declares]] -group_name = "admin" -group_secret = "admin-pw" -"#; - let config = TomlConfigLoader::new_from_str(declares_only).unwrap(); - let group = config.get_acl().unwrap().acl_v1.unwrap().group.unwrap(); - assert_eq!(group.declares.len(), 1); - assert!(group.members.is_empty()); - - let members_only = r#" -[acl.acl_v1.group] -members = ["admin"] -"#; - let config = TomlConfigLoader::new_from_str(members_only).unwrap(); - let group = config.get_acl().unwrap().acl_v1.unwrap().group.unwrap(); - assert!(group.declares.is_empty()); - assert_eq!(group.members, vec!["admin"]); - } - - #[test] - fn test_network_config_source_user_is_implicit() { - let config = TomlConfigLoader::default(); - config.set_network_config_source(Some(ConfigSource::User)); - let dumped = config.dump(); - - assert!(!dumped.contains("[source]")); - - let loaded = TomlConfigLoader::new_from_str(&dumped).unwrap(); - assert_eq!(loaded.get_network_config_source(), ConfigSource::User); - - let explicit_user = TomlConfigLoader::new_from_str( - r#" -[source] -source = "user" -"#, - ) - .unwrap(); - assert_eq!( - explicit_user.get_network_config_source(), - ConfigSource::User - ); - assert!(!explicit_user.dump().contains("[source]")); - } - - #[test] - fn test_ipv6_public_addr_config_roundtrip() { - let config = TomlConfigLoader::default(); - let prefix: cidr::Ipv6Cidr = "2001:db8:100::/64".parse().unwrap(); - - config.set_ipv6_public_addr_provider(true); - config.set_ipv6_public_addr_auto(true); - config.set_ipv6_public_addr_prefix(Some(prefix)); - - assert!(config.get_ipv6_public_addr_provider()); - assert!(config.get_ipv6_public_addr_auto()); - assert_eq!(config.get_ipv6_public_addr_prefix(), Some(prefix)); - - let dumped = config.dump(); - let loaded = TomlConfigLoader::new_from_str(&dumped).unwrap(); - assert!(loaded.get_ipv6_public_addr_provider()); - assert!(loaded.get_ipv6_public_addr_auto()); - assert_eq!(loaded.get_ipv6_public_addr_prefix(), Some(prefix)); } #[tokio::test] - async fn full_example_test() { - let config_str = r#" -instance_name = "default" -instance_id = "87ede5a2-9c3d-492d-9bbe-989b9d07e742" -ipv4 = "10.144.144.10" -listeners = [ "tcp://0.0.0.0:11010", "udp://0.0.0.0:11010" ] -routes = [ "192.168.0.0/16" ] + async fn file_adapter_keeps_os_metadata_outside_core_config() { + let mut file = NamedTempFile::new().unwrap(); + writeln!(file, "instance_name = \"from-file\"").unwrap(); -[network_identity] -network_name = "default" -network_secret = "" + let (config, control) = load_config_from_file(&file.path().to_path_buf(), None, true) + .await + .unwrap(); -[[peer]] -uri = "tcp://public.kkrainbow.top:11010" - -[[peer]] -uri = "udp://192.168.94.33:11010" - -[[proxy_network]] -cidr = "10.147.223.0/24" -allow = ["tcp", "udp", "icmp"] - -[[proxy_network]] -cidr = "10.1.1.0/24" -allow = ["tcp", "icmp"] - -[file_logger] -level = "info" -file = "easytier" -dir = "/tmp/easytier" - -[console_logger] -level = "warn" - -[[port_forward]] -bind_addr = "0.0.0.0:11011" -dst_addr = "192.168.94.33:11011" -proto = "tcp" -"#; - let ret = TomlConfigLoader::new_from_str(config_str); - if let Err(e) = &ret { - println!("{}", e); - } else { - println!("{:?}", ret.as_ref().unwrap()); - } - assert!(ret.is_ok()); - - let ret = ret.unwrap(); - assert_eq!("10.144.144.10/24", ret.get_ipv4().unwrap().to_string()); - - assert_eq!( - vec!["tcp://0.0.0.0:11010", "udp://0.0.0.0:11010"], - ret.get_listener_uris() - .iter() - .map(|u| u.to_string()) - .collect::>() - ); - - assert_eq!( - vec![PortForwardConfig { - bind_addr: "0.0.0.0:11011".parse().unwrap(), - dst_addr: "192.168.94.33:11011".parse().unwrap(), - proto: "tcp".to_string(), - }], - ret.get_port_forwards() - ); - println!("{}", ret.dump()); + assert_eq!(config.get_inst_name(), "from-file"); + assert_eq!(control.path.as_deref(), Some(file.path())); } +} +#[cfg(test)] +mod compatibility_tests { + use std::{io::Write as _, path::PathBuf}; + + use tempfile::NamedTempFile; + + use super::*; + use crate::tests::{remove_env_var, set_env_var}; /// 配置文件环境变量解析功能的集成测试 /// /// 测试范围: @@ -1972,10 +295,7 @@ network_secret = "${DISABLED_TEST_VAR}" "Env var should not be expanded when parsing is disabled" ); - // 验证配置不因环境变量而被标记为只读 - // 注:文件系统权限可能使其只读,但不应因环境变量而只读 - // 这里我们主要验证 NO_DELETE 标记的逻辑 - // 由于没有 config_dir,文件会被标记为 NO_DELETE,但不是因为环境变量 + assert!(!control.is_read_only()); assert!( control.is_no_delete(), "Config should be NO_DELETE due to no config_dir, not env vars" @@ -2068,7 +388,7 @@ network_secret = "${INSTANCE_SECRET}" #[tokio::test] async fn test_real_config_fields_expansion() { // 设置各种实际场景的环境变量 - set_env_var("ET_SECRET", "production-secret-key"); + set_env_var("CONFIG_REAL_SECRET", "production-secret-key"); set_env_var("PEER_HOST", "peer.example.com"); set_env_var("PEER_PORT", "11011"); set_env_var("LISTEN_PORT", "11010"); @@ -2083,7 +403,7 @@ listeners = ["tcp://0.0.0.0:${LISTEN_PORT}"] [network_identity] network_name = "${NETWORK_NAME}" -network_secret = "${ET_SECRET}" +network_secret = "${CONFIG_REAL_SECRET}" [[peer]] uri = "tcp://${PEER_HOST}:${PEER_PORT}" @@ -2120,7 +440,7 @@ uri = "tcp://${PEER_HOST}:${PEER_PORT}" assert!(control.is_no_delete()); // 清理环境变量 - remove_env_var("ET_SECRET"); + remove_env_var("CONFIG_REAL_SECRET"); remove_env_var("PEER_HOST"); remove_env_var("PEER_PORT"); remove_env_var("LISTEN_PORT"); @@ -2204,10 +524,7 @@ network_secret = "${COMPLETELY_UNDEFINED}" "${COMPLETELY_UNDEFINED}" ); - // 注意:由于没有实际替换发生,控制标记不应因环境变量而设置 - // 但会因为其他原因(如没有 config_dir)被标记为 NO_DELETE - // 这里我们主要验证 NO_DELETE 标记的逻辑 - // 由于没有 config_dir,文件会被标记为 NO_DELETE,但不是因为环境变量 + assert!(!control.is_read_only()); assert!(control.is_no_delete()); } diff --git a/easytier/src/common/constants.rs b/easytier/src/common/constants.rs index c4b875fa..5262bb71 100644 --- a/easytier/src/common/constants.rs +++ b/easytier/src/common/constants.rs @@ -1,36 +1,13 @@ -macro_rules! define_global_var { - ($name:ident, $type:ty, $init:expr) => { - pub static $name: once_cell::sync::Lazy> = - once_cell::sync::Lazy::new(|| std::sync::Mutex::new($init)); - }; -} +pub const MANUAL_CONNECTOR_RECONNECT_INTERVAL_MS: u64 = 1000; -#[macro_export] -macro_rules! use_global_var { - ($name:ident) => { - $crate::common::constants::$name.lock().unwrap().to_owned() - }; -} +pub const OSPF_UPDATE_MY_GLOBAL_FOREIGN_NETWORK_INTERVAL_SEC: u64 = 10; -#[macro_export] -macro_rules! set_global_var { - ($name:ident, $val:expr) => { - *$crate::common::constants::$name.lock().unwrap() = $val - }; -} +pub const MAX_DIRECT_CONNS_PER_PEER_IN_FOREIGN_NETWORK: u32 = 3; -define_global_var!(MANUAL_CONNECTOR_RECONNECT_INTERVAL_MS, u64, 1000); - -define_global_var!(OSPF_UPDATE_MY_GLOBAL_FOREIGN_NETWORK_INTERVAL_SEC, u64, 10); - -define_global_var!(MAX_DIRECT_CONNS_PER_PEER_IN_FOREIGN_NETWORK, u32, 3); - -define_global_var!(DIRECT_CONNECT_TO_PUBLIC_SERVER, bool, true); +pub const DIRECT_CONNECT_TO_PUBLIC_SERVER: bool = true; // must make it true in future. -define_global_var!(HMAC_SECRET_DIGEST, bool, false); - -pub const UDP_HOLE_PUNCH_CONNECTOR_SERVICE_ID: u32 = 2; +pub const HMAC_SECRET_DIGEST: bool = false; pub const WIN_SERVICE_WORK_DIR_REG_KEY: &str = "SOFTWARE\\EasyTier\\Service\\WorkDir"; diff --git a/easytier/src/common/credential_manager.rs b/easytier/src/common/credential_manager.rs new file mode 100644 index 00000000..52d740da --- /dev/null +++ b/easytier/src/common/credential_manager.rs @@ -0,0 +1,48 @@ +use std::{path::PathBuf, sync::Arc}; + +use easytier_core::peers::credential_manager::CredentialStorage; + +struct FileCredentialStorage { + path: PathBuf, +} + +impl CredentialStorage for FileCredentialStorage { + fn load(&self) -> anyhow::Result> { + let Ok(serialized) = std::fs::read_to_string(&self.path) else { + return Ok(None); + }; + tracing::info!(path = %self.path.display(), "loaded credentials"); + Ok(Some(serialized)) + } + + fn store(&self, serialized_credentials: &str) -> anyhow::Result<()> { + std::fs::write(&self.path, serialized_credentials)?; + Ok(()) + } +} + +pub(crate) fn runtime_credential_storage( + path: Option, +) -> Option> { + path.map(|path| Arc::new(FileCredentialStorage { path }) as Arc) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn file_storage_round_trips_serialized_credentials() { + let directory = tempfile::tempdir().unwrap(); + let storage = FileCredentialStorage { + path: directory.path().join("credentials.json"), + }; + + assert_eq!(storage.load().unwrap(), None); + storage.store("{\"credential\":true}").unwrap(); + assert_eq!( + storage.load().unwrap().as_deref(), + Some("{\"credential\":true}") + ); + } +} diff --git a/easytier/src/common/dns.rs b/easytier/src/common/dns.rs index 5c047a46..ab382a72 100644 --- a/easytier/src/common/dns.rs +++ b/easytier/src/common/dns.rs @@ -1,19 +1,35 @@ -use std::net::SocketAddr; -use std::sync::Arc; -use std::sync::atomic::AtomicBool; +use std::net::{IpAddr, SocketAddr, ToSocketAddrs as _}; +#[cfg(feature = "dns-resolver")] +use std::{future::Future, io, pin::Pin, sync::Arc, time::Duration}; use anyhow::Context; -use hickory_proto::runtime::TokioRuntimeProvider; +use async_trait::async_trait; +use easytier_core::host::dns::{DnsQuery, DnsRecordResolver, DnsResolver, DnsSrvRecord}; +use easytier_core::socket::SocketContext; +#[cfg(feature = "dns-resolver")] +use hickory_proto::runtime::{RuntimeProvider, TokioRuntimeProvider, iocompat::AsyncIoTokioAsStd}; +#[cfg(feature = "dns-resolver")] use hickory_proto::xfer::Protocol; +#[cfg(feature = "dns-resolver")] use hickory_resolver::config::{LookupIpStrategy, NameServerConfig, ResolverConfig, ResolverOpts}; +#[cfg(feature = "dns-resolver")] use hickory_resolver::name_server::{GenericConnector, TokioConnectionProvider}; +#[cfg(feature = "dns-resolver")] use hickory_resolver::system_conf::read_system_conf; +#[cfg(feature = "dns-resolver")] use hickory_resolver::{Resolver, TokioResolver}; +#[cfg(feature = "dns-resolver")] use once_cell::sync::Lazy; use tokio::net::lookup_host; +#[cfg(feature = "dns-resolver")] +use tokio::net::{TcpSocket, TcpStream, UdpSocket}; use super::error::Error; +use super::netns::NetNS; +#[cfg(feature = "dns-resolver")] +use crate::tunnel::common::apply_socket_mark; +#[cfg(feature = "dns-resolver")] pub fn get_default_resolver_config() -> ResolverConfig { let mut default_resolve_config = ResolverConfig::new(); default_resolve_config.add_name_server(NameServerConfig::new( @@ -27,26 +43,312 @@ pub fn get_default_resolver_config() -> ResolverConfig { default_resolve_config } -pub static ALLOW_USE_SYSTEM_DNS_RESOLVER: Lazy = Lazy::new(|| AtomicBool::new(true)); - -pub static RESOLVER: Lazy>>> = - Lazy::new(|| { - let system_cfg = read_system_conf(); - let mut cfg = get_default_resolver_config(); - let mut opt = ResolverOpts::default(); - if let Ok(s) = system_cfg { - for ns in s.0.name_servers() { - cfg.add_name_server(ns.clone()); - } - opt = s.1; +#[cfg(feature = "dns-resolver")] +fn resolver_config() -> (ResolverConfig, ResolverOpts) { + let system_cfg = read_system_conf(); + let mut config = get_default_resolver_config(); + let mut options = ResolverOpts::default(); + if let Ok(system) = system_cfg { + for name_server in system.0.name_servers() { + config.add_name_server(name_server.clone()); } - opt.ip_strategy = LookupIpStrategy::Ipv4AndIpv6; - let builder = TokioResolver::builder_with_config(cfg, TokioConnectionProvider::default()) - .with_options(opt); - Arc::new(builder.build()) - }); + options = system.1; + } + options.ip_strategy = LookupIpStrategy::Ipv4AndIpv6; + (config, options) +} -pub async fn resolve_txt_record(domain_name: &str) -> Result { +#[cfg(feature = "dns-resolver")] +static RESOLVER: Lazy>>> = Lazy::new(|| { + let (config, options) = resolver_config(); + let builder = TokioResolver::builder_with_config(config, TokioConnectionProvider::default()) + .with_options(options); + Arc::new(builder.build()) +}); + +#[derive(Clone, Debug, PartialEq, Eq, Hash)] +struct RuntimeDnsIoContext { + netns: Option, + socket_mark: Option, +} + +impl RuntimeDnsIoContext { + fn from_socket_context(context: &SocketContext) -> Self { + Self { + netns: context + .netns + .as_ref() + .map(|namespace| namespace.token().to_owned()), + socket_mark: context.socket_mark, + } + } + + fn netns(&self) -> NetNS { + NetNS::new(self.netns.clone()) + } + + #[cfg(feature = "dns-resolver")] + fn is_process_default(&self) -> bool { + self.netns.is_none() && self.socket_mark.is_none() + } +} + +#[cfg(feature = "dns-resolver")] +#[derive(Clone)] +struct RuntimeDnsIoProvider { + inner: TokioRuntimeProvider, + context: RuntimeDnsIoContext, +} + +#[cfg(feature = "dns-resolver")] +impl RuntimeDnsIoProvider { + fn new(context: RuntimeDnsIoContext) -> Self { + Self { + inner: TokioRuntimeProvider::new(), + context, + } + } +} + +#[cfg(feature = "dns-resolver")] +fn create_dns_tcp_socket( + context: &RuntimeDnsIoContext, + server_addr: SocketAddr, + bind_addr: Option, +) -> io::Result { + context.netns().run(|| { + let socket = if server_addr.is_ipv4() { + TcpSocket::new_v4()? + } else { + TcpSocket::new_v6()? + }; + apply_socket_mark(&socket2::SockRef::from(&socket), context.socket_mark) + .map_err(io::Error::other)?; + if let Some(bind_addr) = bind_addr { + socket.bind(bind_addr)?; + } + socket.set_nodelay(true)?; + Ok(socket) + }) +} + +#[cfg(feature = "dns-resolver")] +fn create_dns_udp_socket( + context: &RuntimeDnsIoContext, + local_addr: SocketAddr, +) -> io::Result { + context.netns().run(|| { + let socket = socket2::Socket::new( + socket2::Domain::for_address(local_addr), + socket2::Type::DGRAM, + Some(socket2::Protocol::UDP), + )?; + socket.set_nonblocking(true)?; + apply_socket_mark(&socket, context.socket_mark).map_err(io::Error::other)?; + socket.bind(&socket2::SockAddr::from(local_addr))?; + let socket: std::net::UdpSocket = socket.into(); + UdpSocket::from_std(socket) + }) +} + +#[cfg(feature = "dns-resolver")] +impl RuntimeProvider for RuntimeDnsIoProvider { + type Handle = ::Handle; + type Timer = ::Timer; + type Udp = UdpSocket; + type Tcp = AsyncIoTokioAsStd; + + fn create_handle(&self) -> Self::Handle { + self.inner.create_handle() + } + + fn connect_tcp( + &self, + server_addr: SocketAddr, + bind_addr: Option, + wait_for: Option, + ) -> Pin>>> { + // setns is thread-local. Create the socket synchronously while the + // guard is active, then perform only descriptor I/O after it is gone. + let socket = create_dns_tcp_socket(&self.context, server_addr, bind_addr); + Box::pin(async move { + let socket = socket?; + let wait_for = wait_for.unwrap_or(Duration::from_secs(5)); + match tokio::time::timeout(wait_for, socket.connect(server_addr)).await { + Ok(Ok(stream)) => Ok(AsyncIoTokioAsStd(stream)), + Ok(Err(error)) => Err(error), + Err(_) => Err(io::Error::new( + io::ErrorKind::TimedOut, + format!("connection to {server_addr:?} timed out after {wait_for:?}"), + )), + } + }) + } + + fn bind_udp( + &self, + local_addr: SocketAddr, + _server_addr: SocketAddr, + ) -> Pin>>> { + // Keep namespace switching out of the returned future for the same + // reason as TCP above. + let socket = create_dns_udp_socket(&self.context, local_addr); + Box::pin(async move { socket }) + } +} + +#[cfg(feature = "dns-resolver")] +type ContextualResolver = Resolver>; + +#[derive(Debug, Clone, Copy, Default)] +pub(crate) struct RuntimeDnsResolver; + +impl RuntimeDnsResolver { + pub(crate) fn new() -> Self { + Self + } + + // A netns token can be deleted and recreated. Keep contextual resolvers + // request-scoped so pooled DNS sockets cannot outlive that namespace. + #[cfg(feature = "dns-resolver")] + fn contextual_resolver(context: RuntimeDnsIoContext) -> ContextualResolver { + let (config, options) = resolver_config(); + let provider = GenericConnector::new(RuntimeDnsIoProvider::new(context)); + Resolver::builder_with_config(config, provider) + .with_options(options) + .build() + } + + #[cfg(feature = "dns-resolver")] + async fn resolve_contextual_ips( + &self, + context: RuntimeDnsIoContext, + host: String, + ) -> anyhow::Result> { + if context.socket_mark.is_none() { + // libc DNS cannot attach SO_MARK. It remains usable for a + // namespace-only context when confined to one blocking thread. + let lookup_host = host.clone(); + let netns = context.netns(); + match tokio::task::spawn_blocking(move || { + netns.run(|| { + (lookup_host.as_str(), 0) + .to_socket_addrs() + .map(|addrs| addrs.map(|addr| addr.ip()).collect()) + }) + }) + .await + .context("contextual system DNS task failed")? + { + Ok(addresses) => return Ok(addresses), + Err(error) => tracing::error!(?error, "contextual system dns lookup failed"), + } + } + + let resolver = Self::contextual_resolver(context); + let response = resolver + .lookup_ip(&host) + .await + .with_context(|| format!("contextual hickory lookup_ip failed, host: {host}"))?; + Ok(response.iter().collect()) + } +} + +#[async_trait] +impl DnsResolver for RuntimeDnsResolver { + async fn resolve(&self, query: DnsQuery) -> anyhow::Result> { + let context = RuntimeDnsIoContext::from_socket_context(&query.context); + #[cfg(feature = "dns-resolver")] + { + if context.is_process_default() { + return Ok(resolve_ips(&query.host).await?); + } + return self.resolve_contextual_ips(context, query.host).await; + } + #[cfg(not(feature = "dns-resolver"))] + { + if context.socket_mark.is_some() { + anyhow::bail!("socket-marked DNS requires DNS resolver support"); + } + if context.netns.is_none() { + return Ok(resolve_ips(&query.host).await?); + } + let host = query.host; + let netns = context.netns(); + return tokio::task::spawn_blocking(move || { + netns.run(|| { + (host.as_str(), 0) + .to_socket_addrs() + .map(|addrs| addrs.map(|addr| addr.ip()).collect()) + }) + }) + .await + .context("contextual system DNS task failed")? + .map_err(Into::into); + } + } +} + +#[cfg(feature = "dns-resolver")] +#[async_trait] +impl DnsRecordResolver for RuntimeDnsResolver { + async fn resolve_txt(&self, query: DnsQuery) -> anyhow::Result { + let context = RuntimeDnsIoContext::from_socket_context(&query.context); + if context.is_process_default() { + return Ok(resolve_txt_record(&query.host).await?); + } + + let resolver = Self::contextual_resolver(context); + let response = resolver + .txt_lookup(&query.host) + .await + .with_context(|| format!("txt_lookup failed, domain_name: {}", query.host))?; + let record = response + .iter() + .next() + .with_context(|| format!("no txt record found, domain_name: {}", query.host))?; + let data = record + .txt_data() + .first() + .with_context(|| format!("empty txt record, domain_name: {}", query.host))?; + Ok(String::from_utf8_lossy(data).into_owned()) + } + + async fn resolve_srv(&self, query: DnsQuery) -> anyhow::Result> { + let context = RuntimeDnsIoContext::from_socket_context(&query.context); + let response = if context.is_process_default() { + RESOLVER.srv_lookup(&query.host).await? + } else { + Self::contextual_resolver(context) + .srv_lookup(&query.host) + .await? + }; + Ok(response + .iter() + .map(|record| DnsSrvRecord { + priority: record.priority(), + weight: record.weight(), + port: record.port(), + target: record.target().to_utf8(), + }) + .collect()) + } +} + +#[cfg(not(feature = "dns-resolver"))] +#[async_trait] +impl DnsRecordResolver for RuntimeDnsResolver { + async fn resolve_txt(&self, _query: DnsQuery) -> anyhow::Result { + anyhow::bail!("this build does not include TXT DNS resolution") + } + + async fn resolve_srv(&self, _query: DnsQuery) -> anyhow::Result> { + anyhow::bail!("this build does not include SRV DNS resolution") + } +} + +#[cfg(feature = "dns-resolver")] +async fn resolve_txt_record(domain_name: &str) -> Result { let r = RESOLVER.clone(); let response = r .txt_lookup(domain_name) @@ -67,6 +369,14 @@ pub async fn resolve_txt_record(domain_name: &str) -> Result { pub async fn socket_addrs( url: &url::Url, default_port_number: impl Fn() -> Option, +) -> Result, Error> { + socket_addrs_with_system_resolver(url, default_port_number, true).await +} + +async fn socket_addrs_with_system_resolver( + url: &url::Url, + default_port_number: impl Fn() -> Option, + allow_system_resolver: bool, ) -> Result, Error> { let host = url.host().ok_or(Error::InvalidUrl(url.to_string()))?; let port = url @@ -82,7 +392,7 @@ pub async fn socket_addrs( } let host = host.to_string(); - if ALLOW_USE_SYSTEM_DNS_RESOLVER.load(std::sync::atomic::Ordering::Relaxed) { + if allow_system_resolver { let socket_addr = format!("{}:{}", host, port); match lookup_host(socket_addr).await { Ok(a) => { @@ -92,27 +402,76 @@ pub async fn socket_addrs( } Err(e) => { tracing::error!(?e, "system dns lookup failed"); + #[cfg(not(feature = "dns-resolver"))] + return Err(e.into()); } } } // use hickory_resolver - let ret = RESOLVER.lookup_ip(&host).await.with_context(|| { - format!( - "hickory dns lookup_ip failed, host: {}, port: {}", - host, port - ) - })?; - Ok(ret - .iter() - .map(|ip| SocketAddr::new(ip, port)) - .collect::>()) + #[cfg(feature = "dns-resolver")] + { + let ret = RESOLVER.lookup_ip(&host).await.with_context(|| { + format!( + "hickory dns lookup_ip failed, host: {}, port: {}", + host, port + ) + })?; + Ok(ret + .iter() + .map(|ip| SocketAddr::new(ip, port)) + .collect::>()) + } + #[cfg(not(feature = "dns-resolver"))] + unreachable!("the system resolver error returns above") +} + +#[cfg(feature = "dns-resolver")] +async fn resolve_ips(host: &str) -> Result, Error> { + match lookup_host((host, 0)).await { + Ok(a) => { + let a = a.map(|addr| addr.ip()).collect(); + tracing::debug!(?a, "system dns lookup done"); + return Ok(a); + } + Err(e) => { + tracing::error!(?e, "system dns lookup failed"); + } + } + + let ret = RESOLVER + .lookup_ip(host) + .await + .with_context(|| format!("hickory dns lookup_ip failed, host: {}", host))?; + Ok(ret.iter().collect::>()) +} + +#[cfg(not(feature = "dns-resolver"))] +async fn resolve_ips(host: &str) -> Result, Error> { + Ok(lookup_host((host, 0)) + .await? + .map(|addr| addr.ip()) + .collect()) } #[cfg(test)] mod tests { use super::*; - use guarden::defer; + + #[test] + fn runtime_dns_context_preserves_process_routing_inputs() { + let context = SocketContext::default() + .with_socket_mark(Some(0)) + .with_netns(Some(easytier_core::socket::NetNamespace::new("instance-a"))); + + assert_eq!( + RuntimeDnsIoContext::from_socket_context(&context), + RuntimeDnsIoContext { + netns: Some("instance-a".to_owned()), + socket_mark: Some(0), + } + ); + } #[tokio::test] async fn test_socket_addrs() { @@ -121,11 +480,9 @@ mod tests { assert_eq!(2, addrs.len(), "addrs: {:?}", addrs); println!("addrs: {:?}", addrs); - ALLOW_USE_SYSTEM_DNS_RESOLVER.store(false, std::sync::atomic::Ordering::Relaxed); - defer!( - ALLOW_USE_SYSTEM_DNS_RESOLVER.store(true, std::sync::atomic::Ordering::Relaxed); - ); - let addrs = socket_addrs(&url, || Some(80)).await.unwrap(); + let addrs = socket_addrs_with_system_resolver(&url, || Some(80), false) + .await + .unwrap(); assert_eq!(2, addrs.len(), "addrs: {:?}", addrs); println!("addrs2: {:?}", addrs); } diff --git a/easytier/src/common/error.rs b/easytier/src/common/error.rs index 906499c8..e6209f3b 100644 --- a/easytier/src/common/error.rs +++ b/easytier/src/common/error.rs @@ -1,9 +1,7 @@ use std::{io, result}; use thiserror::Error; -use crate::tunnel; - -use super::PeerId; +use easytier_core::tunnel::TunnelError; #[derive(Error, Debug)] pub enum Error { @@ -15,45 +13,17 @@ pub enum Error { TunError(#[from] tun::Error), #[error("tunnel error {0}")] - TunnelError(#[from] tunnel::TunnelError), - #[error("Peer has no conn, PeerId: {0}")] - PeerNoConnectionError(PeerId), - #[error("RouteError: {0:?}")] - RouteError(Option), + TunnelError(#[from] TunnelError), #[error("Not found")] NotFound, #[error("Invalid Url: {0}")] InvalidUrl(String), #[error("Shell Command error: {0}")] ShellCommandError(String), - // #[error("Rpc listen error: {0}")] - // RpcListenError(String), - #[error("Rpc connect error: {0}")] - RpcConnectError(String), #[error("Timeout error: {0}")] Timeout(#[from] tokio::time::error::Elapsed), - #[error("url in blacklist")] - UrlInBlacklist, - #[error("unknown data store error")] - Unknown, #[error("anyhow error: {0}")] AnyhowError(#[from] anyhow::Error), - - #[error("wait resp error: {0}")] - WaitRespError(String), - - #[error("message decode error: {0}")] - MessageDecodeError(String), - - #[error("secret key error: {0}")] - SecretKeyError(String), - - #[error("noise protocol error: {0}")] - NoiseError(#[from] snow::Error), } pub type Result = result::Result; - -pub type ErrorCollection = crate::utils::error::ErrorCollection; - -// impl From for std:: diff --git a/easytier/src/common/global_ctx.rs b/easytier/src/common/global_ctx.rs index 9073e201..ba1bf268 100644 --- a/easytier/src/common/global_ctx.rs +++ b/easytier/src/common/global_ctx.rs @@ -1,44 +1,29 @@ use std::{ - collections::{BTreeSet, HashMap, hash_map::DefaultHasher}, - hash::Hasher, - net::{IpAddr, SocketAddr}, + collections::HashSet, + net::{IpAddr, Ipv6Addr}, sync::{Arc, Mutex}, - time::{SystemTime, UNIX_EPOCH}, }; use arc_swap::ArcSwap; -use dashmap::DashMap; +use async_trait::async_trait; +use easytier_core::config::PeerId; +use easytier_core::connectivity::composite::ConnectorRuntime as _; +use easytier_core::peers::public_ipv6::PublicIpv6Host; +use easytier_core::socket::{NetNamespace, SocketContext}; +use easytier_core::tunnel::effective_encryption_uses_xor; use super::{ - PeerId, - config::{ConfigLoader, Flags}, + config::{ConfigLoader, Flags, NetworkIdentity}, netns::NetNS, - network::IPCollector, - stun::{StunInfoCollector, StunInfoCollectorTrait}, -}; -use crate::{ - common::{ - config::ProxyNetworkConfig, shrink_dashmap, stats_manager::StatsManager, - token_bucket::TokenBucketManager, - }, - peers::{acl_filter::AclFilter, credential_manager::CredentialManager}, - proto::{ - acl::GroupIdentity, - api::{config::InstanceConfigPatch, instance::PeerConnInfo}, - common::{PeerFeatureFlag, PortForwardConfigPb}, - peer_rpc::PeerGroupInfo, - }, - rpc_service::protected_port, - tunnel::matches_protocol, }; +#[cfg(feature = "management")] +use crate::proto::api::config::InstanceConfigPatch; +use crate::proto::api::instance::PeerConnInfo; +use crate::proto::common::PortForwardConfigPb; use crossbeam::atomic::AtomicCell; -use hmac::{Hmac, Mac}; -use sha2::Sha256; -use socket2::Protocol; -pub type NetworkIdentity = crate::common::config::NetworkIdentity; - -#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)] +#[derive(Debug, Clone, PartialEq)] +#[cfg_attr(feature = "management", derive(serde::Serialize, serde::Deserialize))] pub enum GlobalCtxEvent { TunDeviceReady(String), TunDeviceError(String), @@ -73,6 +58,7 @@ pub enum GlobalCtxEvent { PortForwardAdded(PortForwardConfigPb), + #[cfg(feature = "management")] ConfigPatched(InstanceConfigPatch), ProxyCidrsUpdated(Vec, Vec), // (added, removed) @@ -88,114 +74,6 @@ pub enum GlobalCtxEvent { pub type EventBus = tokio::sync::broadcast::Sender; pub type EventBusSubscriber = tokio::sync::broadcast::Receiver; -/// Source of a trusted public key from OSPF route propagation -#[derive(Debug, Clone, Copy, PartialEq, Eq)] -pub enum TrustedKeySource { - /// Peer node's noise static pubkey - OspfNode, - /// Admin-declared trusted credential pubkey - OspfCredential, -} - -/// Metadata for a trusted public key -#[derive(Debug, Clone)] -pub struct TrustedKeyMetadata { - pub source: TrustedKeySource, - /// Expiry time in Unix seconds. None means never expires. - pub expiry_unix: Option, -} - -impl TrustedKeyMetadata { - pub fn is_expired(&self) -> bool { - if let Some(expiry) = self.expiry_unix { - let now = SystemTime::now() - .duration_since(UNIX_EPOCH) - .unwrap() - .as_secs() as i64; - return now >= expiry; - } - false - } -} - -// key is (pubkey, network-name) -pub type TrustedKeyMap = HashMap, TrustedKeyMetadata>; - -struct TrustedKeyMapManager { - network_trusted_keys: DashMap>, -} - -impl TrustedKeyMapManager { - pub fn new() -> Self { - Self { - network_trusted_keys: DashMap::new(), - } - } - - pub fn update_trusted_keys(&self, network_name: &str, trusted_keys: TrustedKeyMap) { - match self.network_trusted_keys.entry(network_name.to_string()) { - dashmap::Entry::Vacant(entry) => { - entry.insert(ArcSwap::new(Arc::new(trusted_keys))); - } - dashmap::Entry::Occupied(entry) => { - entry.get().store(Arc::new(trusted_keys)); - } - } - } - - pub fn remove_trusted_keys(&self, network_name: &str) { - self.network_trusted_keys.remove(network_name); - shrink_dashmap(&self.network_trusted_keys, None); - } - - pub fn verify_trusted_key(&self, pubkey: &[u8], network_name: &str) -> bool { - self.verify_trusted_key_with_source(pubkey, network_name, None) - } - - pub fn verify_trusted_key_with_source( - &self, - pubkey: &[u8], - network_name: &str, - source: Option, - ) -> bool { - let Some(trusted_keys) = self - .network_trusted_keys - .get(network_name) - .map(|v| v.load_full()) - else { - return false; - }; - - let Some(metadata) = trusted_keys.get(&pubkey.to_vec()) else { - return false; - }; - - if let Some(source) = source { - metadata.source == source && !metadata.is_expired() - } else { - !metadata.is_expired() - } - } - - pub fn list_trusted_keys(&self, network_name: &str) -> Vec<(Vec, TrustedKeyMetadata)> { - let Some(trusted_keys) = self - .network_trusted_keys - .get(network_name) - .map(|v| v.load_full()) - else { - return Vec::new(); - }; - - let mut items = trusted_keys - .iter() - .filter(|(_, metadata)| !metadata.is_expired()) - .map(|(pubkey, metadata)| (pubkey.clone(), metadata.clone())) - .collect::>(); - items.sort_by(|left, right| left.0.cmp(&right.0)); - items - } -} - pub struct GlobalCtx { pub inst_name: String, pub id: uuid::Uuid, @@ -207,41 +85,11 @@ pub struct GlobalCtx { cached_ipv4: AtomicCell>, cached_ipv6: AtomicCell>, - public_ipv6_lease: AtomicCell>, - public_ipv6_routes: Mutex>, - cached_proxy_cidrs: AtomicCell>>, - - ip_collector: Mutex>>, - hostname: Mutex, - stun_info_collection: Mutex>, - - running_listeners: Mutex>, - advertised_ipv6_public_addr_prefix: Mutex>, tun_device_name: Mutex>, flags: ArcSwap, - - // Runtime/base advertised feature flags before config-owned fields are - // overlaid by set_flags. Keep this separate so config patches do not erase - // runtime state such as public-server role, IPv6 provider status, or the - // non-whitelist avoid-relay preference. - base_feature_flags: AtomicCell, - - feature_flags: AtomicCell, - - token_bucket_manager: TokenBucketManager, - - stats_manager: Arc, - - acl_filter: Arc, - - credential_manager: Arc, - - /// OSPF propagated trusted keys (peer pubkeys and admin credentials) - /// Stored in ArcSwap for lock-free reads and atomic batch updates - trusted_keys: Arc, } impl std::fmt::Debug for GlobalCtx { @@ -258,60 +106,52 @@ impl std::fmt::Debug for GlobalCtx { pub type ArcGlobalCtx = std::sync::Arc; +#[async_trait] +impl PublicIpv6Host for GlobalCtx { + async fn collect_reserved_public_ipv6_addrs( + &self, + prefix: cidr::Ipv6Cidr, + ) -> HashSet { + let context = SocketContext::default() + .with_socket_mark(self.config.get_flags().socket_mark) + .with_netns(self.net_ns.name().map(NetNamespace::new)); + let ip_list = crate::host_runtime::native_host_runtime() + .collect_ip_addrs(&context) + .await; + let mut reserved = HashSet::new(); + reserved.extend( + ip_list + .interface_ipv6s + .into_iter() + .map(Ipv6Addr::from) + .filter(|addr| prefix.contains(addr)), + ); + reserved.extend( + ip_list + .public_ipv6 + .into_iter() + .map(Ipv6Addr::from) + .filter(|addr| prefix.contains(addr)), + ); + reserved + } +} + impl GlobalCtx { - fn apply_disable_relay_data_flag( - flags: &Flags, - mut feature_flags: PeerFeatureFlag, - ) -> PeerFeatureFlag { - if flags.disable_relay_data { - feature_flags.avoid_relay_data = true; - } - feature_flags - } - - fn derive_feature_flags(flags: &Flags, mut feature_flags: PeerFeatureFlag) -> PeerFeatureFlag { - feature_flags.kcp_input = !flags.disable_kcp_input; - feature_flags.no_relay_kcp = flags.disable_relay_kcp; - feature_flags.support_conn_list_sync = true; - feature_flags.quic_input = !flags.disable_quic_input; - feature_flags.no_relay_quic = flags.disable_relay_quic; - feature_flags.need_p2p = flags.need_p2p; - feature_flags.disable_p2p = flags.disable_p2p; - Self::apply_disable_relay_data_flag(flags, feature_flags) - } - pub fn new(config_fs: impl ConfigLoader + 'static) -> Self { let id = config_fs.get_id(); let network = config_fs.get_network_identity(); let net_ns = NetNS::new(config_fs.get_netns()); - let hostname = config_fs.get_hostname(); + let hostname = match config_fs.get_hostname() { + hostname if !hostname.is_empty() => hostname, + _ => gethostname::gethostname().to_string_lossy().to_string(), + }; + let flags = config_fs.get_flags(); + if flags.enable_encryption && effective_encryption_uses_xor(&flags.encryption_algorithm) { + tracing::warn!("using insecure XOR because no AEAD encryption is configured"); + } let (event_bus, _) = tokio::sync::broadcast::channel(16); - - let stun_info_collector = StunInfoCollector::new_with_default_servers(); - - if let Some(stun_servers) = config_fs.get_stun_servers() { - stun_info_collector.set_stun_servers(stun_servers); - } else { - stun_info_collector.set_stun_servers(StunInfoCollector::get_default_servers()); - } - - if let Some(stun_servers) = config_fs.get_stun_servers_v6() { - stun_info_collector.set_stun_servers_v6(stun_servers); - } else { - stun_info_collector.set_stun_servers_v6(StunInfoCollector::get_default_servers_v6()); - } - - let stun_info_collector = Arc::new(stun_info_collector); - - let flags = config_fs.get_flags(); - - let base_feature_flags = PeerFeatureFlag::default(); - let feature_flags = Self::derive_feature_flags(&flags, base_feature_flags); - - let credential_storage_path = config_fs.get_credential_file(); - let credential_manager = Arc::new(CredentialManager::new(credential_storage_path)); - GlobalCtx { inst_name: config_fs.get_inst_name(), id, @@ -322,38 +162,11 @@ impl GlobalCtx { event_bus, cached_ipv4: AtomicCell::new(None), cached_ipv6: AtomicCell::new(None), - public_ipv6_lease: AtomicCell::new(None), - public_ipv6_routes: Mutex::new(BTreeSet::new()), - cached_proxy_cidrs: AtomicCell::new(None), - - ip_collector: Mutex::new(Some(Arc::new(IPCollector::new( - net_ns, - stun_info_collector.clone(), - )))), - hostname: Mutex::new(hostname), - stun_info_collection: Mutex::new(stun_info_collector), - - running_listeners: Mutex::new(Vec::new()), - advertised_ipv6_public_addr_prefix: Mutex::new(None), tun_device_name: Mutex::new(None), flags: ArcSwap::new(Arc::new(flags)), - - base_feature_flags: AtomicCell::new(base_feature_flags), - - feature_flags: AtomicCell::new(feature_flags), - - token_bucket_manager: TokenBucketManager::new(), - - stats_manager: Arc::new(StatsManager::new()), - - acl_filter: Arc::new(AclFilter::new()), - - credential_manager, - - trusted_keys: Arc::new(TrustedKeyMapManager::new()), } } @@ -372,15 +185,18 @@ impl GlobalCtx { } } + #[cfg(any(feature = "tun", test))] fn set_tun_device_name(&self, name: Option) { *self.tun_device_name.lock().unwrap() = name; } + #[cfg(any(feature = "tun", test))] pub(crate) fn set_tun_device_ready(&self, name: String) { self.set_tun_device_name(Some(name.clone())); self.issue_event(GlobalCtxEvent::TunDeviceReady(name)); } + #[cfg(any(feature = "tun", test))] pub(crate) fn set_tun_device_error(&self, error: String) { self.set_tun_device_name(None); self.issue_event(GlobalCtxEvent::TunDeviceError(error)); @@ -390,20 +206,6 @@ impl GlobalCtx { self.tun_device_name.lock().unwrap().clone() } - pub fn check_network_in_whitelist(&self, network_name: &str) -> Result<(), anyhow::Error> { - if self - .get_flags() - .relay_network_whitelist - .split(" ") - .map(wildmatch::WildMatch::new) - .any(|wl| wl.matches(network_name)) - { - Ok(()) - } else { - Err(anyhow::anyhow!("network {} not in whitelist", network_name)) - } - } - pub fn get_ipv4(&self) -> Option { if let Some(ret) = self.cached_ipv4.load() { return Some(ret); @@ -432,43 +234,8 @@ impl GlobalCtx { self.cached_ipv6.store(None); } - pub fn get_public_ipv6_lease(&self) -> Option { - self.public_ipv6_lease.load() - } - - pub fn set_public_ipv6_lease(&self, addr: Option) { - self.public_ipv6_lease.store(addr); - } - - pub fn set_public_ipv6_routes(&self, routes: BTreeSet) { - *self.public_ipv6_routes.lock().unwrap() = - routes.into_iter().map(|route| route.address()).collect(); - } - pub fn is_ip_local_ipv6(&self, ip: &std::net::Ipv6Addr) -> bool { self.get_ipv6().map(|x| x.address() == *ip).unwrap_or(false) - || self - .get_public_ipv6_lease() - .map(|x| x.address() == *ip) - .unwrap_or(false) - } - - pub fn is_ip_easytier_managed_ipv6(&self, ip: &std::net::Ipv6Addr) -> bool { - self.is_ip_local_ipv6(ip) || self.public_ipv6_routes.lock().unwrap().contains(ip) - } - - pub fn get_advertised_ipv6_public_addr_prefix(&self) -> Option { - *self.advertised_ipv6_public_addr_prefix.lock().unwrap() - } - - pub fn set_advertised_ipv6_public_addr_prefix(&self, prefix: Option) -> bool { - let mut guard = self.advertised_ipv6_public_addr_prefix.lock().unwrap(); - if *guard == prefix { - return false; - } - - *guard = prefix; - true } pub fn get_id(&self) -> uuid::Uuid { @@ -493,23 +260,10 @@ impl GlobalCtx { self.config.get_network_identity() } - pub fn get_secret_proof(&self, challenge: &[u8]) -> Option> { - let network_secret = self.get_network_identity().network_secret?; - let key = network_secret.as_bytes(); - let mut mac = Hmac::::new_from_slice(key).unwrap(); - mac.update(b"easytier secret proof"); - mac.update(challenge); - Some(mac) - } - pub fn get_network_name(&self) -> String { self.get_network_identity().network_name } - pub fn get_ip_collector(&self) -> Arc { - self.ip_collector.lock().unwrap().as_ref().unwrap().clone() - } - pub fn get_hostname(&self) -> String { return self.hostname.lock().unwrap().clone(); } @@ -518,32 +272,6 @@ impl GlobalCtx { *self.hostname.lock().unwrap() = hostname; } - pub fn get_stun_info_collector(&self) -> Arc { - self.stun_info_collection.lock().unwrap().clone() - } - - pub fn replace_stun_info_collector(&self, collector: Box) { - let arc_collector: Arc = Arc::new(collector); - *self.stun_info_collection.lock().unwrap() = arc_collector.clone(); - - // rebuild the ip collector - *self.ip_collector.lock().unwrap() = Some(Arc::new(IPCollector::new( - self.net_ns.clone(), - arc_collector, - ))); - } - - pub fn get_running_listeners(&self) -> Vec { - self.running_listeners.lock().unwrap().clone() - } - - pub fn add_running_listener(&self, url: url::Url) { - let mut l = self.running_listeners.lock().unwrap(); - if !l.contains(&url) { - l.push(url); - } - } - pub fn get_vpn_portal_cidr(&self) -> Option { self.config.get_vpn_portal_config().map(|x| x.client_cidr) } @@ -554,10 +282,6 @@ impl GlobalCtx { pub fn set_flags(&self, flags: Flags) { self.config.set_flags(flags.clone()); - self.feature_flags.store(Self::derive_feature_flags( - &flags, - self.base_feature_flags.load(), - )); self.flags.store(Arc::new(flags)); } @@ -565,46 +289,6 @@ impl GlobalCtx { self.flags.load_full() } - pub fn get_128_key(&self) -> [u8; 16] { - let mut key = [0u8; 16]; - let secret = self - .config - .get_network_identity() - .network_secret - .unwrap_or_default(); - // fill key according to network secret - let mut hasher = DefaultHasher::new(); - hasher.write(secret.as_bytes()); - key[0..8].copy_from_slice(&hasher.finish().to_be_bytes()); - hasher.write(&key[0..8]); - key[8..16].copy_from_slice(&hasher.finish().to_be_bytes()); - hasher.write(&key[0..16]); - key - } - - pub fn get_256_key(&self) -> [u8; 32] { - let mut key = [0u8; 32]; - let secret = self - .config - .get_network_identity() - .network_secret - .unwrap_or_default(); - // fill key according to network secret - let mut hasher = DefaultHasher::new(); - hasher.write(secret.as_bytes()); - hasher.write(b"easytier-256bit-key"); // 添加固定盐值以区分128位和256位密钥 - - // 生成32字节密钥 - for i in 0..4 { - let chunk_start = i * 8; - let chunk_end = chunk_start + 8; - hasher.write(&key[0..chunk_start]); - hasher.write(&[i as u8]); // 添加索引以确保每个8字节块都不同 - key[chunk_start..chunk_end].copy_from_slice(&hasher.finish().to_be_bytes()); - } - key - } - pub fn enable_exit_node(&self) -> bool { self.flags.load().enable_exit_node || cfg!(target_env = "ohos") } @@ -616,202 +300,11 @@ impl GlobalCtx { pub fn no_tun(&self) -> bool { self.flags.load().no_tun } - - pub fn get_feature_flags(&self) -> PeerFeatureFlag { - self.feature_flags.load() - } - - /// Replace the runtime/base advertised flags as a complete snapshot. - /// - /// This is intended for foreign scoped contexts that inherit an already - /// computed feature-flag snapshot from their parent. Most callers should use - /// a narrower setter so they do not accidentally overwrite unrelated runtime - /// state. - pub fn set_base_advertised_feature_flags(&self, feature_flags: PeerFeatureFlag) { - self.base_feature_flags.store(feature_flags); - let flags = self.flags.load(); - self.feature_flags - .store(Self::apply_disable_relay_data_flag( - flags.as_ref(), - feature_flags, - )); - } - - /// Set the avoid-relay preference that is independent of disable_relay_data. - /// - /// disable_relay_data still forces the effective advertised flag to true, - /// but this base preference is preserved when that config flag is toggled. - pub fn set_avoid_relay_data_preference(&self, avoid_relay_data: bool) -> bool { - let mut base_feature_flags = self.base_feature_flags.load(); - base_feature_flags.avoid_relay_data = avoid_relay_data; - self.base_feature_flags.store(base_feature_flags); - - let mut feature_flags = self.feature_flags.load(); - let previous = feature_flags.avoid_relay_data; - feature_flags.avoid_relay_data = avoid_relay_data || self.flags.load().disable_relay_data; - self.feature_flags.store(feature_flags); - previous != feature_flags.avoid_relay_data - } - - /// Set the runtime IPv6-provider advertised bit without touching - /// config-derived feature flags. - pub fn set_ipv6_public_addr_provider_feature_flag(&self, enabled: bool) -> bool { - let mut base_feature_flags = self.base_feature_flags.load(); - base_feature_flags.ipv6_public_addr_provider = enabled; - self.base_feature_flags.store(base_feature_flags); - - let mut feature_flags = self.feature_flags.load(); - if feature_flags.ipv6_public_addr_provider == enabled { - return false; - } - - feature_flags.ipv6_public_addr_provider = enabled; - self.feature_flags.store(feature_flags); - true - } - - pub fn token_bucket_manager(&self) -> &TokenBucketManager { - &self.token_bucket_manager - } - - pub fn stats_manager(&self) -> &Arc { - &self.stats_manager - } - - pub fn get_acl_filter(&self) -> &Arc { - &self.acl_filter - } - - pub fn get_credential_manager(&self) -> &Arc { - &self.credential_manager - } - - /// Check if a public key is trusted using two-level lookup: - /// 1. OSPF propagated trusted_keys (lock-free) - /// 2. Local credential_manager - pub fn is_pubkey_trusted(&self, pubkey: &[u8], network_name: &str) -> bool { - // First level: check OSPF propagated keys (lock-free) - if self.trusted_keys.verify_trusted_key(pubkey, network_name) { - return true; - } - - // Second level: check local credential_manager if in the same network - if network_name == self.get_network_name() { - return self.credential_manager.is_pubkey_trusted(pubkey); - } - - false - } - - pub fn is_pubkey_trusted_with_source( - &self, - pubkey: &[u8], - network_name: &str, - source: TrustedKeySource, - ) -> bool { - self.trusted_keys - .verify_trusted_key_with_source(pubkey, network_name, Some(source)) - } - - /// Atomically replace all OSPF trusted keys with a new set - /// Called by OSPF route layer after each route update - pub fn update_trusted_keys(&self, keys: TrustedKeyMap, network_name: &str) { - self.trusted_keys.update_trusted_keys(network_name, keys); - } - - pub fn remove_trusted_keys(&self, network_name: &str) { - self.trusted_keys.remove_trusted_keys(network_name); - } - - pub fn list_trusted_keys(&self, network_name: &str) -> Vec<(Vec, TrustedKeyMetadata)> { - self.trusted_keys.list_trusted_keys(network_name) - } - - pub fn get_acl_groups(&self, peer_id: PeerId) -> Vec { - use std::collections::HashSet; - self.config - .get_acl() - .and_then(|acl| acl.acl_v1) - .and_then(|acl_v1| acl_v1.group) - .map_or_else(Vec::new, |group| { - let memberships: HashSet<_> = group.members.iter().collect(); - group - .declares - .iter() - .filter(|g| memberships.contains(&g.group_name)) - .map(|g| { - PeerGroupInfo::generate_with_proof( - g.group_name.clone(), - g.group_secret.clone(), - peer_id, - ) - }) - .collect() - }) - } - - pub fn get_acl_group_declarations(&self) -> Vec { - self.config - .get_acl() - .and_then(|acl| acl.acl_v1) - .and_then(|acl_v1| acl_v1.group) - .map_or_else(Vec::new, |group| group.declares.to_vec()) - } - - pub fn p2p_only(&self) -> bool { - self.flags.load().p2p_only - } - - pub fn latency_first(&self) -> bool { - // NOTICE: p2p only is conflict with latency first - let flags = self.flags.load(); - flags.latency_first && !flags.p2p_only - } - - fn is_port_in_running_listeners(&self, port: u16, is_udp: bool) -> bool { - self.running_listeners - .lock() - .unwrap() - .iter() - .any(|x| x.port() == Some(port) && matches_protocol!(x, Protocol::UDP) == is_udp) - } - - #[tracing::instrument(ret, skip(self))] - pub fn should_deny_proxy(&self, dst_addr: &SocketAddr, is_udp: bool) -> bool { - let _g = self.net_ns.guard(); - let ip = dst_addr.ip(); - // first check if ip is an EasyTier-managed local address - // then try bind this ip, if succ means it is local ip - let dst_is_local_et_ip = self.is_ip_local_virtual_ip(&ip); - // this is an expensive operation, should be called sparingly - // 1. tcp/kcp/quic call this only after proxy conn is established - // 2. udp cache the result in nat entry - let dst_is_local_phy_ip = std::net::UdpSocket::bind(format!("{}:0", ip)).is_ok(); - - tracing::trace!( - "check should_deny_proxy: dst_addr={}, dst_is_local_et_ip={}, dst_is_local_phy_ip={}, is_udp={}", - dst_addr, - dst_is_local_et_ip, - dst_is_local_phy_ip, - is_udp - ); - - if dst_is_local_et_ip || dst_is_local_phy_ip { - // if is local ip, make sure the port is not one of the listening ports - self.is_port_in_running_listeners(dst_addr.port(), is_udp) - || (!is_udp && protected_port::is_protected_tcp_port(dst_addr.port())) - } else { - false - } - } } #[cfg(test)] pub mod tests { - use crate::{ - common::{config::TomlConfigLoader, new_peer_id, stun::MockStunInfoCollector}, - proto::common::NatType, - }; + use crate::common::config::TomlConfigLoader; use super::*; @@ -821,7 +314,7 @@ pub mod tests { let global_ctx = GlobalCtx::new(config); let mut subscriber = global_ctx.subscribe(); - let peer_id = new_peer_id(); + let peer_id = rand::random(); global_ctx.issue_event(GlobalCtxEvent::PeerAdded(peer_id)); global_ctx.issue_event(GlobalCtxEvent::PeerRemoved(peer_id)); global_ctx.issue_event(GlobalCtxEvent::PeerConnAdded(PeerConnInfo::default())); @@ -875,195 +368,13 @@ pub mod tests { ); } - #[tokio::test] - async fn trusted_key_source_lookup_is_precise() { + #[test] + fn host_hostname_fallback_does_not_materialize_in_toml() { let config = TomlConfigLoader::default(); - let global_ctx = GlobalCtx::new(config); - let network_name = "net1"; - let pubkey = vec![1; 32]; + let global_ctx = GlobalCtx::new(config.clone()); - global_ctx.update_trusted_keys( - HashMap::from([( - pubkey.clone(), - TrustedKeyMetadata { - source: TrustedKeySource::OspfCredential, - expiry_unix: None, - }, - )]), - network_name, - ); - - assert!(global_ctx.is_pubkey_trusted(&pubkey, network_name)); - assert!(!global_ctx.is_pubkey_trusted_with_source( - &pubkey, - network_name, - TrustedKeySource::OspfNode, - )); - assert!(global_ctx.is_pubkey_trusted_with_source( - &pubkey, - network_name, - TrustedKeySource::OspfCredential, - )); - } - - #[tokio::test] - async fn set_flags_keeps_derived_feature_flags_in_sync() { - let config = TomlConfigLoader::default(); - let global_ctx = GlobalCtx::new(config); - - let mut feature_flags = global_ctx.get_feature_flags(); - feature_flags.avoid_relay_data = true; - feature_flags.is_public_server = true; - global_ctx.set_base_advertised_feature_flags(feature_flags); - - let mut flags = global_ctx.get_flags().clone(); - flags.disable_kcp_input = true; - flags.disable_relay_kcp = true; - flags.disable_quic_input = true; - flags.disable_relay_quic = true; - flags.need_p2p = true; - flags.disable_p2p = true; - global_ctx.set_flags(flags); - - let feature_flags = global_ctx.get_feature_flags(); - assert!(!feature_flags.kcp_input); - assert!(feature_flags.no_relay_kcp); - assert!(!feature_flags.quic_input); - assert!(feature_flags.no_relay_quic); - assert!(feature_flags.need_p2p); - assert!(feature_flags.disable_p2p); - assert!(feature_flags.support_conn_list_sync); - assert!(feature_flags.avoid_relay_data); - assert!(feature_flags.is_public_server); - assert!(!feature_flags.ipv6_public_addr_provider); - } - - #[tokio::test] - async fn set_base_advertised_feature_flags_applies_current_values() { - let config = TomlConfigLoader::default(); - let global_ctx = GlobalCtx::new(config); - - let feature_flags = PeerFeatureFlag { - kcp_input: false, - no_relay_kcp: true, - quic_input: false, - no_relay_quic: true, - is_public_server: true, - ..Default::default() - }; - global_ctx.set_base_advertised_feature_flags(feature_flags); - - assert_eq!(global_ctx.get_feature_flags(), feature_flags); - } - - #[tokio::test] - async fn set_base_advertised_feature_flags_keeps_disable_relay_data_effective() { - let config = TomlConfigLoader::default(); - let global_ctx = GlobalCtx::new(config); - - let mut flags = global_ctx.get_flags().clone(); - flags.disable_relay_data = true; - global_ctx.set_flags(flags); - - let mut feature_flags = global_ctx.get_feature_flags(); - feature_flags.avoid_relay_data = false; - feature_flags.is_public_server = true; - global_ctx.set_base_advertised_feature_flags(feature_flags); - - let advertised_feature_flags = global_ctx.get_feature_flags(); - assert!(advertised_feature_flags.avoid_relay_data); - assert!(advertised_feature_flags.is_public_server); - - let mut flags = global_ctx.get_flags().clone(); - flags.disable_relay_data = false; - global_ctx.set_flags(flags); - - let advertised_feature_flags = global_ctx.get_feature_flags(); - assert!(!advertised_feature_flags.avoid_relay_data); - assert!(advertised_feature_flags.is_public_server); - } - - #[tokio::test] - async fn disable_relay_data_sets_avoid_relay_feature_flag() { - let config = TomlConfigLoader::default(); - let global_ctx = GlobalCtx::new(config); - - let mut flags = global_ctx.get_flags().clone(); - flags.disable_relay_data = true; - global_ctx.set_flags(flags); - - assert!(global_ctx.get_feature_flags().avoid_relay_data); - - let mut flags = global_ctx.get_flags().clone(); - flags.disable_relay_data = false; - global_ctx.set_flags(flags); - - assert!(!global_ctx.get_feature_flags().avoid_relay_data); - - global_ctx.set_avoid_relay_data_preference(true); - - let mut flags = global_ctx.get_flags().clone(); - flags.disable_relay_data = true; - global_ctx.set_flags(flags); - - assert!(global_ctx.get_feature_flags().avoid_relay_data); - - let mut flags = global_ctx.get_flags().clone(); - flags.disable_relay_data = false; - global_ctx.set_flags(flags); - - assert!(global_ctx.get_feature_flags().avoid_relay_data); - } - - #[tokio::test] - async fn should_deny_proxy_for_process_wide_rpc_port() { - protected_port::clear_protected_tcp_ports_for_test(); - protected_port::register_protected_tcp_port(15888); - - let config = TomlConfigLoader::default(); - let global_ctx = GlobalCtx::new(config); - let rpc_addr = SocketAddr::from(([127, 0, 0, 1], 15888)); - let other_tcp_addr = SocketAddr::from(([127, 0, 0, 1], 15889)); - - assert!(global_ctx.should_deny_proxy(&rpc_addr, false)); - assert!(!global_ctx.should_deny_proxy(&rpc_addr, true)); - assert!(!global_ctx.should_deny_proxy(&other_tcp_addr, false)); - - protected_port::clear_protected_tcp_ports_for_test(); - } - - #[tokio::test] - async fn virtual_ipv6_and_public_ipv6_lease_are_stored_separately() { - let config = TomlConfigLoader::default(); - let global_ctx = GlobalCtx::new(config); - let virtual_ipv6 = "fd00::1/64".parse().unwrap(); - let public_ipv6 = "2001:db8::2/64".parse().unwrap(); - - global_ctx.set_ipv6(Some(virtual_ipv6)); - global_ctx.set_public_ipv6_lease(Some(public_ipv6)); - - assert_eq!(global_ctx.get_ipv6(), Some(virtual_ipv6)); - assert_eq!(global_ctx.get_public_ipv6_lease(), Some(public_ipv6)); - } - - #[tokio::test] - async fn public_ipv6_lease_is_treated_as_local_ip() { - protected_port::clear_protected_tcp_ports_for_test(); - - let config = TomlConfigLoader::default(); - let global_ctx = GlobalCtx::new(config); - let public_ipv6 = "2001:db8::2/64".parse().unwrap(); - let listener: url::Url = "tcp://[2001:db8::2]:11010".parse().unwrap(); - global_ctx.set_public_ipv6_lease(Some(public_ipv6)); - global_ctx.add_running_listener(listener); - - let ip = std::net::IpAddr::V6(public_ipv6.address()); - let socket = SocketAddr::from((public_ipv6.address(), 11010)); - - assert!(global_ctx.is_ip_local_virtual_ip(&ip)); - assert!(global_ctx.should_deny_proxy(&socket, false)); - - protected_port::clear_protected_tcp_ports_for_test(); + assert!(!global_ctx.get_hostname().is_empty()); + assert!(!config.dump().contains("hostname")); } pub fn get_mock_global_ctx_with_network( @@ -1073,11 +384,7 @@ pub mod tests { config_fs.set_inst_name(format!("test_{}", config_fs.get_id())); config_fs.set_network_identity(network_identy.unwrap_or_default()); - let ctx = Arc::new(GlobalCtx::new(config_fs)); - ctx.replace_stun_info_collector(Box::new(MockStunInfoCollector { - udp_nat_type: NatType::Unknown, - })); - ctx + Arc::new(GlobalCtx::new(config_fs)) } pub fn get_mock_global_ctx() -> ArcGlobalCtx { diff --git a/easytier/src/common/idn.rs b/easytier/src/common/idn.rs deleted file mode 100644 index 4c3a234a..00000000 --- a/easytier/src/common/idn.rs +++ /dev/null @@ -1,70 +0,0 @@ -use idna::domain_to_ascii; -use percent_encoding::percent_decode_str; - -pub fn convert_idn_to_ascii(mut url: url::Url) -> anyhow::Result { - if url.is_special() { - return Ok(url); - } - if let Some(domain) = url.domain() { - let domain = percent_decode_str(domain).decode_utf8()?; - let domain = domain_to_ascii(&domain)?; - url.set_host(Some(&domain))?; - } - Ok(url) -} - -#[cfg(test)] -mod tests { - use super::*; - use rstest::rstest; - - #[rstest] - // test_ascii_only_urls - #[case("example.com", "example.com")] - #[case("test.org:8080/path", "test.org:8080/path")] - // test_unicode_domains - #[case("räksmörgås.nu", "xn--rksmrgs-5wao1o.nu")] - #[case("中文.测试", "xn--fiq228c.xn--0zwm56d")] - // test_unicode_domains_with_port - #[case("räksmörgås.nu:8080", "xn--rksmrgs-5wao1o.nu:8080")] - // test_unicode_domains_with_port_and_path - #[case("例子.测试/path", "xn--fsqu00a.xn--0zwm56d/path")] - #[case("中文.测试:9000/api", "xn--fiq228c.xn--0zwm56d:9000/api")] - #[case("räksmörgås.nu:8080/path", "xn--rksmrgs-5wao1o.nu:8080/path")] - // test_unicode_domains_with_port_and_unicode_path - #[case( - "中文.测试:8000/用户/管理", - "xn--fiq228c.xn--0zwm56d:8000/%E7%94%A8%E6%88%B7/%E7%AE%A1%E7%90%86" - )] - // test_ipv6_literals & test_ipv6_with_unicode_path - #[case("[2001:db8::1]:8080", "[2001:db8::1]:8080")] - #[case("[2001:db8::1]/path", "[2001:db8::1]/path")] - #[case( - "[2001:db8::1]/路径/资源", - "[2001:db8::1]/%E8%B7%AF%E5%BE%84/%E8%B5%84%E6%BA%90" - )] - fn test_convert_idn_to_ascii_cases( - #[case] host_part: &str, - #[case] expected_host_part: &str, - #[values("tcp", "udp", "ws", "wss", "wg", "quic", "http", "https")] protocol: &str, - #[values(false, true)] dual_convert: bool, - ) { - let input = url::Url::parse(&format!("{}://{}", protocol, host_part)).unwrap(); - let input = if dual_convert { - // in case url is serialized/deserialized as string somewhere else - input.to_string().parse().unwrap() - } else { - input - }; - let actual = convert_idn_to_ascii(input.clone()).unwrap().to_string(); - - let mut expected = format!("{}://{}", protocol, expected_host_part); - - // ws and wss protocols may automatically add a trailing slash if there's no path after host/port - if input.is_special() && actual.ends_with("/") && !expected_host_part.ends_with("/") { - expected.push('/'); - } - - assert_eq!(actual, expected); - } -} diff --git a/easytier/src/common/ifcfg/mod.rs b/easytier/src/common/ifcfg/mod.rs index 6139e733..0a64a33c 100644 --- a/easytier/src/common/ifcfg/mod.rs +++ b/easytier/src/common/ifcfg/mod.rs @@ -1,21 +1,30 @@ +#![cfg_attr( + not(any(feature = "public-ipv6-provider", feature = "tun")), + allow(dead_code) +)] + #[cfg(any( all(target_os = "macos", not(feature = "macos-ne")), target_os = "freebsd" ))] mod darwin; -#[cfg(target_os = "linux")] +#[cfg(all(target_os = "linux", feature = "linux-netlink"))] mod netlink; +#[cfg(all(target_os = "linux", feature = "linux-netlink"))] +mod netlink_wire; #[cfg(target_os = "windows")] mod win; #[cfg(target_os = "windows")] mod windows; -mod route; - use std::net::{Ipv4Addr, Ipv6Addr}; use async_trait::async_trait; use cidr::{Ipv4Inet, Ipv6Inet}; +#[cfg(any( + all(target_os = "macos", not(feature = "macos-ne")), + target_os = "freebsd" +))] use tokio::process::Command; use super::error::Error; @@ -89,6 +98,10 @@ pub trait IfConfiguerTrait: Send + Sync { } } +#[cfg(any( + all(target_os = "macos", not(feature = "macos-ne")), + target_os = "freebsd" +))] fn cidr_to_subnet_mask(prefix_length: u8) -> Ipv4Addr { if prefix_length > 32 { panic!("Invalid CIDR prefix length"); @@ -105,6 +118,10 @@ fn cidr_to_subnet_mask(prefix_length: u8) -> Ipv4Addr { ) } +#[cfg(any( + all(target_os = "macos", not(feature = "macos-ne")), + target_os = "freebsd" +))] async fn run_shell_cmd(cmd: &str) -> Result<(), Error> { let cmd_out: std::process::Output; let stdout: String; @@ -144,8 +161,10 @@ pub struct DummyIfConfiger {} #[async_trait] impl IfConfiguerTrait for DummyIfConfiger {} -#[cfg(target_os = "linux")] +#[cfg(all(target_os = "linux", feature = "linux-netlink"))] pub type IfConfiger = netlink::NetlinkIfConfiger; +#[cfg(all(target_os = "linux", not(feature = "linux-netlink")))] +pub type IfConfiger = DummyIfConfiger; #[cfg(any( all(target_os = "macos", not(feature = "macos-ne")), @@ -167,28 +186,36 @@ pub type IfConfiger = DummyIfConfiger; #[cfg(target_os = "windows")] pub use windows::RegistryManager; -#[cfg(target_os = "linux")] -pub(crate) fn list_ipv6_route_messages() --> Result, Error> { +#[cfg(all(target_os = "linux", feature = "linux-netlink"))] +pub(crate) use netlink_wire::RouteMessage; +#[cfg(all( + target_os = "linux", + feature = "linux-netlink", + feature = "public-ipv6-provider" +))] +pub(crate) use netlink_wire::RouteType; + +#[cfg(all(target_os = "linux", feature = "linux-netlink"))] +pub(crate) fn list_ipv6_route_messages() -> Result, Error> { netlink::NetlinkIfConfiger::list_ipv6_route_messages() } -#[cfg(target_os = "linux")] +#[cfg(all(target_os = "linux", feature = "linux-netlink"))] pub(crate) fn get_interface_index(name: &str) -> Result { netlink::NetlinkIfConfiger::get_interface_index(name) } -#[cfg(target_os = "linux")] +#[cfg(all(target_os = "linux", feature = "linux-netlink"))] pub(crate) fn add_ipv6_ndp_proxy(name: &str, address: Ipv6Addr) -> Result<(), Error> { netlink::NetlinkIfConfiger::add_ipv6_ndp_proxy(name, address) } -#[cfg(target_os = "linux")] +#[cfg(all(target_os = "linux", feature = "linux-netlink"))] pub(crate) fn remove_ipv6_ndp_proxy(name: &str, address: Ipv6Addr) -> Result<(), Error> { netlink::NetlinkIfConfiger::remove_ipv6_ndp_proxy(name, address) } -#[cfg(target_os = "linux")] +#[cfg(all(target_os = "linux", feature = "linux-netlink"))] pub(crate) fn list_ipv6_ndp_proxy( name: &str, ) -> Result, Error> { diff --git a/easytier/src/common/ifcfg/netlink.rs b/easytier/src/common/ifcfg/netlink.rs index 63cdd916..5cd8f728 100644 --- a/easytier/src/common/ifcfg/netlink.rs +++ b/easytier/src/common/ifcfg/netlink.rs @@ -1,41 +1,40 @@ +#![cfg_attr( + not(any(feature = "public-ipv6-provider", feature = "tun")), + allow(dead_code) +)] + use std::{ collections::BTreeSet, ffi::CString, fmt::Debug, net::{IpAddr, Ipv4Addr, Ipv6Addr}, - num::NonZero, os::fd::AsRawFd, }; use anyhow::Context; use async_trait::async_trait; use cidr::{IpInet, Ipv4Inet, Ipv6Inet}; -use netlink_packet_core::{ - NLM_F_ACK, NLM_F_CREATE, NLM_F_DUMP, NLM_F_EXCL, NLM_F_REQUEST, NetlinkDeserializable, - NetlinkHeader, NetlinkMessage, NetlinkPayload, NetlinkSerializable, -}; -use netlink_packet_route::{ - AddressFamily, RouteNetlinkMessage, - address::{AddressAttribute, AddressMessage}, - neighbour::{ - NeighbourAddress, NeighbourAttribute, NeighbourFlags, NeighbourHeader, NeighbourMessage, - NeighbourState, - }, - route::{ - RouteAddress, RouteAttribute, RouteHeader, RouteMessage, RouteProtocol, RouteScope, - RouteType, - }, -}; use netlink_sys::{Socket, SocketAddr, protocols::NETLINK_ROUTE}; +#[cfg(test)] +use nix::libc::SIOCGIFMTU; use nix::{ ifaddrs::getifaddrs, - libc::{self, Ioctl, SIOCGIFFLAGS, SIOCGIFMTU, SIOCSIFFLAGS, SIOCSIFMTU, ifreq, ioctl}, + libc::{self, Ioctl, SIOCGIFFLAGS, SIOCSIFFLAGS, SIOCSIFMTU, ifreq, ioctl}, net::if_::InterfaceFlags, sys::socket::SockaddrLike as _, }; use pnet::ipnetwork::ip_mask_to_prefix; -use super::{Error, IfConfiguerTrait, route::Route}; +use super::{ + Error, IfConfiguerTrait, + netlink_wire::{ + AddressMessage, MessageBuilder, MessageIter, NLM_F_ACK, NLM_F_CREATE, NLM_F_DUMP, + NLM_F_DUMP_INTR, NLM_F_EXCL, NLM_F_REQUEST, NLMSG_DONE, NLMSG_ERROR, NeighborMessage, + NetlinkDecode, NetlinkEncode, RTM_DELADDR, RTM_DELNEIGH, RTM_DELROUTE, RTM_GETNEIGH, + RTM_GETROUTE, RTM_NEWADDR, RTM_NEWNEIGH, RTM_NEWROUTE, RouteMessage, RouteMessageBuilder, + RouteType, netlink_error_code, + }, +}; pub(crate) fn dummy_socket() -> Result { Ok(std::net::UdpSocket::bind("0:0")?) @@ -51,115 +50,88 @@ fn build_ifreq(name: &str) -> ifreq { ifr } -fn send_netlink_req( - req: T, - flags: u16, -) -> Result { +fn send_netlink_req(builder: MessageBuilder) -> Result { let mut socket = Socket::new(NETLINK_ROUTE)?; socket.bind_auto()?; socket.connect(&SocketAddr::new(0, 0))?; - let mut req: NetlinkMessage = - NetlinkMessage::new(NetlinkHeader::default(), NetlinkPayload::InnerMessage(req)); - req.header.flags = flags; - - req.finalize(); - let mut buf = vec![0; req.header.length as _]; - req.serialize(&mut buf); - - tracing::debug!("net link request >>> {:?}", req); + let buf = builder.finish()?; + tracing::debug!(request_len = buf.len(), "sending netlink request"); socket.send(&buf, 0)?; Ok(socket) } -fn send_netlink_req_and_wait_one_resp( - req: T, - is_remove: bool, -) -> Result<(), Error> { - let socket = send_netlink_req( - req, - NLM_F_ACK | NLM_F_CREATE | NLM_F_REQUEST | if !is_remove { NLM_F_EXCL } else { 0 }, - )?; - let resp = socket.recv_from_full()?; - let ret = NetlinkMessage::::deserialize(&resp.0) - .with_context(|| "Failed to deserialize netlink message")?; - - tracing::debug!("net link response <<< {:?}", ret); - - match ret.payload { - NetlinkPayload::Error(e) => { - if e.code == NonZero::new(0) { - Ok(()) - } else { - Err(e.to_io().into()) +fn send_netlink_req_and_wait_ack(builder: MessageBuilder) -> Result<(), Error> { + let socket = send_netlink_req(builder)?; + loop { + let (response, _) = socket.recv_from_full()?; + for frame in MessageIter::new(&response) { + let (header, payload) = frame?; + if header.message_type == NLMSG_ERROR { + return match netlink_error_code(payload)? { + 0 => Ok(()), + errno => Err(std::io::Error::from_raw_os_error(errno.abs()).into()), + }; + } + if header.message_type == NLMSG_DONE { + return Ok(()); } - } - p => { - tracing::error!("Unexpected netlink response: {:?}", p); - Err(anyhow::anyhow!("Unexpected netlink response").into()) } } } -fn addr_to_ip(addr: RouteAddress) -> Option { - match addr { - RouteAddress::Inet(addr) => Some(addr.into()), - RouteAddress::Inet6(addr) => Some(addr.into()), - _ => None, +fn message_request( + message_type: u16, + flags: u16, + message: &T, +) -> Result { + let mut builder = MessageBuilder::new(message_type, flags); + let mut message_bytes = Vec::new(); + message.write_to(&mut message_bytes)?; + builder.append_bytes(&message_bytes); + Ok(builder) +} + +fn receive_netlink_dump(builder: MessageBuilder) -> Result, Error> { + let socket = send_netlink_req(builder)?; + let mut messages = Vec::new(); + + loop { + let (response, _) = socket.recv_from_full()?; + for frame in MessageIter::new(&response) { + let (header, payload) = frame?; + if header.flags & NLM_F_DUMP_INTR != 0 { + return Err(std::io::Error::new( + std::io::ErrorKind::Interrupted, + "netlink dump was interrupted", + ) + .into()); + } + if header.message_type == NLMSG_DONE { + return Ok(messages); + } + if header.message_type == NLMSG_ERROR { + let error = netlink_error_code(payload)?; + if error == 0 { + continue; + } + return Err(std::io::Error::from_raw_os_error(error.abs()).into()); + } + if header.message_type == T::MESSAGE_TYPE { + messages.push(T::from_bytes(payload)?); + } + } } } -impl From for Route { - fn from(msg: RouteMessage) -> Self { - let mut gateway = None; - let mut source = None; - let mut source_hint = None; - let mut destination = None; - let mut ifindex = None; - let mut metric = None; - - for attr in msg.attributes { - match attr { - RouteAttribute::Source(addr) => { - source = addr_to_ip(addr); - } - RouteAttribute::PrefSource(addr) => { - source_hint = addr_to_ip(addr); - } - RouteAttribute::Destination(addr) => { - destination = addr_to_ip(addr); - } - RouteAttribute::Gateway(addr) => { - gateway = addr_to_ip(addr); - } - RouteAttribute::Oif(i) => { - ifindex = Some(i); - } - RouteAttribute::Priority(priority) => { - metric = Some(priority); - } - _ => {} - } - } - // rtnetlink gives None instead of 0.0.0.0 for the default route, but we'll convert to 0 here to make it match the other platforms - let destination = destination.unwrap_or_else(|| match msg.header.address_family { - AddressFamily::Inet => Ipv4Addr::UNSPECIFIED.into(), - AddressFamily::Inet6 => Ipv6Addr::UNSPECIFIED.into(), - _ => panic!("invalid destination family"), - }); - Self { - destination, - prefix: msg.header.destination_prefix_length, - source, - source_prefix: msg.header.source_prefix_length, - source_hint, - gateway, - ifindex, - table: msg.header.table, - metric, - } - } +fn dump_netlink_messages( + message_type: u16, + dump_header: &[u8], +) -> Result, Error> { + let mut builder = MessageBuilder::new(message_type, NLM_F_REQUEST | NLM_F_DUMP); + builder.append_bytes(dump_header); + receive_netlink_dump(builder) } pub struct NetlinkIfConfiger {} @@ -184,19 +156,14 @@ impl NetlinkIfConfiger { } fn remove_one_ip(name: &str, ip: Ipv4Addr, prefix_len: u8) -> Result<(), Error> { - let mut message = AddressMessage::default(); - message.header.prefix_len = prefix_len; - message.header.index = NetlinkIfConfiger::get_interface_index(name)?; - message.header.family = AddressFamily::Inet; - - message - .attributes - .push(AddressAttribute::Address(std::net::IpAddr::V4(ip))); - - send_netlink_req_and_wait_one_resp::( - RouteNetlinkMessage::DelAddress(message), - true, - ) + let message = AddressMessage::new( + libc::AF_INET as u8, + Self::get_interface_index(name)?, + prefix_len, + IpAddr::V4(ip), + ); + let request = message_request(RTM_DELADDR, NLM_F_ACK | NLM_F_REQUEST, &message)?; + send_netlink_req_and_wait_ack(request) } fn get_prefix_len_ipv6(name: &str, ip: Ipv6Addr) -> Result { @@ -210,19 +177,14 @@ impl NetlinkIfConfiger { } fn remove_one_ipv6(name: &str, ip: Ipv6Addr, prefix_len: u8) -> Result<(), Error> { - let mut message = AddressMessage::default(); - message.header.prefix_len = prefix_len; - message.header.index = NetlinkIfConfiger::get_interface_index(name)?; - message.header.family = AddressFamily::Inet6; - - message - .attributes - .push(AddressAttribute::Address(std::net::IpAddr::V6(ip))); - - send_netlink_req_and_wait_one_resp::( - RouteNetlinkMessage::DelAddress(message), - true, - ) + let message = AddressMessage::new( + libc::AF_INET6 as u8, + Self::get_interface_index(name)?, + prefix_len, + IpAddr::V6(ip), + ); + let request = message_request(RTM_DELADDR, NLM_F_ACK | NLM_F_REQUEST, &message)?; + send_netlink_req_and_wait_ack(request) } pub(crate) fn mtu_op>( @@ -249,6 +211,7 @@ impl NetlinkIfConfiger { Ok(unsafe { ifr.ifr_ifru.ifru_mtu as u32 }) } + #[cfg(test)] fn mtu(name: &str) -> Result { Self::mtu_op(name, SIOCGIFMTU, 0) } @@ -316,166 +279,60 @@ impl NetlinkIfConfiger { Self::set_flags_op(name, SIOCGIFFLAGS, InterfaceFlags::empty()) } - fn list_route_messages(address_family: AddressFamily) -> Result, Error> { - let mut message = RouteMessage::default(); - - message.header.table = RouteHeader::RT_TABLE_UNSPEC; - message.header.protocol = RouteProtocol::Unspec; - - message.header.scope = RouteScope::Universe; - message.header.kind = RouteType::Unicast; - - message.header.address_family = address_family; - message.header.destination_prefix_length = 0; - message.header.source_prefix_length = 0; - - let s = send_netlink_req( - RouteNetlinkMessage::GetRoute(message), - NLM_F_REQUEST | NLM_F_DUMP, - )?; - - let mut ret_vec = vec![]; - - let mut resp = Vec::::new(); - loop { - if resp.is_empty() { - let (new_resp, _) = s.recv_from_full()?; - resp = new_resp; - } - let ret = NetlinkMessage::::deserialize(&resp) - .with_context(|| "Failed to deserialize netlink message")?; - resp = resp.split_off(ret.buffer_len()); - - tracing::debug!("net link response <<< {:?}", ret); - - match ret.payload { - NetlinkPayload::Error(e) => { - if e.code == NonZero::new(0) { - continue; - } else { - return Err(e.to_io().into()); - } - } - NetlinkPayload::InnerMessage(RouteNetlinkMessage::NewRoute(m)) => { - tracing::debug!("net link response <<< {:?}", m); - ret_vec.push(m); - } - NetlinkPayload::Done(_) => { - break; - } - p => { - tracing::error!("Unexpected netlink response: {:?}", p); - return Err(anyhow::anyhow!("Unexpected netlink response").into()); - } - } - } - - Ok(ret_vec) + fn list_route_messages(address_family: u8) -> Result, Error> { + Ok( + dump_netlink_messages::(RTM_GETROUTE, &RouteMessage::dump_header())? + .into_iter() + .filter(|message| message.family() == address_family) + .collect(), + ) } fn list_routes() -> Result, Error> { - Self::list_route_messages(AddressFamily::Inet) + Self::list_route_messages(libc::AF_INET as u8) } pub(crate) fn list_ipv6_route_messages() -> Result, Error> { - Self::list_route_messages(AddressFamily::Inet6) + Self::list_route_messages(libc::AF_INET6 as u8) } - fn ipv6_ndp_proxy_message(name: &str, address: Ipv6Addr) -> Result { - let mut message = NeighbourMessage::default(); - message.header = NeighbourHeader { - family: AddressFamily::Inet6, - ifindex: Self::get_interface_index(name)?, - state: NeighbourState::Permanent, - flags: NeighbourFlags::Proxy, - kind: RouteType::Unicast, - }; - message - .attributes - .push(NeighbourAttribute::Destination(NeighbourAddress::Inet6( - address, - ))); - Ok(message) + fn ipv6_ndp_proxy_message(name: &str, address: Ipv6Addr) -> Result { + Ok(NeighborMessage::proxy( + Self::get_interface_index(name)?, + address, + )) } pub(crate) fn add_ipv6_ndp_proxy(name: &str, address: Ipv6Addr) -> Result<(), Error> { - send_netlink_req_and_wait_one_resp( - RouteNetlinkMessage::NewNeighbour(Self::ipv6_ndp_proxy_message(name, address)?), - false, - ) + let message = Self::ipv6_ndp_proxy_message(name, address)?; + let request = message_request( + RTM_NEWNEIGH, + NLM_F_ACK | NLM_F_CREATE | NLM_F_EXCL | NLM_F_REQUEST, + &message, + )?; + send_netlink_req_and_wait_ack(request) } pub(crate) fn remove_ipv6_ndp_proxy(name: &str, address: Ipv6Addr) -> Result<(), Error> { - send_netlink_req_and_wait_one_resp( - RouteNetlinkMessage::DelNeighbour(Self::ipv6_ndp_proxy_message(name, address)?), - true, - ) + let message = Self::ipv6_ndp_proxy_message(name, address)?; + let request = message_request(RTM_DELNEIGH, NLM_F_ACK | NLM_F_REQUEST, &message)?; + send_netlink_req_and_wait_ack(request) } - fn list_neighbour_messages( - address_family: AddressFamily, - ) -> Result, Error> { - let mut message = NeighbourMessage::default(); - message.header.family = address_family; - message.header.flags = NeighbourFlags::Proxy; - - let s = send_netlink_req( - RouteNetlinkMessage::GetNeighbour(message), - NLM_F_REQUEST | NLM_F_DUMP, - )?; - - let mut ret_vec = vec![]; - let mut resp = Vec::::new(); - loop { - if resp.is_empty() { - let (new_resp, _) = s.recv_from_full()?; - resp = new_resp; - } - - let ret = NetlinkMessage::::deserialize(&resp) - .with_context(|| "Failed to deserialize netlink neighbour message")?; - resp = resp.split_off(ret.buffer_len()); - - tracing::debug!("net link response <<< {:?}", ret); - - match ret.payload { - NetlinkPayload::Error(e) => { - if e.code == NonZero::new(0) { - continue; - } else { - return Err(e.to_io().into()); - } - } - NetlinkPayload::InnerMessage(RouteNetlinkMessage::NewNeighbour(m)) => { - ret_vec.push(m); - } - NetlinkPayload::Done(_) => { - break; - } - p => { - tracing::error!("Unexpected netlink response: {:?}", p); - return Err(anyhow::anyhow!("Unexpected netlink response").into()); - } - } - } - - Ok(ret_vec) + fn list_neighbour_messages(address_family: u8) -> Result, Error> { + let mut builder = MessageBuilder::new(RTM_GETNEIGH, NLM_F_REQUEST | NLM_F_DUMP); + builder.append_bytes(&NeighborMessage::proxy_dump_header(address_family)); + receive_netlink_dump(builder) } pub(crate) fn list_ipv6_ndp_proxy(name: &str) -> Result, Error> { let ifindex = Self::get_interface_index(name)?; - - Ok(Self::list_neighbour_messages(AddressFamily::Inet6)? + Ok(Self::list_neighbour_messages(libc::AF_INET6 as u8)? .into_iter() - .filter(|message| { - message.header.ifindex == ifindex - && message.header.flags.contains(NeighbourFlags::Proxy) - }) - .filter_map(|message| { - message.attributes.into_iter().find_map(|attr| match attr { - NeighbourAttribute::Destination(NeighbourAddress::Inet6(addr)) => Some(addr), - _ => None, - }) + .filter(|message| message.ifindex() == ifindex && message.is_proxy()) + .filter_map(|message| match message.destination() { + Some(IpAddr::V6(address)) => Some(*address), + _ => None, }) .collect()) } @@ -490,30 +347,21 @@ impl IfConfiguerTrait for NetlinkIfConfiger { cidr_prefix: u8, cost: Option, ) -> Result<(), Error> { - let mut message = RouteMessage::default(); - - message.header.table = RouteHeader::RT_TABLE_MAIN; - message.header.protocol = RouteProtocol::Static; - message.header.scope = RouteScope::Universe; - message.header.kind = RouteType::Unicast; - message.header.address_family = AddressFamily::Inet; - // metric - message - .attributes - .push(RouteAttribute::Priority(cost.unwrap_or(65535) as u32)); - // output interface - message - .attributes - .push(RouteAttribute::Oif(NetlinkIfConfiger::get_interface_index( - name, - )?)); - // source address - message.header.destination_prefix_length = cidr_prefix; - message - .attributes - .push(RouteAttribute::Destination(RouteAddress::Inet(address))); - - send_netlink_req_and_wait_one_resp(RouteNetlinkMessage::NewRoute(message), false) + let message = RouteMessageBuilder::new(libc::AF_INET as u8) + .destination(IpAddr::V4(address), cidr_prefix) + .oif(Self::get_interface_index(name)?) + .priority(cost.unwrap_or(65535) as u32) + .table(libc::RT_TABLE_MAIN.into()) + .static_protocol() + .universe_scope() + .route_type(RouteType::Unicast) + .build(); + let request = message_request( + RTM_NEWROUTE, + NLM_F_ACK | NLM_F_CREATE | NLM_F_EXCL | NLM_F_REQUEST, + &message, + )?; + send_netlink_req_and_wait_ack(request) } async fn remove_ipv4_route( @@ -526,12 +374,16 @@ impl IfConfiguerTrait for NetlinkIfConfiger { let ifidx = NetlinkIfConfiger::get_interface_index(name)?; for msg in routes { - let other_route: Route = msg.clone().into(); - if other_route.destination == std::net::IpAddr::V4(address) - && other_route.prefix == cidr_prefix - && other_route.ifindex == Some(ifidx) + let destination = msg + .destination() + .copied() + .unwrap_or(IpAddr::V4(Ipv4Addr::UNSPECIFIED)); + if destination == IpAddr::V4(address) + && msg.dst_len() == cidr_prefix + && msg.oif() == Some(ifidx) { - send_netlink_req_and_wait_one_resp(RouteNetlinkMessage::DelRoute(msg), true)?; + let request = message_request(RTM_DELROUTE, NLM_F_ACK | NLM_F_REQUEST, &msg)?; + send_netlink_req_and_wait_ack(request)?; return Ok(()); } } @@ -545,37 +397,26 @@ impl IfConfiguerTrait for NetlinkIfConfiger { address: Ipv4Addr, cidr_prefix: u8, ) -> Result<(), Error> { - let mut message = AddressMessage::default(); - - message.header.prefix_len = cidr_prefix; - message.header.index = NetlinkIfConfiger::get_interface_index(name)?; - message.header.family = AddressFamily::Inet; - - message - .attributes - .push(AddressAttribute::Address(std::net::IpAddr::V4(address))); - - // for IPv4 the IFA_LOCAL address can be set to the same value as - // IFA_ADDRESS - message - .attributes - .push(AddressAttribute::Local(std::net::IpAddr::V4(address))); - - // set the IFA_BROADCAST address as well - if cidr_prefix == 32 { - message - .attributes - .push(AddressAttribute::Broadcast(address)); + let broadcast = if cidr_prefix == 32 { + address } else { let ip_addr = u32::from(address); - let brd = Ipv4Addr::from((0xffff_ffff_u32) >> u32::from(cidr_prefix) | ip_addr); - message.attributes.push(AddressAttribute::Broadcast(brd)); + Ipv4Addr::from((0xffff_ffff_u32) >> u32::from(cidr_prefix) | ip_addr) }; - - send_netlink_req_and_wait_one_resp::( - RouteNetlinkMessage::NewAddress(message), - false, + let message = AddressMessage::new( + libc::AF_INET as u8, + Self::get_interface_index(name)?, + cidr_prefix, + IpAddr::V4(address), ) + .local(IpAddr::V4(address)) + .broadcast(broadcast); + let request = message_request( + RTM_NEWADDR, + NLM_F_ACK | NLM_F_CREATE | NLM_F_EXCL | NLM_F_REQUEST, + &message, + )?; + send_netlink_req_and_wait_ack(request) } async fn set_link_status(&self, name: &str, up: bool) -> Result<(), Error> { @@ -613,21 +454,18 @@ impl IfConfiguerTrait for NetlinkIfConfiger { address: std::net::Ipv6Addr, cidr_prefix: u8, ) -> Result<(), Error> { - let mut message = AddressMessage::default(); - - message.header.prefix_len = cidr_prefix; - message.header.index = NetlinkIfConfiger::get_interface_index(name)?; - message.header.family = AddressFamily::Inet6; - - message - .attributes - .push(AddressAttribute::Address(std::net::IpAddr::V6(address))); - - // For IPv6, we don't need IFA_LOCAL or IFA_BROADCAST - send_netlink_req_and_wait_one_resp::( - RouteNetlinkMessage::NewAddress(message), - false, - ) + let message = AddressMessage::new( + libc::AF_INET6 as u8, + Self::get_interface_index(name)?, + cidr_prefix, + IpAddr::V6(address), + ); + let request = message_request( + RTM_NEWADDR, + NLM_F_ACK | NLM_F_CREATE | NLM_F_EXCL | NLM_F_REQUEST, + &message, + )?; + send_netlink_req_and_wait_ack(request) } async fn remove_ipv6(&self, name: &str, ip: Option) -> Result<(), Error> { @@ -654,32 +492,23 @@ impl IfConfiguerTrait for NetlinkIfConfiger { cidr_prefix: u8, cost: Option, ) -> Result<(), Error> { - let mut message = RouteMessage::default(); - - message.header.address_family = AddressFamily::Inet6; - message.header.destination_prefix_length = cidr_prefix; - message.header.table = RouteHeader::RT_TABLE_MAIN; - message.header.protocol = RouteProtocol::Static; - message.header.scope = RouteScope::Universe; - message.header.kind = RouteType::Unicast; - - message - .attributes - .push(RouteAttribute::Priority(cost.unwrap_or(65535) as u32)); - - message - .attributes - .push(RouteAttribute::Oif(NetlinkIfConfiger::get_interface_index( - name, - )?)); - + let mut builder = RouteMessageBuilder::new(libc::AF_INET6 as u8) + .oif(Self::get_interface_index(name)?) + .priority(cost.unwrap_or(65535) as u32) + .table(libc::RT_TABLE_MAIN.into()) + .static_protocol() + .universe_scope() + .route_type(RouteType::Unicast); if cidr_prefix != 0 { - message - .attributes - .push(RouteAttribute::Destination(RouteAddress::Inet6(address))); + builder = builder.destination(IpAddr::V6(address), cidr_prefix); } - - send_netlink_req_and_wait_one_resp(RouteNetlinkMessage::NewRoute(message), false) + let message = builder.build(); + let request = message_request( + RTM_NEWROUTE, + NLM_F_ACK | NLM_F_CREATE | NLM_F_EXCL | NLM_F_REQUEST, + &message, + )?; + send_netlink_req_and_wait_ack(request) } async fn remove_ipv6_route( @@ -688,16 +517,20 @@ impl IfConfiguerTrait for NetlinkIfConfiger { address: std::net::Ipv6Addr, cidr_prefix: u8, ) -> Result<(), Error> { - let routes = Self::list_route_messages(AddressFamily::Inet6)?; + let routes = Self::list_route_messages(libc::AF_INET6 as u8)?; let ifidx = NetlinkIfConfiger::get_interface_index(name)?; for msg in routes { - let other_route: Route = msg.clone().into(); - if other_route.destination == std::net::IpAddr::V6(address) - && other_route.prefix == cidr_prefix - && other_route.ifindex == Some(ifidx) + let destination = msg + .destination() + .copied() + .unwrap_or(IpAddr::V6(Ipv6Addr::UNSPECIFIED)); + if destination == IpAddr::V6(address) + && msg.dst_len() == cidr_prefix + && msg.oif() == Some(ifidx) { - send_netlink_req_and_wait_one_resp(RouteNetlinkMessage::DelRoute(msg), true)?; + let request = message_request(RTM_DELROUTE, NLM_F_ACK | NLM_F_REQUEST, &msg)?; + send_netlink_req_and_wait_ack(request)?; return Ok(()); } } @@ -848,8 +681,7 @@ mod tests { let routes = NetlinkIfConfiger::list_routes() .unwrap() .into_iter() - .map(Route::from) - .map(|x| x.destination) + .filter_map(|route| route.destination().copied()) .collect::>(); assert!(routes.contains(&IpAddr::V4("10.5.5.0".parse().unwrap()))); @@ -860,8 +692,7 @@ mod tests { let routes = NetlinkIfConfiger::list_routes() .unwrap() .into_iter() - .map(Route::from) - .map(|x| x.destination) + .filter_map(|route| route.destination().copied()) .collect::>(); assert!(!routes.contains(&IpAddr::V4("10.5.5.0".parse().unwrap()))); } @@ -912,39 +743,24 @@ mod tests { let routes = NetlinkIfConfiger::list_ipv6_route_messages().unwrap(); assert!(routes.iter().any(|route| { - route.header.kind == RouteType::Unicast - && route.header.source_prefix_length == 56 - && route.attributes.iter().any(|attr| { - matches!( - attr, - RouteAttribute::Source(RouteAddress::Inet6(addr)) - if *addr == "2001:db8:100::".parse::().unwrap() - ) - }) - && route - .attributes - .iter() - .any(|attr| matches!(attr, RouteAttribute::Oif(index) if *index == wan_ifindex)) - && !route - .attributes - .iter() - .any(|attr| matches!(attr, RouteAttribute::Destination(_))) + route.route_type() == RouteType::Unicast + && route.src_len() == 56 + && route.source() + == Some(&IpAddr::V6( + "2001:db8:100::".parse::().unwrap(), + )) + && route.oif() == Some(wan_ifindex) + && route.destination().is_none() })); assert!(routes.iter().any(|route| { - route.header.kind == RouteType::Unicast - && route.header.destination_prefix_length == 56 - && route.attributes.iter().any(|attr| { - matches!( - attr, - RouteAttribute::Destination(RouteAddress::Inet6(addr)) - if *addr == "2001:db8:100::".parse::().unwrap() - ) - }) - && route - .attributes - .iter() - .any(|attr| matches!(attr, RouteAttribute::Oif(index) if *index == lan_ifindex)) + route.route_type() == RouteType::Unicast + && route.dst_len() == 56 + && route.destination() + == Some(&IpAddr::V6( + "2001:db8:100::".parse::().unwrap(), + )) + && route.oif() == Some(lan_ifindex) })); } @@ -964,17 +780,9 @@ mod tests { let ifindex = NetlinkIfConfiger::get_interface_index(&iface).unwrap(); let has_route = |routes: &[RouteMessage]| { routes.iter().any(|route| { - route.header.destination_prefix_length == 56 - && route.attributes.iter().any(|attr| { - matches!( - attr, - RouteAttribute::Destination(RouteAddress::Inet6(addr)) if *addr == route_addr - ) - }) - && route - .attributes - .iter() - .any(|attr| matches!(attr, RouteAttribute::Oif(index) if *index == ifindex)) + route.dst_len() == 56 + && route.destination() == Some(&IpAddr::V6(route_addr)) + && route.oif() == Some(ifindex) }) }; diff --git a/easytier/src/common/ifcfg/netlink_wire.rs b/easytier/src/common/ifcfg/netlink_wire.rs new file mode 100644 index 00000000..6bd92925 --- /dev/null +++ b/easytier/src/common/ifcfg/netlink_wire.rs @@ -0,0 +1,682 @@ +use std::{ + io, + net::{IpAddr, Ipv4Addr, Ipv6Addr}, +}; + +use nix::libc; + +// Minimal Linux rtnetlink wire support used by ifcfg. Fixed headers and integer +// attributes use native byte order; IP address attributes contain network-order octets. +pub(crate) const NLM_F_REQUEST: u16 = 0x01; +pub(crate) const NLM_F_ACK: u16 = 0x04; +pub(crate) const NLM_F_DUMP_INTR: u16 = 0x10; +pub(crate) const NLM_F_DUMP: u16 = 0x300; +pub(crate) const NLM_F_EXCL: u16 = 0x200; +pub(crate) const NLM_F_CREATE: u16 = 0x400; + +pub(crate) const NLMSG_ERROR: u16 = 2; +pub(crate) const NLMSG_DONE: u16 = 3; +pub(crate) const RTM_NEWADDR: u16 = 20; +pub(crate) const RTM_DELADDR: u16 = 21; +pub(crate) const RTM_NEWROUTE: u16 = 24; +pub(crate) const RTM_DELROUTE: u16 = 25; +pub(crate) const RTM_GETROUTE: u16 = 26; +pub(crate) const RTM_NEWNEIGH: u16 = 28; +pub(crate) const RTM_DELNEIGH: u16 = 29; +pub(crate) const RTM_GETNEIGH: u16 = 30; + +const NLMSG_HEADER_LEN: usize = 16; +const ATTRIBUTE_HEADER_LEN: usize = 4; +const NLA_TYPE_MASK: u16 = 0x3fff; + +const IFA_ADDRESS: u16 = 1; +const IFA_LOCAL: u16 = 2; +const IFA_BROADCAST: u16 = 4; + +const RTA_DST: u16 = 1; +const RTA_SRC: u16 = 2; +const RTA_OIF: u16 = 4; +const RTA_PRIORITY: u16 = 6; +const RTA_TABLE: u16 = 15; + +const NDA_DST: u16 = 1; +const NTF_PROXY: u8 = 0x08; +const NUD_PERMANENT: u16 = 0x80; + +fn align4(len: usize) -> Option { + len.checked_add(3).map(|len| len & !3) +} + +fn invalid_data(message: &'static str) -> io::Error { + io::Error::new(io::ErrorKind::InvalidData, message) +} + +fn read_u16(bytes: &[u8]) -> io::Result { + bytes + .get(..2) + .and_then(|bytes| bytes.try_into().ok()) + .map(u16::from_ne_bytes) + .ok_or_else(|| invalid_data("truncated netlink u16")) +} + +fn read_u32(bytes: &[u8]) -> io::Result { + bytes + .get(..4) + .and_then(|bytes| bytes.try_into().ok()) + .map(u32::from_ne_bytes) + .ok_or_else(|| invalid_data("truncated netlink u32")) +} + +fn read_i32(bytes: &[u8]) -> io::Result { + bytes + .get(..4) + .and_then(|bytes| bytes.try_into().ok()) + .map(i32::from_ne_bytes) + .ok_or_else(|| invalid_data("truncated netlink i32")) +} + +#[derive(Debug)] +pub(crate) struct MessageBuilder { + bytes: Vec, +} + +impl MessageBuilder { + pub(crate) fn new(message_type: u16, flags: u16) -> Self { + let mut bytes = Vec::with_capacity(64); + bytes.extend_from_slice(&0_u32.to_ne_bytes()); + bytes.extend_from_slice(&message_type.to_ne_bytes()); + bytes.extend_from_slice(&flags.to_ne_bytes()); + bytes.extend_from_slice(&0_u32.to_ne_bytes()); + bytes.extend_from_slice(&0_u32.to_ne_bytes()); + Self { bytes } + } + + pub(crate) fn append_bytes(&mut self, bytes: &[u8]) { + self.bytes.extend_from_slice(bytes); + } + + pub(crate) fn finish(mut self) -> io::Result> { + let len = u32::try_from(self.bytes.len()) + .map_err(|_| invalid_data("netlink message is too large"))?; + self.bytes[..4].copy_from_slice(&len.to_ne_bytes()); + Ok(self.bytes) + } +} + +#[derive(Clone, Copy, Debug)] +pub(crate) struct MessageHeader { + pub(crate) message_type: u16, + pub(crate) flags: u16, +} + +pub(crate) struct MessageIter<'a> { + bytes: &'a [u8], +} + +impl<'a> MessageIter<'a> { + pub(crate) fn new(bytes: &'a [u8]) -> Self { + Self { bytes } + } +} + +impl<'a> Iterator for MessageIter<'a> { + type Item = io::Result<(MessageHeader, &'a [u8])>; + + fn next(&mut self) -> Option { + if self.bytes.is_empty() { + return None; + } + if self.bytes.len() < NLMSG_HEADER_LEN { + self.bytes = &[]; + return Some(Err(invalid_data("truncated netlink header"))); + } + + let len = match read_u32(self.bytes) { + Ok(len) => len as usize, + Err(err) => { + self.bytes = &[]; + return Some(Err(err)); + } + }; + if len < NLMSG_HEADER_LEN || len > self.bytes.len() { + self.bytes = &[]; + return Some(Err(invalid_data("invalid netlink message length"))); + } + let aligned_len = match align4(len) { + Some(len) => len, + None => { + self.bytes = &[]; + return Some(Err(invalid_data("netlink message length overflow"))); + } + }; + if aligned_len > self.bytes.len() { + self.bytes = &[]; + return Some(Err(invalid_data("truncated netlink message padding"))); + } + + let message_type = read_u16(&self.bytes[4..]).expect("header length checked"); + let flags = read_u16(&self.bytes[6..]).expect("header length checked"); + let payload = &self.bytes[NLMSG_HEADER_LEN..len]; + self.bytes = &self.bytes[aligned_len..]; + Some(Ok(( + MessageHeader { + message_type, + flags, + }, + payload, + ))) + } +} + +pub(crate) fn netlink_error_code(payload: &[u8]) -> io::Result { + read_i32(payload) +} + +#[derive(Clone, Debug)] +struct Attribute { + kind: u16, + value: Vec, +} + +impl Attribute { + fn new(kind: u16, value: impl Into>) -> Self { + Self { + kind, + value: value.into(), + } + } + + fn write_to(&self, bytes: &mut Vec) -> io::Result<()> { + let len = ATTRIBUTE_HEADER_LEN + .checked_add(self.value.len()) + .ok_or_else(|| invalid_data("netlink attribute length overflow"))?; + let len = u16::try_from(len).map_err(|_| invalid_data("netlink attribute is too large"))?; + bytes.extend_from_slice(&len.to_ne_bytes()); + bytes.extend_from_slice(&self.kind.to_ne_bytes()); + bytes.extend_from_slice(&self.value); + let aligned_len = align4(bytes.len()) + .ok_or_else(|| invalid_data("netlink attribute alignment overflow"))?; + bytes.resize(aligned_len, 0); + Ok(()) + } +} + +fn parse_attributes(mut bytes: &[u8]) -> io::Result> { + let mut attributes = Vec::new(); + while !bytes.is_empty() { + if bytes.len() < ATTRIBUTE_HEADER_LEN { + return Err(invalid_data("truncated netlink attribute header")); + } + let len = read_u16(bytes)? as usize; + if len < ATTRIBUTE_HEADER_LEN || len > bytes.len() { + return Err(invalid_data("invalid netlink attribute length")); + } + let kind = read_u16(&bytes[2..])?; + attributes.push(Attribute::new( + kind, + bytes[ATTRIBUTE_HEADER_LEN..len].to_vec(), + )); + + let aligned_len = + align4(len).ok_or_else(|| invalid_data("netlink attribute length overflow"))?; + if aligned_len > bytes.len() { + return Err(invalid_data("truncated netlink attribute padding")); + } + bytes = &bytes[aligned_len..]; + } + Ok(attributes) +} + +fn write_attributes(attributes: &[Attribute], bytes: &mut Vec) -> io::Result<()> { + for attribute in attributes { + attribute.write_to(bytes)?; + } + Ok(()) +} + +fn ip_bytes(address: IpAddr) -> Vec { + match address { + IpAddr::V4(address) => address.octets().to_vec(), + IpAddr::V6(address) => address.octets().to_vec(), + } +} + +fn parse_ip(family: u8, bytes: &[u8]) -> Option { + match family as i32 { + libc::AF_INET if bytes.len() == 4 => Some(IpAddr::V4(Ipv4Addr::new( + bytes[0], bytes[1], bytes[2], bytes[3], + ))), + libc::AF_INET6 if bytes.len() == 16 => Some(IpAddr::V6(Ipv6Addr::from( + <[u8; 16]>::try_from(bytes).ok()?, + ))), + _ => None, + } +} + +pub(crate) trait NetlinkEncode { + fn write_to(&self, bytes: &mut Vec) -> io::Result<()>; +} + +pub(crate) trait NetlinkDecode: Sized { + const MESSAGE_TYPE: u16; + + fn from_bytes(bytes: &[u8]) -> io::Result; +} + +#[derive(Clone, Debug)] +pub(crate) struct AddressMessage { + family: u8, + prefix_len: u8, + ifindex: u32, + attributes: Vec, +} + +impl AddressMessage { + pub(crate) fn new(family: u8, ifindex: u32, prefix_len: u8, address: IpAddr) -> Self { + Self { + family, + prefix_len, + ifindex, + attributes: vec![Attribute::new(IFA_ADDRESS, ip_bytes(address))], + } + } + + pub(crate) fn local(mut self, address: IpAddr) -> Self { + self.attributes + .push(Attribute::new(IFA_LOCAL, ip_bytes(address))); + self + } + + pub(crate) fn broadcast(mut self, address: Ipv4Addr) -> Self { + self.attributes + .push(Attribute::new(IFA_BROADCAST, address.octets().to_vec())); + self + } +} + +impl NetlinkEncode for AddressMessage { + fn write_to(&self, bytes: &mut Vec) -> io::Result<()> { + bytes.extend_from_slice(&[self.family, self.prefix_len, 0, 0]); + bytes.extend_from_slice(&self.ifindex.to_ne_bytes()); + write_attributes(&self.attributes, bytes) + } +} + +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub(crate) enum RouteType { + Unicast, + Blackhole, + Other(u8), +} + +impl RouteType { + fn from_raw(value: u8) -> Self { + match value { + 1 => Self::Unicast, + 6 => Self::Blackhole, + value => Self::Other(value), + } + } + + fn raw(self) -> u8 { + match self { + Self::Unicast => 1, + Self::Blackhole => 6, + Self::Other(value) => value, + } + } +} + +#[derive(Clone, Debug)] +pub(crate) struct RouteMessage { + family: u8, + dst_len: u8, + src_len: u8, + tos: u8, + table: u8, + protocol: u8, + scope: u8, + route_type: u8, + flags: u32, + attributes: Vec, + destination: Option, + source: Option, + oif: Option, +} + +impl RouteMessage { + pub(crate) fn family(&self) -> u8 { + self.family + } + + pub(crate) fn dst_len(&self) -> u8 { + self.dst_len + } + + pub(crate) fn src_len(&self) -> u8 { + self.src_len + } + + pub(crate) fn route_type(&self) -> RouteType { + RouteType::from_raw(self.route_type) + } + + pub(crate) fn destination(&self) -> Option<&IpAddr> { + self.destination.as_ref() + } + + pub(crate) fn source(&self) -> Option<&IpAddr> { + self.source.as_ref() + } + + pub(crate) fn oif(&self) -> Option { + self.oif + } +} + +impl RouteMessage { + pub(crate) fn dump_header() -> Vec { + vec![0; 12] + } +} + +impl NetlinkDecode for RouteMessage { + const MESSAGE_TYPE: u16 = RTM_NEWROUTE; + + fn from_bytes(bytes: &[u8]) -> io::Result { + if bytes.len() < 12 { + return Err(invalid_data("truncated route message")); + } + let family = bytes[0]; + let attributes = parse_attributes(&bytes[12..])?; + let destination = attributes + .iter() + .find(|attribute| attribute.kind & NLA_TYPE_MASK == RTA_DST) + .and_then(|attribute| parse_ip(family, &attribute.value)); + let source = attributes + .iter() + .find(|attribute| attribute.kind & NLA_TYPE_MASK == RTA_SRC) + .and_then(|attribute| parse_ip(family, &attribute.value)); + let oif = attributes + .iter() + .find(|attribute| attribute.kind & NLA_TYPE_MASK == RTA_OIF) + .and_then(|attribute| read_u32(&attribute.value).ok()); + + Ok(Self { + family, + dst_len: bytes[1], + src_len: bytes[2], + tos: bytes[3], + table: bytes[4], + protocol: bytes[5], + scope: bytes[6], + route_type: bytes[7], + flags: read_u32(&bytes[8..])?, + attributes, + destination, + source, + oif, + }) + } +} + +impl NetlinkEncode for RouteMessage { + fn write_to(&self, bytes: &mut Vec) -> io::Result<()> { + bytes.extend_from_slice(&[ + self.family, + self.dst_len, + self.src_len, + self.tos, + self.table, + self.protocol, + self.scope, + self.route_type, + ]); + bytes.extend_from_slice(&self.flags.to_ne_bytes()); + write_attributes(&self.attributes, bytes) + } +} + +#[derive(Debug)] +pub(crate) struct RouteMessageBuilder { + message: RouteMessage, +} + +impl RouteMessageBuilder { + pub(crate) fn new(family: u8) -> Self { + Self { + message: RouteMessage { + family, + dst_len: 0, + src_len: 0, + tos: 0, + table: 0, + protocol: 0, + scope: 0, + route_type: 0, + flags: 0, + attributes: Vec::new(), + destination: None, + source: None, + oif: None, + }, + } + } + + pub(crate) fn destination(mut self, address: IpAddr, prefix_len: u8) -> Self { + self.message.dst_len = prefix_len; + self.message.destination = Some(address); + self.message + .attributes + .push(Attribute::new(RTA_DST, ip_bytes(address))); + self + } + + pub(crate) fn oif(mut self, ifindex: u32) -> Self { + self.message.oif = Some(ifindex); + self.message + .attributes + .push(Attribute::new(RTA_OIF, ifindex.to_ne_bytes().to_vec())); + self + } + + pub(crate) fn priority(mut self, priority: u32) -> Self { + self.message.attributes.push(Attribute::new( + RTA_PRIORITY, + priority.to_ne_bytes().to_vec(), + )); + self + } + + pub(crate) fn table(mut self, table: u32) -> Self { + if let Ok(table) = u8::try_from(table) { + self.message.table = table; + } else { + self.message + .attributes + .push(Attribute::new(RTA_TABLE, table.to_ne_bytes().to_vec())); + } + self + } + + pub(crate) fn static_protocol(mut self) -> Self { + self.message.protocol = libc::RTPROT_STATIC; + self + } + + pub(crate) fn universe_scope(mut self) -> Self { + self.message.scope = libc::RT_SCOPE_UNIVERSE; + self + } + + pub(crate) fn route_type(mut self, route_type: RouteType) -> Self { + self.message.route_type = route_type.raw(); + self + } + + pub(crate) fn build(self) -> RouteMessage { + self.message + } +} + +#[derive(Clone, Debug)] +pub(crate) struct NeighborMessage { + family: u8, + ifindex: u32, + state: u16, + flags: u8, + kind: u8, + attributes: Vec, + destination: Option, +} + +impl NeighborMessage { + pub(crate) fn proxy(ifindex: u32, address: Ipv6Addr) -> Self { + Self { + family: libc::AF_INET6 as u8, + ifindex, + state: NUD_PERMANENT, + flags: NTF_PROXY, + kind: 0, + attributes: vec![Attribute::new(NDA_DST, address.octets().to_vec())], + destination: Some(IpAddr::V6(address)), + } + } + + pub(crate) fn proxy_dump_header(family: u8) -> Vec { + let mut bytes = vec![family, 0, 0, 0]; + bytes.extend_from_slice(&0_u32.to_ne_bytes()); + bytes.extend_from_slice(&0_u16.to_ne_bytes()); + bytes.extend_from_slice(&[NTF_PROXY, 0]); + bytes + } + + pub(crate) fn ifindex(&self) -> u32 { + self.ifindex + } + + pub(crate) fn is_proxy(&self) -> bool { + self.flags & NTF_PROXY != 0 + } + + pub(crate) fn destination(&self) -> Option<&IpAddr> { + self.destination.as_ref() + } +} + +impl NetlinkDecode for NeighborMessage { + const MESSAGE_TYPE: u16 = RTM_NEWNEIGH; + + fn from_bytes(bytes: &[u8]) -> io::Result { + if bytes.len() < 12 { + return Err(invalid_data("truncated neighbor message")); + } + let family = bytes[0]; + let attributes = parse_attributes(&bytes[12..])?; + let destination = attributes + .iter() + .find(|attribute| attribute.kind & NLA_TYPE_MASK == NDA_DST) + .and_then(|attribute| parse_ip(family, &attribute.value)); + Ok(Self { + family, + ifindex: read_u32(&bytes[4..])?, + state: read_u16(&bytes[8..])?, + flags: bytes[10], + kind: bytes[11], + attributes, + destination, + }) + } +} + +impl NetlinkEncode for NeighborMessage { + fn write_to(&self, bytes: &mut Vec) -> io::Result<()> { + bytes.extend_from_slice(&[self.family, 0, 0, 0]); + bytes.extend_from_slice(&self.ifindex.to_ne_bytes()); + bytes.extend_from_slice(&self.state.to_ne_bytes()); + bytes.extend_from_slice(&[self.flags, self.kind]); + write_attributes(&self.attributes, bytes) + } +} + +#[cfg(test)] +mod tests { + use super::*; + + fn encode(message: &T) -> Vec { + let mut bytes = Vec::new(); + message.write_to(&mut bytes).unwrap(); + bytes + } + + #[test] + fn route_round_trip_preserves_unknown_attributes() { + let mut message = RouteMessageBuilder::new(libc::AF_INET6 as u8) + .destination("2001:db8::".parse().unwrap(), 64) + .oif(7) + .priority(42) + .table(libc::RT_TABLE_MAIN.into()) + .static_protocol() + .universe_scope() + .route_type(RouteType::Unicast) + .build(); + message + .attributes + .push(Attribute::new(0x4321, vec![1, 2, 3, 4, 5])); + + let bytes = encode(&message); + let decoded = RouteMessage::from_bytes(&bytes).unwrap(); + assert_eq!(decoded.destination(), message.destination()); + assert_eq!(decoded.oif(), Some(7)); + assert_eq!(decoded.route_type(), RouteType::Unicast); + assert_eq!(encode(&decoded), bytes); + } + + #[test] + fn route_parser_reads_ipv6_source_prefix() { + let mut bytes = vec![ + libc::AF_INET6 as u8, + 0, + 56, + 0, + libc::RT_TABLE_MAIN, + libc::RTPROT_STATIC, + libc::RT_SCOPE_UNIVERSE, + RouteType::Unicast.raw(), + ]; + bytes.extend_from_slice(&0_u32.to_ne_bytes()); + Attribute::new( + RTA_SRC, + "2001:db8:1::" + .parse::() + .unwrap() + .octets() + .to_vec(), + ) + .write_to(&mut bytes) + .unwrap(); + + let message = RouteMessage::from_bytes(&bytes).unwrap(); + assert_eq!(message.src_len(), 56); + assert_eq!(message.source(), Some(&"2001:db8:1::".parse().unwrap())); + } + + #[test] + fn neighbor_proxy_round_trip() { + let message = NeighborMessage::proxy(9, "2001:db8::1".parse().unwrap()); + let decoded = NeighborMessage::from_bytes(&encode(&message)).unwrap(); + assert_eq!(decoded.ifindex(), 9); + assert!(decoded.is_proxy()); + assert_eq!(decoded.destination(), message.destination()); + } + + #[test] + fn message_iterator_rejects_truncated_padding() { + let mut bytes = MessageBuilder::new(RTM_GETROUTE, NLM_F_REQUEST) + .finish() + .unwrap(); + bytes[0..4].copy_from_slice(&17_u32.to_ne_bytes()); + bytes.push(0); + assert!(MessageIter::new(&bytes).next().unwrap().is_err()); + } +} diff --git a/easytier/src/common/ifcfg/route.rs b/easytier/src/common/ifcfg/route.rs index 5e428b0b..d3bd788b 100644 --- a/easytier/src/common/ifcfg/route.rs +++ b/easytier/src/common/ifcfg/route.rs @@ -1,4 +1,4 @@ -use std::net::{IpAddr, Ipv4Addr, Ipv6Addr}; +use std::net::IpAddr; #[derive(Debug, Clone, PartialEq, Eq)] pub struct Route { @@ -44,90 +44,3 @@ pub struct Route { /// If luid is specified, ifindex is optional. pub luid: Option, } - -impl Route { - /// Create a route that matches a given destination network. - /// - /// Either the gateway or interface should be set before attempting to add to a routing table. - pub fn new(destination: IpAddr, prefix: u8) -> Self { - Self { - destination, - prefix, - gateway: None, - ifindex: None, - #[cfg(target_os = "linux")] - // default to main table - table: 254, - #[cfg(target_os = "linux")] - source: None, - #[cfg(target_os = "linux")] - source_prefix: 0, - #[cfg(target_os = "linux")] - source_hint: None, - #[cfg(any(target_os = "windows", target_os = "linux"))] - metric: None, - #[cfg(target_os = "windows")] - luid: None, - } - } - - /// Set the next next hop gateway for this route. - pub fn with_gateway(mut self, gateway: IpAddr) -> Self { - self.gateway = Some(gateway); - self - } - - /// Set the index of the local interface through which the next hop of this route should be reached. - pub fn with_ifindex(mut self, ifindex: u32) -> Self { - self.ifindex = Some(ifindex); - self - } - - /// Set table the route will be installed in. - #[cfg(target_os = "linux")] - pub fn with_table(mut self, table: u8) -> Self { - self.table = table; - self - } - - /// Set source. - #[cfg(target_os = "linux")] - pub fn with_source(mut self, source: IpAddr, prefix: u8) -> Self { - self.source = Some(source); - self.source_prefix = prefix; - self - } - - /// Set source hint. - #[cfg(target_os = "linux")] - pub fn with_source_hint(mut self, hint: IpAddr) -> Self { - self.source_hint = Some(hint); - self - } - - /// Set route metric. - #[cfg(any(target_os = "windows", target_os = "linux"))] - pub fn with_metric(mut self, metric: u32) -> Self { - self.metric = Some(metric); - self - } - - /// Set luid of the local interface through which the next hop of this route should be reached. - #[cfg(target_os = "windows")] - pub fn with_luid(mut self, luid: u64) -> Self { - self.luid = Some(luid); - self - } - - /// Get the netmask covering the network portion of the destination address. - pub fn mask(&self) -> IpAddr { - match self.destination { - IpAddr::V4(_) => IpAddr::V4(Ipv4Addr::from( - u32::MAX.checked_shl(32 - self.prefix as u32).unwrap_or(0), - )), - IpAddr::V6(_) => IpAddr::V6(Ipv6Addr::from( - u128::MAX.checked_shl(128 - self.prefix as u32).unwrap_or(0), - )), - } - } -} diff --git a/easytier/src/common/ifcfg/win/luid.rs b/easytier/src/common/ifcfg/win/luid.rs index 67405440..d4832e25 100644 --- a/easytier/src/common/ifcfg/win/luid.rs +++ b/easytier/src/common/ifcfg/win/luid.rs @@ -5,107 +5,18 @@ // ATTENTION: NOT included are DNS() and SetDNS() - functions to query and set DNS servers for a network interface. // -use super::netsh; use super::types::*; use cidr::Ipv4Inet; use cidr::Ipv6Inet; use std::net::{Ipv4Addr, Ipv6Addr}; use std::ptr; -use winapi::shared::{ - guiddef::GUID, ifdef::NET_LUID, netioapi::*, nldef::*, winerror::*, ws2def::*, ws2ipdef::*, -}; +use winapi::shared::{ifdef::NET_LUID, netioapi::*, nldef::*, winerror::*, ws2def::*}; pub struct InterfaceLuid { luid: NET_LUID, } impl InterfaceLuid { - pub fn new(luid_value: u64) -> Self { - InterfaceLuid { - luid: winapi::shared::ifdef::NET_LUID_LH { Value: luid_value }, - } - } - - pub fn luid(&self) -> NET_LUID { - self.luid - } - - /// get_ip_interface method retrieves IP information for the specified interface on the local computer. - pub fn get_ip_interface( - &self, - family: ADDRESS_FAMILY, - ) -> Result { - let mut row = MIB_IPINTERFACE_ROW::default(); - unsafe { InitializeIpInterfaceEntry(&mut row) }; - - row.InterfaceLuid = self.luid; - row.Family = family; - - let result = unsafe { GetIpInterfaceEntry(&mut row) }; - if NO_ERROR == result { - Ok(row) - } else { - Err(result) - } - } - - /// https://learn.microsoft.com/en-us/windows/win32/api/netioapi/nf-netioapi-setipinterfaceentry - /// If only InterfaceIndex was specified, SetIpInterfaceEntry() will modify ipif with a correct InterfaceLuid - pub fn set_ip_interface(&self, ipif: *mut MIB_IPINTERFACE_ROW) -> Result<(), NETIO_STATUS> { - let result = unsafe { SetIpInterfaceEntry(ipif) }; - if NO_ERROR == result { - Ok(()) - } else { - Err(result) - } - } - - /// get_interface method retrieves information for the specified adapter on the local computer. - /// https://docs.microsoft.com/en-us/windows/desktop/api/netioapi/nf-netioapi-getifentry2 - pub fn get_interface(&self) -> Result { - let mut row = MIB_IF_ROW2 { - InterfaceLuid: self.luid, - ..MIB_IF_ROW2::default() - }; - - let result = unsafe { GetIfEntry2(&mut row) }; - if NO_ERROR == result { - Ok(row) - } else { - Err(result) - } - } - - /// GUID method converts a locally unique identifier (LUID) for a network interface to a globally unique identifier (GUID) for the interface. - /// https://docs.microsoft.com/en-us/windows/desktop/api/netioapi/nf-netioapi-convertinterfaceluidtoguid - pub fn get_guid(&self) -> Result { - let mut interface_guid = GUID::default(); - - let result = unsafe { ConvertInterfaceLuidToGuid(&self.luid, &mut interface_guid) }; - - if NO_ERROR == result { - Ok(interface_guid) - } else { - Err(result) - } - } - - /// luid_from_guid function converts a globally unique identifier (GUID) for a network interface to the locally unique identifier (LUID) for the interface. - /// https://docs.microsoft.com/en-us/windows/desktop/api/netioapi/nf-netioapi-convertinterfaceguidtoluid - pub fn luid_from_guid(interface_guid: &GUID) -> Result { - let mut interface_luid = NET_LUID::default(); - - let result = unsafe { ConvertInterfaceGuidToLuid(interface_guid, &mut interface_luid) }; - - if NO_ERROR == result { - Ok(Self { - luid: interface_luid, - }) - } else { - Err(result) - } - } - /// luid_from_index function converts a local index for a network interface to the locally unique identifier (LUID) for the interface. /// https://docs.microsoft.com/en-us/windows/desktop/api/netioapi/nf-netioapi-convertinterfaceindextoluid pub fn luid_from_index(interface_index: u32) -> Result { @@ -122,46 +33,6 @@ impl InterfaceLuid { } } - /// get_from_ipv4_address method returns MibUnicastIPAddressRow struct that matches to provided 'ip' argument. Corresponds to GetUnicastIpAddressEntry - /// (https://docs.microsoft.com/en-us/windows/desktop/api/netioapi/nf-netioapi-getunicastipaddressentry) - pub fn get_from_ipv4_address( - &self, - ip: &Ipv4Addr, - ) -> Result { - let mut row = MIB_UNICASTIPADDRESS_ROW::default(); - unsafe { InitializeUnicastIpAddressEntry(&mut row) }; - - unsafe { *row.Address.Ipv4_mut() = convert_ipv4addr_to_sockaddr(ip) }; - - let result = unsafe { GetUnicastIpAddressEntry(&mut row) }; - - if NO_ERROR == result { - Ok(row) - } else { - Err(result) - } - } - - /// get_from_ipv6_address method returns MibUnicastIPAddressRow struct that matches to provided 'ip' argument. Corresponds to GetUnicastIpAddressEntry - /// (https://docs.microsoft.com/en-us/windows/desktop/api/netioapi/nf-netioapi-getunicastipaddressentry) - pub fn get_from_ipv6_address( - &self, - ip: &Ipv6Addr, - ) -> Result { - let mut row = MIB_UNICASTIPADDRESS_ROW::default(); - unsafe { InitializeUnicastIpAddressEntry(&mut row) }; - - unsafe { *row.Address.Ipv6_mut() = convert_ipv6addr_to_sockaddr(ip) }; - - let result = unsafe { GetUnicastIpAddressEntry(&mut row) }; - - if NO_ERROR == result { - Ok(row) - } else { - Err(result) - } - } - /// add_ipv4_address method adds new unicast IP address to the interface. Corresponds to CreateUnicastIpAddressEntry function /// (https://docs.microsoft.com/en-us/windows/desktop/api/netioapi/nf-netioapi-createunicastipaddressentry). pub fn add_ipv4_address(&self, address: &Ipv4Inet) -> Result<(), NETIO_STATUS> { @@ -208,50 +79,6 @@ impl InterfaceLuid { } } - /// add_ipv4_addresses method adds multiple new unicast IP addresses to the interface. Corresponds to CreateUnicastIpAddressEntry function - /// (https://docs.microsoft.com/en-us/windows/desktop/api/netioapi/nf-netioapi-createunicastipaddressentry). - pub fn add_ipv4_addresses( - &self, - addresses: impl IntoIterator, - ) -> Result<(), NETIO_STATUS> { - for ip in addresses.into_iter().enumerate() { - self.add_ipv4_address(&ip.1)?; - } - Ok(()) - } - - /// add_ipv6_addresses method adds multiple new unicast IP addresses to the interface. Corresponds to CreateUnicastIpAddressEntry function - /// (https://docs.microsoft.com/en-us/windows/desktop/api/netioapi/nf-netioapi-createunicastipaddressentry). - pub fn add_ipv6_addresses( - &self, - addresses: impl IntoIterator, - ) -> Result<(), NETIO_STATUS> { - for ip in addresses.into_iter().enumerate() { - self.add_ipv6_address(&ip.1)?; - } - Ok(()) - } - - /// set_ipv4_addresses method sets new unicast IP addresses to the interface. - pub fn set_ipv4_addresses( - &self, - addresses: impl IntoIterator, - ) -> Result<(), NETIO_STATUS> { - self.flush_ipv4_addresses()?; - self.add_ipv4_addresses(addresses)?; - Ok(()) - } - - /// set_ipv6_addresses method sets new unicast IP addresses to the interface. - pub fn set_ipv6_addresses( - &self, - addresses: impl IntoIterator, - ) -> Result<(), NETIO_STATUS> { - self.flush_ipv6_addresses()?; - self.add_ipv6_addresses(addresses)?; - Ok(()) - } - /// delete_ipv4_address method deletes interface's unicast IP address. Corresponds to DeleteUnicastIpAddressEntry function /// (https://docs.microsoft.com/en-us/windows/desktop/api/netioapi/nf-netioapi-deleteunicastipaddressentry). pub fn delete_ipv4_address(&self, address: &Ipv4Inet) -> Result<(), NETIO_STATUS> { @@ -275,34 +102,6 @@ impl InterfaceLuid { } } - /// delete_ipv4_address method deletes interface's unicast IP address. Corresponds to DeleteUnicastIpAddressEntry function - /// (https://docs.microsoft.com/en-us/windows/desktop/api/netioapi/nf-netioapi-deleteunicastipaddressentry). - pub fn delete_ipv4_address2( - &self, - address: *const SOCKADDR_IN, - prefix_len: u8, - ) -> Result<(), NETIO_STATUS> { - let mut row = MIB_UNICASTIPADDRESS_ROW::default(); - unsafe { InitializeUnicastIpAddressEntry(&mut row) }; - - row.InterfaceLuid = self.luid; - row.DadState = IpDadStatePreferred; - row.ValidLifetime = 0xffffffff; - row.PreferredLifetime = 0xffffffff; - - assert!(!address.is_null()); - unsafe { *row.Address.Ipv4_mut() = *address }; - row.OnLinkPrefixLength = prefix_len; - - let result = unsafe { DeleteUnicastIpAddressEntry(&row) }; - - if NO_ERROR == result { - Ok(()) - } else { - Err(result) - } - } - /// delete_ipv6_address method deletes interface's unicast IP address. Corresponds to DeleteUnicastIpAddressEntry function /// (https://docs.microsoft.com/en-us/windows/desktop/api/netioapi/nf-netioapi-deleteunicastipaddressentry). pub fn delete_ipv6_address(&self, address: &Ipv6Inet) -> Result<(), NETIO_STATUS> { @@ -326,34 +125,6 @@ impl InterfaceLuid { } } - /// delete_ipv6_address method deletes interface's unicast IP address. Corresponds to DeleteUnicastIpAddressEntry function - /// (https://docs.microsoft.com/en-us/windows/desktop/api/netioapi/nf-netioapi-deleteunicastipaddressentry). - pub fn delete_ipv6_address2( - &self, - address: *const SOCKADDR_IN6, - prefix_len: u8, - ) -> Result<(), NETIO_STATUS> { - let mut row = MIB_UNICASTIPADDRESS_ROW::default(); - unsafe { InitializeUnicastIpAddressEntry(&mut row) }; - - row.InterfaceLuid = self.luid; - row.DadState = IpDadStatePreferred; - row.ValidLifetime = 0xffffffff; - row.PreferredLifetime = 0xffffffff; - - assert!(!address.is_null()); - unsafe { *row.Address.Ipv6_mut() = *address }; - row.OnLinkPrefixLength = prefix_len; - - let result = unsafe { DeleteUnicastIpAddressEntry(&row) }; - - if NO_ERROR == result { - Ok(()) - } else { - Err(result) - } - } - /// flush_ip_addresses method deletes all interface's unicast IP addresses. pub fn flush_ip_addresses(&self, address_family: ADDRESS_FAMILY) -> Result<(), NETIO_STATUS> { let mut p_table: PMIB_UNICASTIPADDRESS_TABLE = ptr::null_mut(); @@ -387,70 +158,6 @@ impl InterfaceLuid { self.flush_ip_addresses(AF_INET6 as _) } - /// route_ipv4 method returns route determined with the input arguments. Corresponds to GetIpForwardEntry2 function - /// (https://docs.microsoft.com/en-us/windows/desktop/api/netioapi/nf-netioapi-getipforwardentry2). - /// NOTE: If the corresponding route isn't found, the method will return error. - pub fn route_ipv4( - &self, - destination: &Ipv4Inet, - next_hop: &Ipv4Addr, - ) -> Result { - let mut row = MIB_IPFORWARD_ROW2::default(); - unsafe { InitializeIpForwardEntry(&mut row) }; - - row.InterfaceLuid = self.luid; - row.ValidLifetime = 0xffffffff; - row.PreferredLifetime = 0xffffffff; - - unsafe { - *row.DestinationPrefix.Prefix.Ipv4_mut() = - convert_ipv4addr_to_sockaddr(&destination.address()) - }; - row.DestinationPrefix.PrefixLength = destination.network_length(); - - unsafe { *row.NextHop.Ipv4_mut() = convert_ipv4addr_to_sockaddr(next_hop) }; - - let result = unsafe { GetIpForwardEntry2(&mut row) }; - - if NO_ERROR == result { - Ok(row) - } else { - Err(result) - } - } - - /// route_ipv6 method returns route determined with the input arguments. Corresponds to GetIpForwardEntry2 function - /// (https://docs.microsoft.com/en-us/windows/desktop/api/netioapi/nf-netioapi-getipforwardentry2). - /// NOTE: If the corresponding route isn't found, the method will return error. - pub fn route_ipv6( - &self, - destination: &Ipv6Inet, - next_hop: &Ipv6Addr, - ) -> Result { - let mut row = MIB_IPFORWARD_ROW2::default(); - unsafe { InitializeIpForwardEntry(&mut row) }; - - row.InterfaceLuid = self.luid; - row.ValidLifetime = 0xffffffff; - row.PreferredLifetime = 0xffffffff; - - unsafe { - *row.DestinationPrefix.Prefix.Ipv6_mut() = - convert_ipv6addr_to_sockaddr(&destination.address()) - }; - row.DestinationPrefix.PrefixLength = destination.network_length(); - - unsafe { *row.NextHop.Ipv6_mut() = convert_ipv6addr_to_sockaddr(next_hop) }; - - let result = unsafe { GetIpForwardEntry2(&mut row) }; - - if NO_ERROR == result { - Ok(row) - } else { - Err(result) - } - } - /// add_route_ipv4 method adds a route to the interface. Corresponds to CreateIpForwardEntry2 function, with added splitDefault feature. /// (https://docs.microsoft.com/en-us/windows/desktop/api/netioapi/nf-netioapi-createipforwardentry2) pub fn add_route_ipv4( @@ -541,26 +248,6 @@ impl InterfaceLuid { Ok(()) } - /// set_routes_ipv4 method sets (flush than add) multiple routes to the interface. - pub fn set_routes_ipv4( - &self, - routes_data: impl IntoIterator, - ) -> Result<(), NETIO_STATUS> { - self.flush_routes_ipv4()?; - self.add_routes_ipv4(routes_data)?; - Ok(()) - } - - /// set_routes_ipv6 method sets (flush than add) multiple routes to the interface. - pub fn set_routes_ipv6( - &self, - routes_data: impl IntoIterator, - ) -> Result<(), NETIO_STATUS> { - self.flush_routes_ipv6()?; - self.add_routes_ipv6(routes_data)?; - Ok(()) - } - /// delete_route_ipv4 method deletes a route that matches the criteria. Corresponds to DeleteIpForwardEntry2 function /// (https://docs.microsoft.com/en-us/windows/desktop/api/netioapi/nf-netioapi-deleteipforwardentry2). pub fn delete_route_ipv4( @@ -628,73 +315,6 @@ impl InterfaceLuid { } } - /// flush_routes method deletes all interface's routes. - /// It continues on failures, and returns the last error afterwards. - pub fn flush_routes(&self, address_family: ADDRESS_FAMILY) -> Result<(), NETIO_STATUS> { - let mut last_error: NETIO_STATUS = NO_ERROR; - - let mut p_table: PMIB_IPFORWARD_TABLE2 = ptr::null_mut(); - let result = unsafe { GetIpForwardTable2(address_family, &mut p_table) }; - if NO_ERROR != result { - return Err(result); - } - - assert!(!p_table.is_null()); - let num_entries = unsafe { *p_table }.NumEntries; - let x_table = unsafe { *p_table }.Table.as_ptr(); - for i in 0..num_entries { - let current_entry = unsafe { x_table.add(i as _) }; - if unsafe { (*current_entry).InterfaceLuid.Value } == self.luid.Value { - let result = unsafe { DeleteIpForwardEntry2(current_entry) }; - if NO_ERROR != result { - last_error = result; - } - } - } - - unsafe { FreeMibTable(p_table as _) }; - - if NO_ERROR == last_error { - Ok(()) - } else { - Err(result) - } - } - - /// flush_routes_ipv4 method deletes all interface's routes. - /// It continues on failures, and returns the last error afterwards. - pub fn flush_routes_ipv4(&self) -> Result<(), NETIO_STATUS> { - self.flush_routes(AF_INET as _) - } - - /// flush_routes_ipv6 method deletes all interface's routes. - /// It continues on failures, and returns the last error afterwards. - pub fn flush_routes_ipv6(&self) -> Result<(), NETIO_STATUS> { - self.flush_routes(AF_INET6 as _) - } - - /// flush_dns method clears all DNS servers associated with the adapter. - fn flush_dns(&self, family: ADDRESS_FAMILY) -> Result<(), String> { - let ip_itf = match self.get_ip_interface(family) { - Ok(ip_itf) => ip_itf, - Err(_) => { - return Err(String::from("Failed to obtain interface")); - } - }; - - netsh::flush_dns(family, ip_itf.InterfaceIndex) - } - - /// flush_dns_ipv4 method clears all DNS servers associated with the adapter. - pub fn flush_dns_ipv4(&self) -> Result<(), String> { - self.flush_dns(AF_INET as _) - } - - /// flush_dns_ipv6 method clears all DNS servers associated with the adapter. - pub fn flush_dns_ipv6(&self) -> Result<(), String> { - self.flush_dns(AF_INET6 as _) - } - /// Sets MTU on the interface /// TODO: Set IP and other things in here too, so the code is more organized pub fn set_iface_config(&self, mtu: u32) -> Result<(), NETIO_STATUS> { diff --git a/easytier/src/common/ifcfg/win/mod.rs b/easytier/src/common/ifcfg/win/mod.rs index 5c850f46..a24fe627 100644 --- a/easytier/src/common/ifcfg/win/mod.rs +++ b/easytier/src/common/ifcfg/win/mod.rs @@ -1,3 +1,2 @@ pub mod luid; -pub mod netsh; pub mod types; diff --git a/easytier/src/common/ifcfg/win/netsh.rs b/easytier/src/common/ifcfg/win/netsh.rs deleted file mode 100644 index 267d404f..00000000 --- a/easytier/src/common/ifcfg/win/netsh.rs +++ /dev/null @@ -1,118 +0,0 @@ -// -// Port supporting code for wireguard-nt from wireguard-windows v0.5.3 to Rust -// This file replicates the functionality of wireguard-windows/tunnel/winipcfg/netsh.go -// - -use std::{ - net::{Ipv4Addr, Ipv6Addr}, - process::{Command, Stdio}, -}; -use winapi::shared::ws2def::{ADDRESS_FAMILY, AF_INET, AF_INET6}; - -pub fn flush_dns(family: ADDRESS_FAMILY, if_index: u32) -> Result<(), String> { - let proto_name = match family as i32 { - AF_INET => "ipv4", - AF_INET6 => "ipv6", - _ => { - return Err(String::from("Invalid address family")); - } - }; - - //let netsh_params = format!("interface {proto} set dnsservers name={itf} source=static address=none validate=no register=both", proto=proto_name, itf=ip_itf.InterfaceIndex); - let ret_netsh = Command::new("netsh.exe") - .arg("interface") - .arg(proto_name) - .arg("set") - .arg("dnsservers") - .arg(format!("name={}", if_index)) - .arg("source=static") - .arg("address=none") - .arg("validate=no") - .arg("register=both") - .stdin(Stdio::piped()) - .stdout(Stdio::piped()) - .output(); - match ret_netsh { - Ok(output) => { - // netsh.exe returns error messages only and is silent upon success. BUT then it will return \r\n, so we need to look at the lines. - if let Ok(stdout_str) = String::from_utf8(output.stdout) { - if stdout_str.is_empty() || stdout_str == "\r\n" { - Ok(()) - } else { - // TODO: ignore "There are no Domain Name Servers (DNS) configured on this computer." - // Is this string localized? - Err(stdout_str) - } - } else { - Err(String::from("Could not parse netsh output")) - } - } - Err(_) => Err(String::from("Failed to execute command")), - } -} - -// Please execute flush_dns() first, as written in the original source code. -fn add_dns(family: ADDRESS_FAMILY, if_index: u32, dnses: &[String]) -> Result<(), String> { - let proto_name = match family as i32 { - AF_INET => "ipv4", - AF_INET6 => "ipv6", - _ => { - return Err(String::from("Invalid address family")); - } - }; - - // "interface ipv4 add dnsservers name=%d address=%s validate=no" - let ret_netsh = Command::new("netsh.exe") - .arg("interface") - .arg(proto_name) - .arg("add") - .arg("dnsservers") - .arg(format!("name={}", if_index)) - .arg(format!( - "address={}", - dnses - .iter() - .map(|x| x.to_string()) - .collect::>() - .join(",") - )) - .arg("validate=no") - .stdin(Stdio::piped()) - .stdout(Stdio::piped()) - .output(); - match ret_netsh { - Ok(output) => { - // netsh.exe returns error messages only and is silent upon success. BUT then it will return \r\n, so we need to look at the lines. - if let Ok(stdout_str) = String::from_utf8(output.stdout) { - if stdout_str.is_empty() || stdout_str == "\r\n" { - Ok(()) - } else { - // TODO: ignore "There are no Domain Name Servers (DNS) configured on this computer." - // Is this string localized? - Err(stdout_str) - } - } else { - Err(String::from("Could not parse netsh output")) - } - } - Err(_) => Err(String::from("Failed to execute command")), - } -} - -pub fn add_dns_ipv4(if_index: u32, dnses: &[Ipv4Addr]) -> Result<(), String> { - flush_dns(AF_INET as _, if_index)?; - if dnses.is_empty() { - return Ok(()); - } - let dnses_str: Vec = dnses.iter().map(|addr| addr.to_string()).collect(); - add_dns(AF_INET as _, if_index, &dnses_str) -} - -pub fn add_dns_ipv6(if_index: u32, dnses: &[Ipv6Addr]) -> Result<(), String> { - flush_dns(AF_INET6 as _, if_index)?; - if dnses.is_empty() { - return Ok(()); - } - let dnses_str: Vec = dnses.iter().map(|addr| addr.to_string()).collect(); - add_dns(AF_INET6 as _, if_index, &dnses_str) -} diff --git a/easytier/src/common/ifcfg/win/types.rs b/easytier/src/common/ifcfg/win/types.rs index dae0be94..1e54702f 100644 --- a/easytier/src/common/ifcfg/win/types.rs +++ b/easytier/src/common/ifcfg/win/types.rs @@ -4,9 +4,7 @@ // use cidr::{Ipv4Inet, Ipv6Inet}; -use std::ffi::OsString; use std::net::{Ipv4Addr, Ipv6Addr}; -use std::os::windows::prelude::*; use winapi::shared::ws2def::*; use winapi::shared::ws2ipdef::*; @@ -66,37 +64,3 @@ pub fn convert_ipv6addr_to_sockaddr(ip: &Ipv6Addr) -> SOCKADDR_IN6 { ..Default::default() } } - -/// This function converts winapi::shared::ws2def::SOCKADDR_IN to std::net::Ipv4Addr -pub fn convert_sockaddr_to_ipv4addr(sockaddr: &SOCKADDR_IN) -> Ipv4Addr { - unsafe { - Ipv4Addr::new( - sockaddr.sin_addr.S_un.S_un_b().s_b1, - sockaddr.sin_addr.S_un.S_un_b().s_b2, - sockaddr.sin_addr.S_un.S_un_b().s_b3, - sockaddr.sin_addr.S_un.S_un_b().s_b4, - ) - } -} - -/// This function converts a null-terminated Windows Unicode PWCHAR/LPWSTR to an OsString -pub fn u16_ptr_to_osstring(ptr: *const u16) -> OsString { - assert!(!ptr.is_null()); - let len = (0..) - .take_while(|&i| unsafe { *ptr.offset(i) } != 0) - .count(); - let slice = unsafe { std::slice::from_raw_parts(ptr, len) }; - - OsString::from_wide(slice) -} - -/// This function converts a null-terminated Windows PWCHAR/LPWSTR to a String -pub fn u16_ptr_to_string(ptr: *const u16) -> String { - assert!(!ptr.is_null()); - let len = (0..) - .take_while(|&i| unsafe { *ptr.offset(i) } != 0) - .count(); - let slice = unsafe { std::slice::from_raw_parts(ptr, len) }; - - String::from_utf16_lossy(slice) -} diff --git a/easytier/src/common/log.rs b/easytier/src/common/log.rs deleted file mode 100644 index cfab48d9..00000000 --- a/easytier/src/common/log.rs +++ /dev/null @@ -1,501 +0,0 @@ -use crate::common::config::{FileLoggerConfig, LoggingConfigLoader}; -use crate::common::get_logger_timer_rfc3339; -use crate::common::tracing_rolling_appender::{FileAppenderWrapper, RollingFileAppenderBase}; -use crate::rpc_service::logger::{CURRENT_LOG_LEVEL, LOGGER_LEVEL_SENDER}; -use anyhow::Context; -use paste::paste; -use std::io::IsTerminal; -use tracing::level_filters::LevelFilter; -use tracing::{Level, Metadata}; -use tracing_subscriber::Registry; -use tracing_subscriber::filter::{FilterExt, filter_fn}; -use tracing_subscriber::fmt::format::FmtSpan; -use tracing_subscriber::fmt::layer; -use tracing_subscriber::layer::SubscriberExt; -use tracing_subscriber::util::SubscriberInitExt; -use tracing_subscriber::{EnvFilter, Layer}; - -macro_rules! __log__ { - (const $var:ident = $target:expr) => { - const $var: &'static str = $target; - __log__!(@impl $target, $); - }; - - (@impl $target:expr, $_:tt) => { - __log__!(@impl $_, $target, error, warn, info, debug, trace); - }; - - (@impl $_:tt, $target:expr, $($lvl:ident),+) => { - paste! { - $( - macro_rules! [< __ $lvl __ >] { - (category: $cat:expr, $_ ($arg:tt)+) => { - tracing::$lvl!(target: concat!($target, "::", $cat), $_ ($arg)+) - }; - ($_ ($arg:tt)+) => { - tracing::$lvl!(target: $target, $_ ($arg)+) - }; - } - - #[allow(unused_imports)] - pub(crate) use [< __ $lvl __ >] as $lvl; - )+ - } - }; -} - -__log__!(const LOG_TARGET = "CORE"); - -fn parse_env_filter(default_level: Option) -> Result { - let directive = match default_level { - Some(level) => level.into(), - None => format!("{LOG_TARGET}=info").parse()?, - }; - - EnvFilter::builder() - .with_default_directive(directive) - .from_env() - .with_context(|| "failed to create env filter") -} - -fn parse_static_filter(level: LevelFilter) -> Result { - EnvFilter::builder() - .with_default_directive(level.into()) - .parse("") - .with_context(|| "failed to create static filter") -} - -fn parse_file_filter(level: LevelFilter) -> Result { - if matches!(level, LevelFilter::OFF) { - parse_static_filter(level) - } else { - parse_env_filter(Some(level)) - } -} - -fn is_log(meta: &Metadata) -> bool { - meta.target() == LOG_TARGET || meta.target().starts_with(&format!("{LOG_TARGET}::")) -} - -pub type NewFilterSender = std::sync::mpsc::Sender; - -macro_rules! tracing_layer { - ($layer:expr) => { - $layer.with_filter(filter_fn(is_log).not()).boxed() - }; -} - -macro_rules! log_layer { - ($layer:expr) => { - $layer - .with_file(false) - .with_line_number(false) - .with_filter(filter_fn(is_log)) - .boxed() - }; -} - -pub fn init( - config: impl LoggingConfigLoader, - reload: bool, -) -> Result, anyhow::Error> { - let mut layers = Vec::new(); - - let console_layers = console_layers( - config - .get_console_logger_config() - .level - .map(|s| s.parse().unwrap()), - )?; - layers.extend(console_layers); - - let sender = if cfg!(not(test)) { - let (file_layers, sender) = file_layers(config.get_file_logger_config(), reload)?; - layers.extend(file_layers); - sender - } else { - None - }; - - Registry::default() - .with(layers) - .try_init() - .map(|_| sender) - .map_err(Into::into) -} - -type BoxLayer = Box + Send + Sync>; - -fn console_layers(default_level: Option) -> anyhow::Result> { - let mut layers = Vec::new(); - if matches!(default_level, Some(LevelFilter::OFF)) { - return Ok(layers); - } - - let (console_filter, _) = - tracing_subscriber::reload::Layer::new(parse_env_filter(default_level)?); - - let (stdout, stderr) = cfg_select! { - test => {{ - let w = tracing_subscriber::fmt::TestWriter::new; - (w, w) - }} - _ => (std::io::stdout, std::io::stderr), - }; - - let ansi = std::io::stderr().is_terminal() || cfg!(test); - - let layer = || { - layer() - .compact() - .with_timer(get_logger_timer_rfc3339()) - .with_ansi(ansi) - .with_span_events(FmtSpan::NEW | FmtSpan::CLOSE) - .with_writer(stderr) - }; - - layers.push( - vec![ - tracing_layer!(layer()), - log_layer!(layer()).with_filter(LevelFilter::WARN).boxed(), - log_layer!(layer().with_writer(stdout)) - .with_filter(filter_fn(|metadata| *metadata.level() > Level::WARN)) - .boxed(), - ] - .with_filter(console_filter) - .boxed(), - ); - - #[cfg(feature = "tracing")] - { - layers.push(console_subscriber::ConsoleLayer::builder().spawn().boxed()); - } - - Ok(layers) -} - -fn file_layers( - config: FileLoggerConfig, - reload: bool, -) -> anyhow::Result<(Vec, Option)> { - let mut layers = Vec::new(); - - let level = config - .level - .map(|s| s.parse().unwrap()) - .unwrap_or(LevelFilter::OFF); - - if matches!(level, LevelFilter::OFF) && !reload { - return Ok((layers, None)); - } - - let (file_filter, file_filter_reloader) = - tracing_subscriber::reload::Layer::<_, Registry>::new(parse_file_filter(level)?); - - let layer = |wrapper| { - layer() - .with_ansi(false) - .with_writer(wrapper) - .with_timer(get_logger_timer_rfc3339()) - }; - - let wrapper = { - let path = { - let dir = config.dir.as_deref().unwrap_or("."); - let file = config.file.as_deref().unwrap_or("easytier.log"); - let path = std::path::Path::new(dir).join(file); - path.to_string_lossy().into_owned() - }; - - let builder = RollingFileAppenderBase::builder(); - let file_appender = builder - .filename(path) - .condition_daily() - .max_filecount(config.count.unwrap_or(10)) - .condition_max_file_size(config.size_mb.unwrap_or(100) * 1024 * 1024) - .build() - .with_context(|| "failed to initialize rolling file appender")?; - - FileAppenderWrapper::new(file_appender) - }; - - layers.push( - vec![ - tracing_layer!(layer(wrapper.clone())), - log_layer!(layer(wrapper.clone())), - ] - .with_filter(file_filter) - .boxed(), - ); - - if !reload { - return Ok((layers, None)); - } - - let (tx, rx) = std::sync::mpsc::channel(); - - // 初始化全局状态 - let _ = LOGGER_LEVEL_SENDER.set(std::sync::Mutex::new(tx.clone())); - let _ = CURRENT_LOG_LEVEL.set(std::sync::Mutex::new(level.to_string())); - - std::thread::spawn(move || { - while let Ok(lf) = rx.recv() { - let parsed_level = match lf.parse::() { - Ok(level) => level, - Err(e) => { - error!("Failed to parse new log level {:?}: {}", lf, e); - continue; - } - }; - - let mut new_filter = match parse_file_filter(parsed_level) { - Ok(filter) => Some(filter), - Err(e) => { - error!("Failed to build new log filter for {:?}: {:?}", lf, e); - continue; - } - }; - - match file_filter_reloader.modify(|f| { - *f = new_filter - .take() - .expect("log filter reloader only applies one filter per reload"); - }) { - Ok(()) => { - info!("Reload log filter succeed, new filter level: {:?}", lf); - } - Err(e) => { - error!("Failed to reload log filter: {:?}", e); - } - } - } - info!("Stop log filter reloader"); - }); - - Ok((layers, Some(tx))) -} - -#[cfg(test)] -mod tests { - use super::*; - use crate::common::config::FileLoggerConfig; - - const RUST_LOG: &str = "RUST_LOG"; - - struct EnvVarGuard { - key: &'static str, - previous: Option, - } - - impl EnvVarGuard { - fn set(key: &'static str, value: &str) -> Self { - let previous = std::env::var_os(key); - unsafe { std::env::set_var(key, value) }; - Self { key, previous } - } - - fn unset(key: &'static str) -> Self { - let previous = std::env::var_os(key); - unsafe { std::env::remove_var(key) }; - Self { key, previous } - } - } - - impl Drop for EnvVarGuard { - fn drop(&mut self) { - match &self.previous { - Some(value) => unsafe { std::env::set_var(self.key, value) }, - None => unsafe { std::env::remove_var(self.key) }, - } - } - } - - #[ctor::ctor] - fn init() { - let _ = Registry::default() - .with(console_layers(Some(LevelFilter::WARN)).unwrap()) - .try_init(); - } - - #[test] - fn default_file_logger_level_is_off_without_reload() { - let (layers, sender) = file_layers(FileLoggerConfig::default(), false).unwrap(); - - assert!(layers.is_empty()); - assert!(sender.is_none()); - } - - #[test] - #[serial_test::serial] - fn default_file_logger_level_filters_info_with_reload() { - let _guard = EnvVarGuard::set(RUST_LOG, "info"); - let temp_dir = tempfile::tempdir().unwrap(); - let log_file_name = "default-off-test.log".to_string(); - let log_path = temp_dir.path().join(&log_file_name); - - let cfg = FileLoggerConfig { - file: Some(log_file_name), - dir: Some(temp_dir.path().to_string_lossy().to_string()), - ..Default::default() - }; - - let (layers, _sender) = file_layers(cfg, true).unwrap(); - let marker = "default-file-logger-off-marker"; - let subscriber = Registry::default().with(layers); - - tracing::subscriber::with_default(subscriber, || { - tracing::info!(target: LOG_TARGET, "{}", marker); - std::thread::sleep(std::time::Duration::from_millis(300)); - }); - - let content = std::fs::read_to_string(&log_path).unwrap_or_default(); - assert!( - !content.contains(marker), - "default file logger level should filter info logs" - ); - } - - #[test] - #[serial_test::serial] - fn file_logger_level_uses_env_filter_when_enabled() { - let _guard = EnvVarGuard::set(RUST_LOG, "debug"); - let temp_dir = tempfile::tempdir().unwrap(); - let log_file_name = "env-filter-test.log".to_string(); - let log_path = temp_dir.path().join(&log_file_name); - - let cfg = FileLoggerConfig { - level: Some(LevelFilter::INFO.to_string()), - file: Some(log_file_name), - dir: Some(temp_dir.path().to_string_lossy().to_string()), - ..Default::default() - }; - - let (layers, _sender) = file_layers(cfg, true).unwrap(); - let marker = "file-logger-env-filter-marker"; - let subscriber = Registry::default().with(layers); - - tracing::subscriber::with_default(subscriber, || { - tracing::debug!(target: LOG_TARGET, "{}", marker); - std::thread::sleep(std::time::Duration::from_millis(300)); - }); - - let content = std::fs::read_to_string(&log_path).unwrap_or_default(); - assert!( - content.contains(marker), - "enabled file logger should use RUST_LOG directives" - ); - } - - #[test] - #[serial_test::serial] - fn file_logger_reload_uses_env_filter_when_enabled() { - let _guard = EnvVarGuard::set(RUST_LOG, "debug"); - let temp_dir = tempfile::tempdir().unwrap(); - let log_file_name = "reload-env-filter-test.log".to_string(); - let log_path = temp_dir.path().join(&log_file_name); - - let cfg = FileLoggerConfig { - file: Some(log_file_name), - dir: Some(temp_dir.path().to_string_lossy().to_string()), - ..Default::default() - }; - - let (layers, sender) = file_layers(cfg, true).unwrap(); - let sender = sender.expect("reload=true should return a sender"); - let marker = "file-logger-reload-env-filter-marker"; - let subscriber = Registry::default().with(layers); - - tracing::subscriber::with_default(subscriber, || { - sender.send(LevelFilter::INFO.to_string()).unwrap(); - std::thread::sleep(std::time::Duration::from_millis(300)); - - tracing::debug!(target: LOG_TARGET, "{}", marker); - std::thread::sleep(std::time::Duration::from_millis(300)); - }); - - let content = std::fs::read_to_string(&log_path).unwrap_or_default(); - assert!( - content.contains(marker), - "file logger enabled by reload should use RUST_LOG directives" - ); - } - - #[test] - #[serial_test::serial] - fn file_logger_reload_off_ignores_env_filter() { - let _guard = EnvVarGuard::set(RUST_LOG, "info"); - let temp_dir = tempfile::tempdir().unwrap(); - let log_file_name = "reload-off-test.log".to_string(); - let log_path = temp_dir.path().join(&log_file_name); - - let cfg = FileLoggerConfig { - level: Some(LevelFilter::INFO.to_string()), - file: Some(log_file_name), - dir: Some(temp_dir.path().to_string_lossy().to_string()), - ..Default::default() - }; - - let (layers, sender) = file_layers(cfg, true).unwrap(); - let sender = sender.expect("reload=true should return a sender"); - let marker = "file-logger-reload-off-marker"; - let subscriber = Registry::default().with(layers); - - tracing::subscriber::with_default(subscriber, || { - sender.send(LevelFilter::OFF.to_string()).unwrap(); - std::thread::sleep(std::time::Duration::from_millis(300)); - - tracing::info!(target: LOG_TARGET, "{}", marker); - std::thread::sleep(std::time::Duration::from_millis(300)); - }); - - let content = std::fs::read_to_string(&log_path).unwrap_or_default(); - assert!( - !content.contains(marker), - "disabled file logger should ignore RUST_LOG directives" - ); - } - - #[test] - #[serial_test::serial] - fn test_logger_reload() { - let _guard = EnvVarGuard::unset(RUST_LOG); - let temp_dir = tempfile::tempdir().unwrap(); - let log_file_name = "reload-test.log".to_string(); - let log_path = temp_dir.path().join(&log_file_name); - - let cfg = FileLoggerConfig { - level: Some(LevelFilter::INFO.to_string()), - file: Some(log_file_name), - dir: Some(temp_dir.path().to_string_lossy().to_string()), - size_mb: Some(10), - count: Some(1), - }; - - let (layers, sender) = file_layers(cfg, true).unwrap(); - let sender = sender.expect("reload=true should return a sender"); - - let before_marker = "reload-before-debug-marker"; - let after_marker = "reload-after-debug-marker"; - let subscriber = Registry::default().with(layers); - - tracing::subscriber::with_default(subscriber, || { - tracing::debug!("{}", before_marker); - - sender.send(LevelFilter::DEBUG.to_string()).unwrap(); - std::thread::sleep(std::time::Duration::from_millis(300)); - - tracing::debug!("{}", after_marker); - std::thread::sleep(std::time::Duration::from_millis(300)); - }); - - let content = std::fs::read_to_string(&log_path).unwrap_or_default(); - assert!( - !content.contains(before_marker), - "debug log should be filtered before reload" - ); - assert!( - content.contains(after_marker), - "debug log should be visible after reload" - ); - } -} diff --git a/easytier/src/common/log/file.rs b/easytier/src/common/log/file.rs new file mode 100644 index 00000000..fadc2f14 --- /dev/null +++ b/easytier/src/common/log/file.rs @@ -0,0 +1,221 @@ +use std::sync::atomic::AtomicU8; + +#[cfg(feature = "management")] +use anyhow::Context as _; +use log::{Level, LevelFilter}; +#[cfg(feature = "management")] +use std::sync::atomic::Ordering; + +#[cfg(feature = "management")] +use crate::common::{ + config::FileLoggerConfig, + tracing_rolling_appender::{FileAppenderWrapper, RollingFileAppenderBase}, +}; + +use super::level_rank; +#[cfg(feature = "management")] +use super::{TargetFilter, format_line, level_is_enabled, parse_level}; + +#[cfg(feature = "management")] +pub(super) struct FileSink { + output: Option, + reload: bool, +} + +#[cfg(not(feature = "management"))] +pub(super) struct FileSink; + +#[cfg(feature = "management")] +struct FileOutput { + appender: FileAppenderWrapper, + max_level: AtomicU8, + state: parking_lot::RwLock, +} + +#[cfg(feature = "management")] +struct FileState { + level: LevelFilter, + filter: TargetFilter, +} + +#[cfg(feature = "management")] +impl FileSink { + pub(super) fn disabled() -> Self { + Self { + output: None, + reload: false, + } + } + + pub(super) fn from_config(config: FileLoggerConfig, reload: bool) -> anyhow::Result { + let level = config + .level + .as_deref() + .map(parse_level) + .transpose() + .context("invalid file log level")? + .unwrap_or(LevelFilter::Off); + if level == LevelFilter::Off && !reload { + return Ok(Self::disabled()); + } + + let dir = config.dir.as_deref().unwrap_or("."); + let file = config.file.as_deref().unwrap_or("easytier.log"); + let path = std::path::Path::new(dir).join(file); + let file_appender = RollingFileAppenderBase::builder() + .filename(path.to_string_lossy().into_owned()) + .condition_daily() + .max_filecount(config.count.unwrap_or(10)) + .condition_max_file_size(config.size_mb.unwrap_or(100) * 1024 * 1024) + .build() + .context("failed to initialize rolling file appender")?; + + let filter = file_filter(level)?; + let max_level = AtomicU8::new(level_rank(filter.max_level())); + Ok(Self { + output: Some(FileOutput { + appender: FileAppenderWrapper::new(file_appender), + max_level, + state: parking_lot::RwLock::new(FileState { level, filter }), + }), + reload, + }) + } + + pub(super) fn enabled(&self, target: &str, level: Level) -> bool { + self.output.as_ref().is_some_and(|output| { + level_is_enabled(output.max_level.load(Ordering::Acquire), level) + && output.state.read().filter.enabled(target, level) + }) + } + + pub(super) fn max_level(&self) -> LevelFilter { + self.output + .as_ref() + .map(|output| output.state.read().filter.max_level()) + .unwrap_or(LevelFilter::Off) + } + + pub(super) fn max_level_rank(&self) -> u8 { + self.output + .as_ref() + .map(|output| output.max_level.load(Ordering::Relaxed)) + .unwrap_or_else(|| level_rank(LevelFilter::Off)) + } + + pub(super) fn dynamic(&self) -> bool { + self.reload && self.output.is_some() + } + + pub(super) fn emit(&self, timestamp: &str, level: Level, target: &str, message: &str) { + if let Some(output) = &self.output { + let line = format_line(timestamp, level, target, message, false); + let _ = output.appender.write_all(line.as_bytes()); + } + } + + pub(super) fn set_level( + &self, + level: LevelFilter, + console_max_level: LevelFilter, + active_max_level: &AtomicU8, + ) -> anyhow::Result<()> { + let output = self + .output + .as_ref() + .filter(|_| self.reload) + .context("logger reloader is not initialized")?; + let filter = file_filter(level)?; + let file_max_filter = filter.max_level(); + let file_max_level = level_rank(file_max_filter); + let active_max_filter = console_max_level.max(file_max_filter); + + let mut state = output.state.write(); + *state = FileState { level, filter }; + output.max_level.store(file_max_level, Ordering::Release); + active_max_level.store(level_rank(active_max_filter), Ordering::Release); + log::set_max_level(active_max_filter); + drop(state); + Ok(()) + } + + pub(super) fn level(&self) -> LevelFilter { + self.output + .as_ref() + .map(|output| output.state.read().level) + .unwrap_or(LevelFilter::Info) + } + + pub(super) fn flush(&self) { + if let Some(output) = &self.output { + let _ = output.appender.flush(); + } + } + + #[cfg(test)] + pub(super) fn is_open(&self) -> bool { + self.output.is_some() + } + + #[cfg(test)] + pub(super) fn assert_level_state(&self, console_max_level: u8, active_max_level: u8) { + let output = self.output.as_ref().expect("file logger is open"); + let state = output.state.read(); + let file_max_level = level_rank(state.filter.max_level()); + assert_eq!(output.max_level.load(Ordering::Relaxed), file_max_level); + assert_eq!(active_max_level, console_max_level.max(file_max_level)); + assert_eq!( + level_rank(log::max_level()), + console_max_level.max(file_max_level) + ); + } +} + +#[cfg(not(feature = "management"))] +impl FileSink { + pub(super) fn disabled() -> Self { + Self + } + + pub(super) fn enabled(&self, _target: &str, _level: Level) -> bool { + false + } + + pub(super) fn max_level(&self) -> LevelFilter { + LevelFilter::Off + } + + pub(super) fn max_level_rank(&self) -> u8 { + level_rank(LevelFilter::Off) + } + + pub(super) fn dynamic(&self) -> bool { + false + } + + pub(super) fn emit(&self, _timestamp: &str, _level: Level, _target: &str, _message: &str) {} + + pub(super) fn set_level( + &self, + _level: LevelFilter, + _console_max_level: LevelFilter, + _active_max_level: &AtomicU8, + ) -> anyhow::Result<()> { + anyhow::bail!("file logging is not available in this build") + } + + pub(super) fn level(&self) -> LevelFilter { + LevelFilter::Info + } + + pub(super) fn flush(&self) {} +} + +#[cfg(feature = "management")] +fn file_filter(level: LevelFilter) -> anyhow::Result { + if level == LevelFilter::Off { + Ok(TargetFilter::off()) + } else { + TargetFilter::from_environment(TargetFilter::with_default(level)) + } +} diff --git a/easytier/src/common/log/management.rs b/easytier/src/common/log/management.rs new file mode 100644 index 00000000..3ef39c9b --- /dev/null +++ b/easytier/src/common/log/management.rs @@ -0,0 +1,19 @@ +use anyhow::Context as _; + +use crate::common::config::LoggingConfigLoader; + +use super::{FileSink, Logger, TargetFilter, install, parse_level}; + +pub fn init(config: impl LoggingConfigLoader, reload: bool) -> anyhow::Result<()> { + let console_config = config.get_console_logger_config(); + let console_level = console_config + .level + .as_deref() + .map(parse_level) + .transpose() + .context("invalid console log level")?; + let console = TargetFilter::console(console_level)?; + let file = FileSink::from_config(config.get_file_logger_config(), reload)?; + + install(Logger::new(console, file)) +} diff --git a/easytier/src/common/log/mod.rs b/easytier/src/common/log/mod.rs new file mode 100644 index 00000000..e918f055 --- /dev/null +++ b/easytier/src/common/log/mod.rs @@ -0,0 +1,638 @@ +use std::{ + fmt::{self, Write as _}, + io::{self, IsTerminal, Write as _}, + sync::{ + OnceLock, + atomic::{AtomicU8, Ordering}, + }, + time::{SystemTime, UNIX_EPOCH}, +}; + +use anyhow::Context as _; +use log::{Level, LevelFilter, Metadata as LogMetadata, Record as LogRecord}; +use paste::paste; +use tracing::{ + Event, + field::{Field, Visit}, +}; + +mod file; +#[cfg(feature = "management")] +mod management; +#[cfg(feature = "management")] +pub use management::init; +mod tracing_backend; + +use file::FileSink; + +macro_rules! __log__ { + (const $var:ident = $target:expr) => { + const $var: &'static str = $target; + __log__!(@impl $target, $); + }; + + (@impl $target:expr, $_:tt) => { + __log__!(@impl $_, $target, error, warn, info, debug, trace); + }; + + (@impl $_:tt, $target:expr, $($lvl:ident),+) => { + paste! { + $( + macro_rules! [< __ $lvl __ >] { + (category: $cat:expr, $_ ($arg:tt)+) => { + tracing::$lvl!(target: concat!($target, "::", $cat), $_ ($arg)+) + }; + ($_ ($arg:tt)+) => { + tracing::$lvl!(target: $target, $_ ($arg)+) + }; + } + + #[allow(unused_imports)] + pub(crate) use [< __ $lvl __ >] as $lvl; + )+ + } + }; +} + +__log__!(const LOG_TARGET = "CORE"); + +static LOGGER: OnceLock = OnceLock::new(); + +pub fn init_console() -> anyhow::Result<()> { + install(Logger::new( + TargetFilter::console(Some(LevelFilter::Info))?, + FileSink::disabled(), + )) +} + +pub fn set_file_level(level: &str) -> anyhow::Result<()> { + let level = parse_level(level).context("invalid file log level")?; + LOGGER + .get() + .context("logger is not initialized")? + .set_file_level(level) +} + +pub fn file_level() -> String { + LOGGER + .get() + .map(Logger::file_level) + .unwrap_or(LevelFilter::Info) + .to_string() + .to_ascii_lowercase() +} + +fn install(logger: Logger) -> anyhow::Result<()> { + LOGGER + .set(logger) + .map_err(|_| anyhow::anyhow!("logger is already initialized"))?; + let logger = LOGGER.get().expect("logger was just initialized"); + + log::set_logger(logger).map_err(|_| anyhow::anyhow!("a log logger is already installed"))?; + log::set_max_level(logger.max_level()); + tracing_backend::install(logger).context("failed to install tracing subscriber") +} + +fn parse_level(level: &str) -> anyhow::Result { + level + .parse() + .map_err(|error| anyhow::anyhow!("{error}: {level:?}")) +} + +#[derive(Clone, Debug)] +struct TargetFilter { + default: LevelFilter, + targets: Vec<(Box, LevelFilter)>, +} + +impl TargetFilter { + fn console(level: Option) -> anyhow::Result { + if level == Some(LevelFilter::Off) { + return Ok(Self::off()); + } + + let fallback = match level { + Some(level) => Self::with_default(level), + None => Self { + default: LevelFilter::Off, + targets: vec![(LOG_TARGET.into(), LevelFilter::Info)], + }, + }; + Self::from_environment(fallback) + } + + fn off() -> Self { + Self::with_default(LevelFilter::Off) + } + + fn with_default(default: LevelFilter) -> Self { + Self { + default, + targets: Vec::new(), + } + } + + fn from_environment(fallback: Self) -> anyhow::Result { + let spec = std::env::var("RUST_LOG").unwrap_or_default(); + Self::parse(&spec).map(|filter| filter.unwrap_or(fallback)) + } + + fn parse(spec: &str) -> anyhow::Result> { + let mut filter = Self::off(); + let mut found = false; + + for directive in spec.split(',').map(str::trim).filter(|s| !s.is_empty()) { + found = true; + if directive.contains(['[', ']', '{', '}']) { + anyhow::bail!( + "span and field filters are not supported in RUST_LOG: {directive:?}" + ); + } + + if let Some((target, level)) = directive.rsplit_once('=') { + let target = target.trim(); + if target.is_empty() { + anyhow::bail!("missing target in RUST_LOG directive: {directive:?}"); + } + filter + .targets + .push((target.into(), parse_level(level.trim())?)); + } else if let Ok(level) = directive.parse() { + filter.default = level; + } else { + if directive.chars().any(char::is_whitespace) { + anyhow::bail!("invalid RUST_LOG directive: {directive:?}"); + } + filter.targets.push((directive.into(), LevelFilter::Trace)); + } + } + + Ok(found.then_some(filter)) + } + + fn enabled(&self, target: &str, level: Level) -> bool { + let mut selected = self.default; + let mut selected_len = 0; + for (prefix, filter) in &self.targets { + if prefix.len() >= selected_len && target.starts_with(prefix.as_ref()) { + selected = *filter; + selected_len = prefix.len(); + } + } + selected >= level.to_level_filter() + } + + fn max_level(&self) -> LevelFilter { + self.targets + .iter() + .map(|(_, level)| *level) + .fold(self.default, std::cmp::max) + } +} + +struct Logger { + console: TargetFilter, + console_max_level: u8, + active_max_level: AtomicU8, + color: bool, + file: FileSink, +} + +impl Logger { + fn new(console: TargetFilter, file: FileSink) -> Self { + let console_max_level = level_rank(console.max_level()); + let active_max_level = console_max_level.max(file.max_level_rank()); + Self { + console, + console_max_level, + active_max_level: AtomicU8::new(active_max_level), + color: io::stderr().is_terminal() && std::env::var_os("NO_COLOR").is_none(), + file, + } + } + + fn enabled(&self, target: &str, level: Level) -> bool { + level_is_enabled(self.active_max_level.load(Ordering::Acquire), level) + && (self.console_enabled(target, level) || self.file.enabled(target, level)) + } + + fn console_enabled(&self, target: &str, level: Level) -> bool { + level_is_enabled(self.console_max_level, level) && self.console.enabled(target, level) + } + + fn max_level(&self) -> LevelFilter { + self.console.max_level().max(self.file.max_level()) + } + + fn emit(&self, level: Level, target: &str, message: &str) { + let console_enabled = self.console_enabled(target, level); + let file_enabled = self.file.enabled(target, level); + if !console_enabled && !file_enabled { + return; + } + + let timestamp = timestamp_rfc3339_utc(); + if console_enabled { + let line = format_line(×tamp, level, target, message, self.color); + let _ = if matches!(level, Level::Error | Level::Warn) { + io::stderr().lock().write_all(line.as_bytes()) + } else { + io::stdout().lock().write_all(line.as_bytes()) + }; + } + + if file_enabled { + self.file.emit(×tamp, level, target, message); + } + } + + fn set_file_level(&self, level: LevelFilter) -> anyhow::Result<()> { + self.file + .set_level(level, self.console.max_level(), &self.active_max_level) + } + + fn file_level(&self) -> LevelFilter { + self.file.level() + } + + fn flush_file(&self) { + self.file.flush(); + } +} + +fn level_rank(level: LevelFilter) -> u8 { + match level { + LevelFilter::Off => 0, + LevelFilter::Error => 1, + LevelFilter::Warn => 2, + LevelFilter::Info => 3, + LevelFilter::Debug => 4, + LevelFilter::Trace => 5, + } +} + +fn level_is_enabled(max_level: u8, level: Level) -> bool { + max_level >= level_rank(level.to_level_filter()) +} + +impl log::Log for Logger { + fn enabled(&self, metadata: &LogMetadata<'_>) -> bool { + self.enabled(metadata.target(), metadata.level()) + } + + fn log(&self, record: &LogRecord<'_>) { + if self.enabled(record.target(), record.level()) { + self.emit(record.level(), record.target(), &record.args().to_string()); + } + } + + fn flush(&self) { + self.flush_file(); + } +} + +fn format_line(timestamp: &str, level: Level, target: &str, message: &str, color: bool) -> String { + let mut line = String::with_capacity(timestamp.len() + target.len() + message.len() + 32); + if color { + let color = match level { + Level::Error => "\x1b[31m", + Level::Warn => "\x1b[33m", + Level::Info => "\x1b[32m", + Level::Debug => "\x1b[34m", + Level::Trace => "\x1b[90m", + }; + let _ = writeln!( + line, + "{timestamp} {color}{level:<5}\x1b[0m {target}: {message}" + ); + } else { + let _ = writeln!(line, "{timestamp} {level:<5} {target}: {message}"); + } + line +} + +fn timestamp_rfc3339_utc() -> String { + let duration = SystemTime::now() + .duration_since(UNIX_EPOCH) + .unwrap_or_default(); + let seconds = duration.as_secs(); + let seconds_of_day = seconds % 86_400; + let (year, month, day) = civil_date_from_unix_days((seconds / 86_400) as i64); + let hour = seconds_of_day / 3_600; + let minute = seconds_of_day % 3_600 / 60; + let second = seconds_of_day % 60; + + format!( + "{year:04}-{month:02}-{day:02}T{hour:02}:{minute:02}:{second:02}.{:03}Z", + duration.subsec_millis() + ) +} + +fn civil_date_from_unix_days(days: i64) -> (i64, i64, i64) { + let days = days + 719_468; + let era = if days >= 0 { days } else { days - 146_096 } / 146_097; + let day_of_era = days - era * 146_097; + let year_of_era = + (day_of_era - day_of_era / 1_460 + day_of_era / 36_524 - day_of_era / 146_096) / 365; + let mut year = year_of_era + era * 400; + let day_of_year = day_of_era - (365 * year_of_era + year_of_era / 4 - year_of_era / 100); + let month_prime = (5 * day_of_year + 2) / 153; + let day = day_of_year - (153 * month_prime + 2) / 5 + 1; + let month = month_prime + if month_prime < 10 { 3 } else { -9 }; + year += i64::from(month <= 2); + (year, month, day) +} + +fn tracing_level(level: &tracing::Level) -> Level { + match *level { + tracing::Level::ERROR => Level::Error, + tracing::Level::WARN => Level::Warn, + tracing::Level::INFO => Level::Info, + tracing::Level::DEBUG => Level::Debug, + tracing::Level::TRACE => Level::Trace, + } +} + +fn emit_event(logger: &Logger, event: &Event<'_>) { + let metadata = event.metadata(); + let level = tracing_level(metadata.level()); + if !logger.enabled(metadata.target(), level) { + return; + } + + let mut fields = EventFields::default(); + event.record(&mut fields); + logger.emit(level, metadata.target(), &fields.finish()); +} + +#[derive(Default)] +struct EventFields { + message: Option, + fields: String, +} + +impl EventFields { + fn write_field(&mut self, field: &Field, value: impl fmt::Display) { + if !self.fields.is_empty() { + self.fields.push(' '); + } + let _ = write!(self.fields, "{}={value}", field.name()); + } + + fn finish(self) -> String { + match (self.message, self.fields.is_empty()) { + (Some(message), false) => format!("{message} {}", self.fields), + (Some(message), true) => message, + (None, _) => self.fields, + } + } +} + +impl Visit for EventFields { + fn record_debug(&mut self, field: &Field, value: &dyn fmt::Debug) { + if field.name() == "message" { + self.message = Some(format!("{value:?}")); + } else { + self.write_field(field, format_args!("{value:?}")); + } + } + + fn record_str(&mut self, field: &Field, value: &str) { + if field.name() == "message" { + self.message = Some(value.to_owned()); + } else { + self.write_field(field, format_args!("{value:?}")); + } + } + + fn record_error(&mut self, field: &Field, value: &(dyn std::error::Error + 'static)) { + self.write_field(field, value); + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[cfg(feature = "management")] + use crate::common::config::FileLoggerConfig; + + struct EnvVarGuard { + previous: Option, + } + + impl EnvVarGuard { + fn set(value: Option<&str>) -> Self { + let previous = std::env::var_os("RUST_LOG"); + match value { + Some(value) => unsafe { std::env::set_var("RUST_LOG", value) }, + None => unsafe { std::env::remove_var("RUST_LOG") }, + } + Self { previous } + } + } + + impl Drop for EnvVarGuard { + fn drop(&mut self) { + match &self.previous { + Some(value) => unsafe { std::env::set_var("RUST_LOG", value) }, + None => unsafe { std::env::remove_var("RUST_LOG") }, + } + } + } + + #[test] + #[serial_test::serial] + fn default_console_only_enables_core_info() { + let _env = EnvVarGuard::set(None); + let filter = TargetFilter::console(None).unwrap(); + + assert!(filter.enabled("CORE::peer", Level::Info)); + assert!(!filter.enabled("CORE", Level::Debug)); + assert!(!filter.enabled("other", Level::Error)); + } + + #[test] + fn rust_log_supports_global_and_target_levels() { + let filter = TargetFilter::parse("warn,easytier_core=debug,hyper=off") + .unwrap() + .unwrap(); + + assert!(filter.enabled("easytier_core::peers", Level::Debug)); + assert!(!filter.enabled("hyper::client", Level::Error)); + assert!(!filter.enabled("other", Level::Info)); + assert!(filter.enabled("other", Level::Warn)); + } + + #[test] + fn rust_log_target_without_level_enables_trace() { + let filter = TargetFilter::parse("easytier_core").unwrap().unwrap(); + + assert!(filter.enabled("easytier_core::peer", Level::Trace)); + assert!(!filter.enabled("other", Level::Error)); + } + + #[test] + fn formatting_is_compact_and_color_is_optional() { + let plain = format_line( + "2026-07-22T12:00:00+08:00", + Level::Info, + "CORE", + "ready", + false, + ); + let colored = format_line( + "2026-07-22T12:00:00+08:00", + Level::Info, + "CORE", + "ready", + true, + ); + + assert_eq!(plain, "2026-07-22T12:00:00+08:00 INFO CORE: ready\n"); + assert!(colored.contains("\x1b[32mINFO \x1b[0m")); + } + + #[test] + fn unix_day_conversion_matches_known_dates() { + assert_eq!(civil_date_from_unix_days(0), (1970, 1, 1)); + assert_eq!(civil_date_from_unix_days(20_656), (2026, 7, 22)); + } + + #[test] + #[cfg(feature = "management")] + fn default_file_logger_is_not_opened_without_reload() { + let file = FileSink::from_config(FileLoggerConfig::default(), false).unwrap(); + assert!(!file.is_open()); + } + + #[test] + #[cfg(feature = "management")] + #[serial_test::serial] + fn file_logger_uses_rust_log_when_enabled() { + let _env = EnvVarGuard::set(Some("debug")); + let temp_dir = tempfile::tempdir().unwrap(); + let log_path = temp_dir.path().join("env-filter.log"); + let config = FileLoggerConfig { + level: Some("info".to_owned()), + file: Some("env-filter.log".to_owned()), + dir: Some(temp_dir.path().to_string_lossy().into_owned()), + ..Default::default() + }; + let file = FileSink::from_config(config, true).unwrap(); + let logger = Logger::new(TargetFilter::off(), file); + + logger.emit(Level::Debug, LOG_TARGET, "env-filter-marker"); + logger.flush_file(); + + let content = std::fs::read_to_string(log_path).unwrap(); + assert!(content.contains("env-filter-marker")); + } + + #[test] + #[cfg(feature = "management")] + #[serial_test::serial] + fn reloading_file_logger_preserves_rust_log_and_supports_off() { + let _env = EnvVarGuard::set(Some("debug")); + let temp_dir = tempfile::tempdir().unwrap(); + let log_path = temp_dir.path().join("reload.log"); + let config = FileLoggerConfig { + file: Some("reload.log".to_owned()), + dir: Some(temp_dir.path().to_string_lossy().into_owned()), + ..Default::default() + }; + let file = FileSink::from_config(config, true).unwrap(); + let logger = std::sync::Arc::new(Logger::new(TargetFilter::off(), file)); + + assert_eq!( + logger.active_max_level.load(Ordering::Relaxed), + level_rank(LevelFilter::Off) + ); + + let barrier = std::sync::Arc::new(std::sync::Barrier::new(9)); + let threads = (0..8) + .map(|thread_index| { + let logger = logger.clone(); + let barrier = barrier.clone(); + std::thread::spawn(move || { + barrier.wait(); + for iteration in 0..500 { + let level = if (thread_index + iteration) % 2 == 0 { + LevelFilter::Info + } else { + LevelFilter::Off + }; + logger.set_file_level(level).unwrap(); + } + }) + }) + .collect::>(); + barrier.wait(); + for thread in threads { + thread.join().unwrap(); + } + + logger.file.assert_level_state( + logger.console_max_level, + logger.active_max_level.load(Ordering::Relaxed), + ); + + logger.set_file_level(LevelFilter::Info).unwrap(); + assert_eq!( + logger.active_max_level.load(Ordering::Relaxed), + level_rank(LevelFilter::Debug) + ); + assert_eq!(logger.file.max_level_rank(), level_rank(LevelFilter::Debug)); + logger.emit(Level::Debug, LOG_TARGET, "enabled-by-env"); + logger.set_file_level(LevelFilter::Off).unwrap(); + assert_eq!( + logger.active_max_level.load(Ordering::Relaxed), + level_rank(LevelFilter::Off) + ); + logger.emit(Level::Error, LOG_TARGET, "disabled-despite-env"); + logger.flush_file(); + + let content = std::fs::read_to_string(log_path).unwrap(); + assert!(content.contains("enabled-by-env")); + assert!(!content.contains("disabled-despite-env")); + assert_eq!(logger.file_level(), LevelFilter::Off); + } + + #[test] + #[cfg(all(feature = "management", not(feature = "tracing")))] + #[serial_test::serial] + fn tracing_events_and_direct_log_records_share_the_file_sink() { + let _env = EnvVarGuard::set(None); + let temp_dir = tempfile::tempdir().unwrap(); + let log_path = temp_dir.path().join("shared-sink.log"); + let config = FileLoggerConfig { + level: Some("info".to_owned()), + file: Some("shared-sink.log".to_owned()), + dir: Some(temp_dir.path().to_string_lossy().into_owned()), + ..Default::default() + }; + let file = FileSink::from_config(config, false).unwrap(); + let logger = Box::leak(Box::new(Logger::new(TargetFilter::off(), file))); + let dispatch = tracing::Dispatch::new(tracing_backend::EventSubscriber::new(logger)); + + tracing::dispatcher::with_default(&dispatch, || { + let span = tracing::info_span!(target: LOG_TARGET, "ignored-span", peer = 7); + let _entered = span.enter(); + tracing::info!(target: LOG_TARGET, answer = 42, "tracing-event"); + }); + log::Log::log( + logger, + &LogRecord::builder() + .level(Level::Info) + .target("dependency") + .args(format_args!("direct-log-record")) + .build(), + ); + logger.flush_file(); + + let content = std::fs::read_to_string(log_path).unwrap(); + assert!(content.contains("tracing-event answer=42")); + assert!(content.contains("direct-log-record")); + assert!(!content.contains("ignored-span")); + } +} diff --git a/easytier/src/common/log/tracing_backend.rs b/easytier/src/common/log/tracing_backend.rs new file mode 100644 index 00000000..4d2842ff --- /dev/null +++ b/easytier/src/common/log/tracing_backend.rs @@ -0,0 +1,121 @@ +#[cfg(not(feature = "tracing"))] +use std::sync::atomic::{AtomicUsize, Ordering}; + +#[cfg(feature = "tracing")] +use tracing::{Event, subscriber::SetGlobalDefaultError}; +#[cfg(not(feature = "tracing"))] +use tracing::{ + Metadata, + span::{Attributes, Id, Record as SpanRecord}, + subscriber::{Interest, SetGlobalDefaultError}, +}; +#[cfg(feature = "tracing")] +use tracing_subscriber::prelude::*; + +#[cfg(not(feature = "tracing"))] +use super::tracing_level; +use super::{Logger, emit_event}; + +#[cfg(not(feature = "tracing"))] +pub(super) fn install(logger: &'static Logger) -> Result<(), SetGlobalDefaultError> { + tracing::subscriber::set_global_default(EventSubscriber::new(logger)) +} + +#[cfg(feature = "tracing")] +pub(super) fn install(logger: &'static Logger) -> Result<(), SetGlobalDefaultError> { + let subscriber = tracing_subscriber::registry() + .with(console_subscriber::ConsoleLayer::builder().spawn()) + .with(EventLayer(logger)); + tracing::subscriber::set_global_default(subscriber) +} + +#[cfg(feature = "tracing")] +struct EventLayer(&'static Logger); + +#[cfg(feature = "tracing")] +impl tracing_subscriber::Layer for EventLayer +where + S: tracing::Subscriber, +{ + fn on_event(&self, event: &Event<'_>, _ctx: tracing_subscriber::layer::Context<'_, S>) { + emit_event(self.0, event); + } +} + +#[cfg(not(feature = "tracing"))] +pub(super) struct EventSubscriber { + logger: &'static Logger, + next_span_id: AtomicUsize, +} + +#[cfg(not(feature = "tracing"))] +impl EventSubscriber { + pub(super) fn new(logger: &'static Logger) -> Self { + Self { + logger, + next_span_id: AtomicUsize::new(1), + } + } +} + +#[cfg(not(feature = "tracing"))] +impl tracing::Subscriber for EventSubscriber { + fn enabled(&self, metadata: &Metadata<'_>) -> bool { + metadata.is_event() + && self + .logger + .enabled(metadata.target(), tracing_level(metadata.level())) + } + + fn new_span(&self, _attributes: &Attributes<'_>) -> Id { + let id = self + .next_span_id + .fetch_update(Ordering::Relaxed, Ordering::Relaxed, |id| { + Some(if id == usize::MAX { 1 } else { id + 1 }) + }) + .expect("span id update always succeeds"); + Id::from_u64(id as u64) + } + + fn record(&self, _span: &Id, _values: &SpanRecord<'_>) {} + + fn record_follows_from(&self, _span: &Id, _follows: &Id) {} + + fn event(&self, event: &tracing::Event<'_>) { + emit_event(self.logger, event); + } + + fn enter(&self, _span: &Id) {} + + fn exit(&self, _span: &Id) {} + + fn register_callsite(&self, metadata: &'static Metadata<'static>) -> Interest { + if metadata.is_event() { + Interest::sometimes() + } else { + Interest::never() + } + } + + fn max_level_hint(&self) -> Option { + if self.logger.file.dynamic() { + return Some(tracing::level_filters::LevelFilter::TRACE); + } + Some(match self.logger.max_level() { + log::LevelFilter::Off => tracing::level_filters::LevelFilter::OFF, + log::LevelFilter::Error => tracing::level_filters::LevelFilter::ERROR, + log::LevelFilter::Warn => tracing::level_filters::LevelFilter::WARN, + log::LevelFilter::Info => tracing::level_filters::LevelFilter::INFO, + log::LevelFilter::Debug => tracing::level_filters::LevelFilter::DEBUG, + log::LevelFilter::Trace => tracing::level_filters::LevelFilter::TRACE, + }) + } + + fn clone_span(&self, id: &Id) -> Id { + id.clone() + } + + fn try_close(&self, _id: Id) -> bool { + true + } +} diff --git a/easytier/src/common/mod.rs b/easytier/src/common/mod.rs index e323465e..aa88b3ee 100644 --- a/easytier/src/common/mod.rs +++ b/easytier/src/common/mod.rs @@ -1,101 +1,28 @@ -use std::{ - fmt::Debug, - future, - sync::{Arc, Mutex}, -}; -use time::util::refresh_tz; -use tokio::{task::JoinSet, time::timeout}; -use tracing::Instrument; - -pub mod acl_processor; -pub mod compressor; pub mod config; pub mod constants; +#[cfg(feature = "management")] +pub mod credential_manager; pub mod dns; +#[cfg(feature = "management")] pub mod env_parser; pub mod error; pub mod global_ctx; -pub mod idn; pub mod ifcfg; +#[cfg(feature = "logging")] pub mod log; pub mod machine_id; pub mod netns; pub mod network; +#[cfg(feature = "management")] pub mod os_info; -pub mod stats_manager; pub mod stun; -pub mod stun_codec_ext; -pub mod token_bucket; +#[cfg(feature = "management")] pub mod tracing_rolling_appender; +#[cfg(feature = "upnp")] pub mod upnp; pub use machine_id::{MachineIdOptions, resolve_machine_id}; -pub fn get_logger_timer( - format: F, -) -> tracing_subscriber::fmt::time::OffsetTime { - refresh_tz(); - let local_offset = time::UtcOffset::current_local_offset() - .unwrap_or(time::UtcOffset::from_whole_seconds(0).unwrap()); - tracing_subscriber::fmt::time::OffsetTime::new(local_offset, format) -} - -pub fn get_logger_timer_rfc3339() --> tracing_subscriber::fmt::time::OffsetTime { - get_logger_timer(time::format_description::well_known::Rfc3339) -} - -pub type PeerId = u32; - -pub fn new_peer_id() -> PeerId { - rand::random() -} - -pub fn join_joinset_background( - js: Arc>>, - origin: String, -) { - let js = Arc::downgrade(&js); - let o = origin.clone(); - tokio::spawn( - async move { - while js.strong_count() > 0 { - tokio::time::sleep(std::time::Duration::from_secs(1)).await; - - let fut = future::poll_fn(|cx| { - let Some(js) = js.upgrade() else { - return std::task::Poll::Ready(()); - }; - - let mut js = js.lock().unwrap(); - while !js.is_empty() { - let ret = js.poll_join_next(cx); - match ret { - std::task::Poll::Ready(Some(_)) => { - continue; - } - std::task::Poll::Ready(None) => { - break; - } - std::task::Poll::Pending => { - return std::task::Poll::Pending; - } - } - } - std::task::Poll::Ready(()) - }); - - let _ = timeout(std::time::Duration::from_secs(5), fut).await; - } - tracing::debug!(?o, "joinset task exit"); - } - .instrument(tracing::info_span!( - "join_joinset_background", - origin = origin - )), - ); -} - pub fn shrink_dashmap( map: &dashmap::DashMap, threshold: Option, @@ -105,44 +32,3 @@ pub fn shrink_dashmap( map.shrink_to_fit(); } } - -#[cfg(test)] -mod tests { - use super::*; - - #[tokio::test] - async fn test_join_joinset_backgroud() { - let js = Arc::new(Mutex::new(JoinSet::<()>::new())); - join_joinset_background(js.clone(), "TEST".to_owned()); - js.try_lock().unwrap().spawn(async { - tokio::time::sleep(std::time::Duration::from_secs(1)).await; - }); - tokio::time::sleep(std::time::Duration::from_secs(2)).await; - assert!(js.try_lock().unwrap().is_empty()); - - for _ in 0..5 { - js.try_lock().unwrap().spawn(async { - tokio::time::sleep(std::time::Duration::from_secs(3)).await; - }); - tokio::task::yield_now().await; - } - - tokio::time::sleep(std::time::Duration::from_secs(2)).await; - - for _ in 0..5 { - js.try_lock().unwrap().spawn(async { - tokio::time::sleep(std::time::Duration::from_secs(1)).await; - }); - tokio::task::yield_now().await; - } - - tokio::time::sleep(std::time::Duration::from_secs(2)).await; - assert!(js.try_lock().unwrap().is_empty()); - - let weak_js = Arc::downgrade(&js); - drop(js); - tokio::time::sleep(std::time::Duration::from_secs(2)).await; - assert_eq!(weak_js.weak_count(), 0); - assert_eq!(weak_js.strong_count(), 0); - } -} diff --git a/easytier/src/common/netns.rs b/easytier/src/common/netns.rs index 25e38faf..a998e386 100644 --- a/easytier/src/common/netns.rs +++ b/easytier/src/common/netns.rs @@ -1,3 +1,4 @@ +use easytier_core::socket::SocketContext; use futures::Future; #[cfg(target_os = "linux")] @@ -84,6 +85,15 @@ impl NetNS { NetNS { name } } + pub fn from_socket_context(context: &SocketContext) -> Self { + Self::new( + context + .netns + .as_ref() + .map(|namespace| namespace.token().to_owned()), + ) + } + pub async fn run_async(&self, f: F) -> Ret where F: FnOnce() -> Fut, diff --git a/easytier/src/common/network.rs b/easytier/src/common/network.rs index 9097ad69..f48b8227 100644 --- a/easytier/src/common/network.rs +++ b/easytier/src/common/network.rs @@ -1,4 +1,5 @@ -use std::{net::IpAddr, ops::Deref, sync::Arc}; +#[cfg(target_os = "windows")] +use std::net::IpAddr; #[cfg(target_os = "windows")] use network_interface::{ @@ -7,16 +8,12 @@ use network_interface::{ use pnet::datalink::NetworkInterface; #[cfg(target_os = "windows")] use pnet::{ipnetwork::IpNetwork, util::MacAddr}; -use tokio::{ - sync::{Mutex, RwLock}, - task::JoinSet, -}; +#[cfg(all(target_os = "macos", not(feature = "macos-ne")))] +use tokio::sync::Mutex; use crate::proto::peer_rpc::GetIpListResponse; -use super::{netns::NetNS, stun::StunInfoCollectorTrait}; - -pub const CACHED_IP_LIST_TIMEOUT_SEC: u64 = 60; +use super::netns::NetNS; struct InterfaceFilter { iface: NetworkInterface, @@ -198,224 +195,215 @@ pub async fn local_ipv6() -> std::io::Result { } } -pub struct IPCollector { - cached_ip_list: Arc>, - collect_ip_task: Mutex>, - net_ns: NetNS, - stun_info_collector: Arc>, -} - -impl IPCollector { - pub fn new(net_ns: NetNS, stun_info_collector: T) -> Self { - Self { - cached_ip_list: Arc::new(RwLock::new(GetIpListResponse::default())), - collect_ip_task: Mutex::new(JoinSet::new()), - net_ns, - stun_info_collector: Arc::new(Box::new(stun_info_collector)), - } +pub(crate) async fn collect_interfaces(net_ns: NetNS, filter: bool) -> Vec { + #[cfg(target_os = "linux")] + { + run_in_namespace(net_ns, move || async move { + collect_interfaces_in_current_namespace(filter).await + }) + .await } - pub async fn collect_ip_addrs(&self) -> GetIpListResponse { - let mut task = self.collect_ip_task.lock().await; - if task.is_empty() { - let cached_ip_list = self.cached_ip_list.clone(); - *cached_ip_list.write().await = - Self::do_collect_local_ip_addrs(self.net_ns.clone()).await; - let net_ns = self.net_ns.clone(); - let stun_info_collector = self.stun_info_collector.clone(); - let cached_ip_list = self.cached_ip_list.clone(); - task.spawn(async move { - let mut last_fetch_iface_time = std::time::Instant::now(); - loop { - if last_fetch_iface_time.elapsed().as_secs() > CACHED_IP_LIST_TIMEOUT_SEC { - let ifaces = Self::do_collect_local_ip_addrs(net_ns.clone()).await; - *cached_ip_list.write().await = ifaces; - last_fetch_iface_time = std::time::Instant::now(); - } + #[cfg(not(target_os = "linux"))] + { + let _g = net_ns.guard(); + collect_interfaces_in_current_namespace(filter).await + } +} - let stun_info = stun_info_collector.get_stun_info(); - for ip in stun_info.public_ip.iter() { - let Ok(ip_addr) = ip.parse::() else { - continue; - }; +async fn collect_interfaces_in_current_namespace(filter: bool) -> Vec { + #[cfg(target_os = "windows")] + let ifaces = collect_interfaces_windows(); + #[cfg(not(target_os = "windows"))] + let ifaces = pnet::datalink::interfaces(); + let mut ret = vec![]; + for iface in ifaces { + let f = InterfaceFilter { + iface: iface.clone(), + }; - match ip_addr { - IpAddr::V4(v) => { - cached_ip_list.write().await.public_ipv4.replace(v.into()); - } - IpAddr::V6(v) => { - cached_ip_list.write().await.public_ipv6.replace(v.into()); - } - } - } + if filter && !f.filter_iface().await { + continue; + } - tracing::debug!( - "got public ip: {:?}, {:?}", - cached_ip_list.read().await.public_ipv4, - cached_ip_list.read().await.public_ipv6 + ret.push(iface); + } + + ret +} + +#[cfg(target_os = "linux")] +async fn run_in_namespace(net_ns: NetNS, operation: F) -> T +where + T: Send + 'static, + F: FnOnce() -> Fut + Send + 'static, + Fut: std::future::Future + 'static, +{ + tokio::task::spawn_blocking(move || { + let runtime = tokio::runtime::Builder::new_current_thread() + .enable_all() + .build() + .expect("build namespace-local runtime"); + net_ns.run(|| runtime.block_on(operation())) + }) + .await + .expect("namespace-local network operation panicked") +} + +#[cfg(target_os = "windows")] +fn collect_interfaces_windows() -> Vec { + match SystemNetworkInterface::show() { + Ok(ifaces) => ifaces.into_iter().map(convert_windows_interface).collect(), + Err(e) => { + tracing::warn!( + ?e, + "failed to enumerate interfaces via network-interface, falling back to pnet" + ); + match std::panic::catch_unwind(pnet::datalink::interfaces) { + Ok(ifaces) => ifaces, + Err(_) => { + tracing::error!( + "failed to enumerate interfaces via both network-interface and pnet" ); - - let sleep_sec = if cached_ip_list.read().await.public_ipv4.is_some() { - CACHED_IP_LIST_TIMEOUT_SEC - } else { - 3 - }; - tokio::time::sleep(std::time::Duration::from_secs(sleep_sec)).await; - } - }); - } - - self.cached_ip_list.read().await.deref().clone() - } - - pub async fn collect_interfaces(net_ns: NetNS, filter: bool) -> Vec { - let _g = net_ns.guard(); - #[cfg(target_os = "windows")] - let ifaces = Self::collect_interfaces_windows(); - #[cfg(not(target_os = "windows"))] - let ifaces = pnet::datalink::interfaces(); - let mut ret = vec![]; - for iface in ifaces { - let f = InterfaceFilter { - iface: iface.clone(), - }; - - if filter && !f.filter_iface().await { - continue; - } - - ret.push(iface); - } - - ret - } - - #[cfg(target_os = "windows")] - fn collect_interfaces_windows() -> Vec { - match SystemNetworkInterface::show() { - Ok(ifaces) => ifaces - .into_iter() - .map(Self::convert_windows_interface) - .collect(), - Err(e) => { - tracing::warn!( - ?e, - "failed to enumerate interfaces via network-interface, falling back to pnet" - ); - match std::panic::catch_unwind(pnet::datalink::interfaces) { - Ok(ifaces) => ifaces, - Err(_) => { - tracing::error!( - "failed to enumerate interfaces via both network-interface and pnet" - ); - Vec::new() - } + Vec::new() } } } } - - #[cfg(target_os = "windows")] - fn convert_windows_interface(iface: SystemNetworkInterface) -> NetworkInterface { - let mac = iface.mac_addr.as_deref().and_then(|mac| { - mac.parse::() - .map_err(|e| { - tracing::debug!(iface = %iface.name, mac, ?e, "failed to parse interface mac") - }) - .ok() - }); - - let ips = iface - .addr - .into_iter() - .filter_map(Self::convert_windows_interface_addr) - .collect(); - - NetworkInterface { - name: iface.name, - description: String::new(), - index: iface.index, - mac, - ips, - // pnet does not populate Windows flags either, so keep the existing semantics. - flags: 0, - } - } - - #[cfg(target_os = "windows")] - fn convert_windows_interface_addr(addr: SystemAddr) -> Option { - match addr { - SystemAddr::V4(addr) => { - let netmask = addr - .netmask - .map(IpAddr::V4) - .unwrap_or(IpAddr::V4(std::net::Ipv4Addr::new(255, 255, 255, 255))); - IpNetwork::with_netmask(IpAddr::V4(addr.ip), netmask) - .map_err(|e| { - tracing::debug!(ip = %addr.ip, ?addr.netmask, ?e, "failed to convert ipv4") - }) - .ok() - } - SystemAddr::V6(addr) => { - let netmask = addr - .netmask - .map(IpAddr::V6) - .unwrap_or(IpAddr::V6(std::net::Ipv6Addr::from(u128::MAX))); - IpNetwork::with_netmask(IpAddr::V6(addr.ip), netmask) - .map_err(|e| { - tracing::debug!(ip = %addr.ip, ?addr.netmask, ?e, "failed to convert ipv6") - }) - .ok() - } - } - } - - #[tracing::instrument(skip(net_ns))] - async fn do_collect_local_ip_addrs(net_ns: NetNS) -> GetIpListResponse { - let mut ret = GetIpListResponse::default(); - - let ifaces = Self::collect_interfaces(net_ns.clone(), true).await; - let _g = net_ns.guard(); - for iface in ifaces { - for ip in iface.ips { - let ip: std::net::IpAddr = ip.ip(); - if let std::net::IpAddr::V4(v4) = ip { - if ip.is_loopback() || ip.is_multicast() { - continue; - } - ret.interface_ipv4s.push(v4.into()); - } - } - } - - let ifaces = Self::collect_interfaces(net_ns.clone(), false).await; - let _g = net_ns.guard(); - for iface in ifaces { - for ip in iface.ips { - let ip: std::net::IpAddr = ip.ip(); - if let std::net::IpAddr::V6(v6) = ip { - if v6.is_multicast() || v6.is_loopback() || v6.is_unicast_link_local() { - continue; - } - ret.interface_ipv6s.push(v6.into()); - } - } - } - - if let Ok(v4_addr) = local_ipv4().await { - tracing::trace!("got local ipv4: {}", v4_addr); - if !ret.interface_ipv4s.contains(&v4_addr.into()) { - ret.interface_ipv4s.push(v4_addr.into()); - } - } - - if let Ok(v6_addr) = local_ipv6().await { - tracing::trace!("got local ipv6: {}", v6_addr); - if !ret.interface_ipv6s.contains(&v6_addr.into()) { - ret.interface_ipv6s.push(v6_addr.into()); - } - } - - ret - } +} + +#[cfg(target_os = "windows")] +fn convert_windows_interface(iface: SystemNetworkInterface) -> NetworkInterface { + let mac = iface.mac_addr.as_deref().and_then(|mac| { + mac.parse::() + .map_err( + |e| tracing::debug!(iface = %iface.name, mac, ?e, "failed to parse interface mac"), + ) + .ok() + }); + + let ips = iface + .addr + .into_iter() + .filter_map(convert_windows_interface_addr) + .collect(); + + NetworkInterface { + name: iface.name, + description: String::new(), + index: iface.index, + mac, + ips, + // pnet does not populate Windows flags either, so keep the existing semantics. + flags: 0, + } +} + +#[cfg(target_os = "windows")] +fn convert_windows_interface_addr(addr: SystemAddr) -> Option { + match addr { + SystemAddr::V4(addr) => { + let netmask = addr + .netmask + .map(IpAddr::V4) + .unwrap_or(IpAddr::V4(std::net::Ipv4Addr::new(255, 255, 255, 255))); + IpNetwork::with_netmask(IpAddr::V4(addr.ip), netmask) + .map_err( + |e| tracing::debug!(ip = %addr.ip, ?addr.netmask, ?e, "failed to convert ipv4"), + ) + .ok() + } + SystemAddr::V6(addr) => { + let netmask = addr + .netmask + .map(IpAddr::V6) + .unwrap_or(IpAddr::V6(std::net::Ipv6Addr::from(u128::MAX))); + IpNetwork::with_netmask(IpAddr::V6(addr.ip), netmask) + .map_err( + |e| tracing::debug!(ip = %addr.ip, ?addr.netmask, ?e, "failed to convert ipv6"), + ) + .ok() + } + } +} + +#[tracing::instrument(skip(net_ns))] +pub(crate) async fn collect_local_ip_addrs(net_ns: NetNS) -> GetIpListResponse { + #[cfg(target_os = "linux")] + { + return run_in_namespace(net_ns, || async { + collect_local_ip_addrs_in_current_namespace().await + }) + .await; + } + + #[cfg(not(target_os = "linux"))] + { + let _g = net_ns.guard(); + collect_local_ip_addrs_in_current_namespace().await + } +} + +async fn collect_local_ip_addrs_in_current_namespace() -> GetIpListResponse { + let mut ret = GetIpListResponse::default(); + + let ifaces = collect_interfaces_in_current_namespace(true).await; + for iface in ifaces { + for ip in iface.ips { + let ip: std::net::IpAddr = ip.ip(); + if let std::net::IpAddr::V4(v4) = ip { + if ip.is_loopback() || ip.is_multicast() { + continue; + } + ret.interface_ipv4s.push(v4.into()); + } + } + } + + let ifaces = collect_interfaces_in_current_namespace(false).await; + for iface in ifaces { + for ip in iface.ips { + let ip: std::net::IpAddr = ip.ip(); + if let std::net::IpAddr::V6(v6) = ip { + if v6.is_multicast() || v6.is_loopback() || v6.is_unicast_link_local() { + continue; + } + ret.interface_ipv6s.push(v6.into()); + } + } + } + + if let Ok(v4_addr) = local_ipv4().await { + tracing::trace!("got local ipv4: {}", v4_addr); + if !ret.interface_ipv4s.contains(&v4_addr.into()) { + ret.interface_ipv4s.push(v4_addr.into()); + } + } + + if let Ok(v6_addr) = local_ipv6().await { + tracing::trace!("got local ipv6: {}", v6_addr); + if !ret.interface_ipv6s.contains(&v6_addr.into()) { + ret.interface_ipv6s.push(v6_addr.into()); + } + } + + ret +} + +#[cfg(test)] +mod tests { + use super::*; + + #[cfg(target_os = "linux")] + #[tokio::test] + async fn namespace_operation_does_not_migrate_between_os_threads() { + let (before, after) = run_in_namespace(NetNS::new(None), || async { + let before = std::thread::current().id(); + tokio::task::yield_now().await; + (before, std::thread::current().id()) + }) + .await; + + assert_eq!(before, after); + } } diff --git a/easytier/src/common/stun.rs b/easytier/src/common/stun.rs index 14e99ca4..dfaf3f7a 100644 --- a/easytier/src/common/stun.rs +++ b/easytier/src/common/stun.rs @@ -1,1569 +1,250 @@ -use std::collections::BTreeSet; -use std::net::{IpAddr, Ipv4Addr, Ipv6Addr, SocketAddr}; -use std::sync::atomic::AtomicBool; -use std::sync::{Arc, RwLock}; -use std::time::Duration; +//! Native composition for the portable core STUN collector. -use crate::proto::common::{NatType, StunInfo}; -use anyhow::Context; -use chrono::Local; -use crossbeam::atomic::AtomicCell; -use quanta::Instant; -use rand::seq::IteratorRandom; -use socket2::{SockAddr, SockRef}; -use tokio::io::{AsyncReadExt, AsyncWriteExt}; -use tokio::net::{UdpSocket, lookup_host}; -use tokio::sync::{Mutex, broadcast}; -use tokio::task::JoinSet; -use tracing::{Instrument, Level}; +#[cfg(test)] +use std::{net::SocketAddr, sync::Arc}; -use bytecodec::{DecodeExt, EncodeExt}; -use stun_codec::rfc5389::methods::BINDING; -use stun_codec::{Message, MessageClass, MessageDecoder, MessageEncoder}; +#[cfg(test)] +use async_trait::async_trait; +#[cfg(test)] +use easytier_core::connectivity::stun::{StunInfoProvider, StunSocketMapper}; +use easytier_core::{ + connectivity::stun::StunInfoCollector as CoreStunInfoCollector, socket::SocketContext, +}; +#[cfg(test)] +use easytier_proto::common::{NatType, StunInfo}; -use crate::common::error::Error; +use crate::host_runtime::{NativeHostRuntime, native_host_runtime}; +#[cfg(test)] +use crate::socket::udp::RuntimeUdpSocket; -use super::dns::resolve_txt_record; -use super::stun_codec_ext::*; +pub type StunInfoCollector = CoreStunInfoCollector; -const DEFAULT_UDP_STUN_SERVERS: &[&str] = &[ - "txt:stun.easytier.cn", - "stun.miwifi.com", - "stun.chat.bilibili.com", - "stun.hitv.com", -]; - -const DEFAULT_TCP_STUN_SERVERS: &[&str] = &[ - "stun.hot-chilli.net", - "stun.fitauto.ru", - "fwa.lifesizecloud.com", - "global.turn.twilio.com", - "turn.cloudflare.com", - "stun.voip.blackberry.com", - "stun.radiojar.com", -]; - -const DEFAULT_UDP_V6_STUN_SERVERS: &[&str] = &["txt:stun-v6.easytier.cn"]; - -struct HostResolverIter { - hostnames: Vec, - ips: Vec, - max_ip_per_domain: u32, - use_ipv6: bool, -} - -impl HostResolverIter { - fn new(hostnames: Vec, max_ip_per_domain: u32, use_ipv6: bool) -> Self { - Self { - hostnames, - ips: vec![], - max_ip_per_domain, - use_ipv6, - } - } - - async fn get_txt_record(domain_name: &str) -> Result, Error> { - let txt_data = resolve_txt_record(domain_name).await?; - Ok(txt_data.split(" ").map(|x| x.to_string()).collect()) - } - - #[async_recursion::async_recursion] - async fn next(&mut self) -> Option { - if self.ips.is_empty() { - if self.hostnames.is_empty() { - return None; - } - - let host = self.hostnames.remove(0); - let host = if host.contains(':') { - host - } else { - format!("{}:3478", host) - }; - - if host.starts_with("txt:") { - let domain_name = host.trim_start_matches("txt:"); - match Self::get_txt_record(domain_name).await { - Ok(hosts) => { - tracing::info!( - ?domain_name, - ?hosts, - "get txt record success when resolve stun server" - ); - // insert hosts to the head of hostnames - self.hostnames.splice(0..0, hosts.into_iter()); - } - Err(e) => { - tracing::warn!( - ?domain_name, - ?e, - "get txt record failed when resolve stun server" - ); - } - } - return self.next().await; - } - - let use_ipv6 = self.use_ipv6; - - match lookup_host(&host).await { - Ok(ips) => { - self.ips = ips - .filter(|x| if use_ipv6 { x.is_ipv6() } else { x.is_ipv4() }) - .choose_multiple(&mut rand::thread_rng(), self.max_ip_per_domain as usize); - - if self.ips.is_empty() { - return self.next().await; - } - } - Err(e) => { - tracing::warn!(?host, ?e, "lookup host for stun failed"); - return self.next().await; - } - }; - } - - Some(self.ips.remove(0)) +pub fn default_udp_stun_servers() -> Vec { + if cfg!(test) { + Vec::new() + } else { + StunInfoCollector::get_default_servers() } } -#[derive(Debug, Clone)] -struct StunPacket { - data: Vec, - addr: SocketAddr, -} - -type StunPacketReceiver = tokio::sync::broadcast::Receiver; - -#[derive(Debug, Clone, Copy)] -struct BindRequestResponse { - local_addr: SocketAddr, - stun_server_addr: SocketAddr, - - recv_from_addr: SocketAddr, - mapped_socket_addr: Option, - changed_socket_addr: Option, - - change_ip: bool, - change_port: bool, - - real_ip_changed: bool, - real_port_changed: bool, - - latency_us: u32, -} - -impl BindRequestResponse { - pub fn get_mapped_addr_no_check(&self) -> &SocketAddr { - self.mapped_socket_addr.as_ref().unwrap() +pub fn default_tcp_stun_servers() -> Vec { + if cfg!(test) { + Vec::new() + } else { + StunInfoCollector::get_default_tcp_servers() } } -#[derive(Debug, Clone)] -struct StunClient { - stun_server: SocketAddr, - resp_timeout: Duration, - req_repeat: u32, - socket: Arc, - stun_packet_receiver: Arc>, -} - -impl StunClient { - pub fn new( - stun_server: SocketAddr, - socket: Arc, - stun_packet_receiver: StunPacketReceiver, - ) -> Self { - Self { - stun_server, - resp_timeout: Duration::from_millis(3000), - req_repeat: 2, - socket, - stun_packet_receiver: Arc::new(Mutex::new(stun_packet_receiver)), - } - } - - #[tracing::instrument(skip(self, buf))] - async fn wait_stun_response<'a, const N: usize>( - &self, - buf: &'a mut [u8; N], - tids: &Vec, - expected_ip_changed: bool, - expected_port_changed: bool, - stun_host: &SocketAddr, - ) -> Result<(Message, SocketAddr), Error> { - let mut now = tokio::time::Instant::now(); - let deadline = now + self.resp_timeout; - - while now < deadline { - let mut locked_receiver = self.stun_packet_receiver.lock().await; - let stun_packet_raw = tokio::time::timeout(deadline - now, locked_receiver.recv()) - .await? - .with_context(|| "recv stun packet from broadcast channel error")?; - now = tokio::time::Instant::now(); - - let (len, remote_addr) = (stun_packet_raw.data.len(), stun_packet_raw.addr); - - if len < 20 { - continue; - } - - let udp_buf = stun_packet_raw.data; - - // TODO:: we cannot borrow `buf` directly in udp recv_from, so we copy it here - unsafe { std::ptr::copy(udp_buf.as_ptr(), buf.as_ptr() as *mut u8, len) }; - - let mut decoder = MessageDecoder::::new(); - let Ok(msg) = decoder - .decode_from_bytes(&buf[..len]) - .with_context(|| format!("decode stun msg {:?}", buf))? - else { - continue; - }; - - tracing::trace!(b = ?&udp_buf[..len], ?tids, ?remote_addr, ?stun_host, "recv stun response, msg: {:#?}", msg); - - if msg.class() != MessageClass::SuccessResponse - || msg.method() != BINDING - || !tids.contains(&tid_to_u32(&msg.transaction_id())) - { - continue; - } - - return Ok((msg, remote_addr)); - } - - Err(Error::Unknown) - } - - fn extrace_mapped_addr(msg: &Message) -> Option { - let mut mapped_addr = None; - for x in msg.attributes() { - match x { - Attribute::MappedAddress(addr) if mapped_addr.is_none() => { - let _ = mapped_addr.insert(addr.address()); - } - Attribute::XorMappedAddress(addr) if mapped_addr.is_none() => { - let _ = mapped_addr.insert(addr.address()); - } - _ => {} - } - } - mapped_addr - } - - fn extract_changed_addr(msg: &Message) -> Option { - let mut changed_addr = None; - for x in msg.attributes() { - match x { - Attribute::OtherAddress(m) if changed_addr.is_none() => { - let _ = changed_addr.insert(m.address()); - } - Attribute::ChangedAddress(m) if changed_addr.is_none() => { - let _ = changed_addr.insert(m.address()); - } - _ => {} - } - } - changed_addr - } - - #[tracing::instrument(ret, level = Level::TRACE)] - pub async fn bind_request( - self, - change_ip: bool, - change_port: bool, - ) -> Result { - let stun_host = self.stun_server; - // repeat req in case of packet loss - let mut tids = vec![]; - for _ in 0..self.req_repeat { - let tid = rand::random::(); - // let tid = 1; - let mut buf = [0u8; 28]; - // memset buf - unsafe { std::ptr::write_bytes(buf.as_mut_ptr(), 0, buf.len()) }; - - let mut message = - Message::::new(MessageClass::Request, BINDING, u32_to_tid(tid)); - message.add_attribute(ChangeRequest::new(change_ip, change_port)); - - // Encodes the message - let mut encoder = MessageEncoder::new(); - let msg = encoder - .encode_into_bytes(message.clone()) - .with_context(|| "encode stun message")?; - tids.push(tid); - tracing::trace!(?message, ?msg, tid, "send stun request"); - self.socket.send_to(msg.as_slice(), &stun_host).await?; - } - - let now = Instant::now(); - - tracing::trace!("waiting stun response"); - let mut buf = [0; 1620]; - let (msg, recv_addr) = self - .wait_stun_response(&mut buf, &tids, change_ip, change_port, &stun_host) - .await?; - - let changed_socket_addr = Self::extract_changed_addr(&msg); - let real_ip_changed = stun_host.ip() != recv_addr.ip(); - let real_port_changed = stun_host.port() != recv_addr.port(); - - let resp = BindRequestResponse { - local_addr: self.socket.local_addr()?, - stun_server_addr: stun_host, - recv_from_addr: recv_addr, - mapped_socket_addr: Self::extrace_mapped_addr(&msg), - changed_socket_addr, - change_ip, - change_port, - - real_ip_changed, - real_port_changed, - - latency_us: now.elapsed().as_micros() as u32, - }; - - tracing::trace!( - ?stun_host, - ?recv_addr, - ?changed_socket_addr, - "finish stun bind request" - ); - - Ok(resp) - } -} - -struct StunClientBuilder { - udp: Arc, - task_set: JoinSet<()>, - stun_packet_sender: broadcast::Sender, -} - -impl StunClientBuilder { - pub fn new(udp: Arc) -> Self { - let (stun_packet_sender, _) = broadcast::channel(1024); - let mut task_set = JoinSet::new(); - - let udp_clone = udp.clone(); - let stun_packet_sender_clone = stun_packet_sender.clone(); - task_set.spawn( - async move { - let mut buf = [0; 1620]; - tracing::trace!("start stun packet listener"); - loop { - let Ok((len, addr)) = udp_clone.recv_from(&mut buf).await else { - tracing::error!("udp recv_from error"); - break; - }; - let data = buf[..len].to_vec(); - tracing::trace!(?addr, ?data, "recv udp stun packet"); - let _ = stun_packet_sender_clone.send(StunPacket { data, addr }); - } - } - .instrument(tracing::info_span!("stun_packet_listener")), - ); - - Self { - udp, - task_set, - stun_packet_sender, - } - } - - pub fn new_stun_client(&self, stun_server: SocketAddr) -> StunClient { - StunClient::new( - stun_server, - self.udp.clone(), - self.stun_packet_sender.subscribe(), - ) - } - - pub async fn stop(&mut self) { - self.task_set.abort_all(); - while self.task_set.join_next().await.is_some() {} - } -} - -#[derive(Debug, Clone)] -pub enum StunTransport { - Udp, - Tcp, -} - -#[derive(Debug, Clone)] -pub struct StunNatTypeDetectResult { - transport: StunTransport, - source_addr: SocketAddr, - stun_resps: Vec, - // if we are easy symmetric nat, we need to test with another port to check inc or dec - extra_bind_test: Option, -} - -impl StunNatTypeDetectResult { - fn new( - transport: StunTransport, - source_addr: SocketAddr, - stun_resps: Vec, - ) -> Self { - Self { - transport, - source_addr, - stun_resps, - extra_bind_test: None, - } - } - - fn has_ip_changed_resp(&self) -> bool { - for resp in self.stun_resps.iter() { - if resp.real_ip_changed { - return true; - } - } - false - } - - fn has_port_changed_resp(&self) -> bool { - for resp in self.stun_resps.iter() { - if resp.real_port_changed { - return true; - } - } - false - } - - fn is_open_internet(&self) -> bool { - for resp in self.stun_resps.iter() { - if resp.mapped_socket_addr == Some(self.source_addr) { - return true; - } - } - false - } - - fn is_no_pat(&self) -> bool { - for resp in self.stun_resps.iter() { - if resp.mapped_socket_addr.map(|x| x.port()) == Some(self.source_addr.port()) { - return true; - } - } - false - } - - fn stun_server_count(&self) -> usize { - // find resp with distinct stun server - self.stun_resps - .iter() - .map(|x| x.recv_from_addr) - .collect::>() - .len() - } - - fn is_cone(&self) -> bool { - // if unique mapped addr count is less than stun server count, it is cone - let mapped_addr_count = self - .stun_resps - .iter() - .filter_map(|x| x.mapped_socket_addr) - .collect::>() - .len(); - mapped_addr_count == 1 - } - - fn nat_type_udp(&self) -> NatType { - if self.stun_server_count() < 2 { - return NatType::Unknown; - } - - if self.is_cone() { - if self.has_ip_changed_resp() { - if self.is_open_internet() { - NatType::OpenInternet - } else if self.is_no_pat() { - NatType::NoPat - } else { - NatType::FullCone - } - } else if self.has_port_changed_resp() { - NatType::Restricted - } else { - NatType::PortRestricted - } - } else if !self.stun_resps.is_empty() { - if self.public_ips().len() != 1 - || self.usable_stun_resp_count() <= 1 - || self.max_port() - self.min_port() > 15 - { - NatType::Symmetric - } else if let Some(extra_bind_mapped) = self - .extra_bind_test - .as_ref() - .and_then(|extra| extra.mapped_socket_addr) - { - let extra_port = extra_bind_mapped.port(); - - let max_port_diff = extra_port.saturating_sub(self.max_port()); - let min_port_diff = self.min_port().saturating_sub(extra_port); - if max_port_diff != 0 && max_port_diff < 100 { - NatType::SymmetricEasyInc - } else if min_port_diff != 0 && min_port_diff < 100 { - NatType::SymmetricEasyDec - } else { - NatType::Symmetric - } - } else { - NatType::Symmetric - } - } else { - NatType::Unknown - } - } - - fn nat_type_tcp(&self) -> NatType { - if self.is_open_internet() { - return NatType::OpenInternet; - } - - if self.stun_server_count() < 2 || self.stun_resps.is_empty() { - return NatType::Unknown; - } - - if self.is_cone() { - if self.is_no_pat() { - NatType::NoPat - } else { - NatType::FullCone - } - } else { - NatType::Symmetric - } - } - - pub fn nat_type(&self) -> NatType { - match self.transport { - StunTransport::Udp => self.nat_type_udp(), - StunTransport::Tcp => self.nat_type_tcp(), - } - } - - pub fn public_ips(&self) -> Vec { - self.stun_resps - .iter() - .filter_map(|x| x.mapped_socket_addr.map(|x| x.ip())) - .collect::>() - .into_iter() - .collect() - } - - pub fn collect_available_stun_server(&self) -> Vec { - let mut ret = vec![]; - for resp in self.stun_resps.iter() { - if !ret.contains(&resp.stun_server_addr) { - ret.push(resp.stun_server_addr); - } - } - ret - } - - pub fn local_addr(&self) -> SocketAddr { - self.source_addr - } - - pub fn extend_result(&mut self, other: StunNatTypeDetectResult) { - self.stun_resps.extend(other.stun_resps); - } - - pub fn min_port(&self) -> u16 { - self.stun_resps - .iter() - .filter_map(|x| x.mapped_socket_addr.map(|x| x.port())) - .min() - .unwrap_or(0) - } - - pub fn max_port(&self) -> u16 { - self.stun_resps - .iter() - .filter_map(|x| x.mapped_socket_addr.map(|x| x.port())) - .max() - .unwrap_or(u16::MAX) - } - - pub fn usable_stun_resp_count(&self) -> usize { - self.stun_resps - .iter() - .filter(|x| x.mapped_socket_addr.is_some()) - .count() - } -} - -pub struct UdpNatTypeDetector { - stun_server_hosts: Vec, - max_ip_per_domain: u32, -} - -impl UdpNatTypeDetector { - pub fn new(stun_server_hosts: Vec, max_ip_per_domain: u32) -> Self { - Self { - stun_server_hosts, - max_ip_per_domain, - } - } - - async fn get_extra_bind_result( - &self, - source_port: u16, - stun_server: SocketAddr, - ) -> Result { - let udp = Arc::new(UdpSocket::bind(format!("0.0.0.0:{}", source_port)).await?); - let client_builder = StunClientBuilder::new(udp.clone()); - client_builder - .new_stun_client(stun_server) - .bind_request(false, false) - .await - } - - pub async fn detect_nat_type( - &self, - source_port: u16, - ) -> Result { - let udp = Arc::new(UdpSocket::bind(format!("0.0.0.0:{}", source_port)).await?); - self.detect_nat_type_with_socket(udp).await - } - - #[tracing::instrument(skip(self))] - pub async fn detect_nat_type_with_socket( - &self, - udp: Arc, - ) -> Result { - let mut stun_servers = vec![]; - let mut host_resolver = HostResolverIter::new( - self.stun_server_hosts.clone(), - self.max_ip_per_domain, - false, - ); - while let Some(addr) = host_resolver.next().await { - stun_servers.push(addr); - } - - let client_builder = StunClientBuilder::new(udp.clone()); - let mut stun_task_set = JoinSet::new(); - - for stun_server in stun_servers.iter() { - stun_task_set.spawn( - client_builder - .new_stun_client(*stun_server) - .bind_request(false, false), - ); - stun_task_set.spawn( - client_builder - .new_stun_client(*stun_server) - .bind_request(false, true), - ); - stun_task_set.spawn( - client_builder - .new_stun_client(*stun_server) - .bind_request(true, true), - ); - } - - let mut bind_resps = vec![]; - while let Some(resp) = stun_task_set.join_next().await { - if let Ok(Ok(resp)) = resp { - bind_resps.push(resp); - } - } - - Ok(StunNatTypeDetectResult::new( - StunTransport::Udp, - udp.local_addr()?, - bind_resps, - )) - } -} - -#[derive(Debug, Clone)] -struct TcpStunClient { - stun_server: SocketAddr, - conn_timeout: Duration, - io_timeout: Duration, - source_port: u16, -} - -impl TcpStunClient { - pub fn new(stun_server: SocketAddr, source_port: u16) -> Self { - Self { - stun_server, - conn_timeout: Duration::from_millis(1500), - io_timeout: Duration::from_millis(3000), - source_port, - } - } - - fn extract_mapped_addr(msg: &Message) -> Option { - let mut mapped_addr = None; - for x in msg.attributes() { - match x { - Attribute::MappedAddress(addr) if mapped_addr.is_none() => { - let _ = mapped_addr.insert(addr.address()); - } - Attribute::XorMappedAddress(addr) if mapped_addr.is_none() => { - let _ = mapped_addr.insert(addr.address()); - } - _ => {} - } - } - mapped_addr - } - - fn message_size_from_header(header: &[u8; 20]) -> Result { - if (header[0] & 0b1100_0000) != 0 { - return Err(Error::MessageDecodeError( - "invalid stun message type".to_string(), - )); - } - let msg_len = u16::from_be_bytes([header[2], header[3]]) as usize; - if !msg_len.is_multiple_of(4) { - return Err(Error::MessageDecodeError( - "invalid stun message length".to_string(), - )); - } - let total = 20usize - .checked_add(msg_len) - .ok_or_else(|| Error::MessageDecodeError("invalid stun message size".to_string()))?; - if total > 4096 { - return Err(Error::MessageDecodeError( - "stun message too large".to_string(), - )); - } - Ok(total) - } - - async fn tcp_read_stun_message( - stream: &mut tokio::net::TcpStream, - timeout: Duration, - ) -> Result, Error> { - let mut header = [0u8; 20]; - tokio::time::timeout(timeout, stream.read_exact(&mut header)).await??; - let total_size = Self::message_size_from_header(&header)?; - let mut buf = vec![0u8; total_size]; - buf[..20].copy_from_slice(&header); - if total_size > 20 { - tokio::time::timeout(timeout, stream.read_exact(&mut buf[20..])).await??; - } - - let mut decoder = MessageDecoder::::new(); - let Ok(msg) = decoder - .decode_from_bytes(&buf) - .with_context(|| "decode tcp stun message")? - else { - return Err(Error::MessageDecodeError( - "invalid stun message".to_string(), - )); - }; - Ok(msg) - } - - async fn connect(&self) -> Result { - let bind_addr = match self.stun_server { - SocketAddr::V4(_) => { - SocketAddr::new(IpAddr::V4(Ipv4Addr::UNSPECIFIED), self.source_port) - } - SocketAddr::V6(_) => { - SocketAddr::new(IpAddr::V6(Ipv6Addr::UNSPECIFIED), self.source_port) - } - }; - - let socket2_socket = socket2::Socket::new( - socket2::Domain::for_address(self.stun_server), - socket2::Type::STREAM, - Some(socket2::Protocol::TCP), - )?; - - if bind_addr.is_ipv6() { - socket2_socket.set_only_v6(true)?; - } - - socket2_socket.set_nonblocking(true)?; - socket2_socket.set_reuse_address(true)?; - - #[cfg(all(unix, not(target_os = "solaris"), not(target_os = "illumos")))] - { - let _ = socket2_socket.set_reuse_port(true); - } - - socket2_socket.bind(&SockAddr::from(bind_addr))?; - - let socket = tokio::net::TcpSocket::from_std_stream(socket2_socket.into()); - let stream = - tokio::time::timeout(self.conn_timeout, socket.connect(self.stun_server)).await??; - - let _ = SockRef::from(&stream).set_linger(Some(Duration::ZERO)); - - Ok(stream) - } - - #[tracing::instrument(ret, level = Level::TRACE)] - pub async fn bind_request(self) -> Result { - let mut stream = self.connect().await?; - let local_addr = stream.local_addr()?; - let stun_host = self.stun_server; - - let tid = rand::random::(); - let message = Message::::new(MessageClass::Request, BINDING, u32_to_tid(tid)); - let mut encoder = MessageEncoder::new(); - let msg = encoder - .encode_into_bytes(message.clone()) - .with_context(|| "encode tcp stun message")?; - tokio::time::timeout(self.io_timeout, stream.write_all(msg.as_slice())).await??; - - let now = Instant::now(); - let msg = Self::tcp_read_stun_message(&mut stream, self.io_timeout).await?; - if msg.class() != MessageClass::SuccessResponse - || msg.method() != BINDING - || tid_to_u32(&msg.transaction_id()) != tid - { - return Err(Error::MessageDecodeError( - "unexpected stun response".to_string(), - )); - } - - Ok(BindRequestResponse { - local_addr, - stun_server_addr: stun_host, - recv_from_addr: stun_host, - mapped_socket_addr: Self::extract_mapped_addr(&msg), - changed_socket_addr: None, - change_ip: false, - change_port: false, - real_ip_changed: false, - real_port_changed: false, - latency_us: now.elapsed().as_micros() as u32, - }) - } -} - -pub struct TcpNatTypeDetector { - stun_server_hosts: Vec, - max_ip_per_domain: u32, -} - -impl TcpNatTypeDetector { - pub fn new(stun_server_hosts: Vec, max_ip_per_domain: u32) -> Self { - Self { - stun_server_hosts, - max_ip_per_domain, - } - } - - #[tracing::instrument(skip(self))] - pub async fn detect_nat_type( - &self, - source_port: u16, - ) -> Result { - let mut stun_servers = vec![]; - let mut host_resolver = HostResolverIter::new( - self.stun_server_hosts.clone(), - self.max_ip_per_domain, - false, - ); - while let Some(addr) = host_resolver.next().await { - stun_servers.push(addr); - } - - let mut bind_resps = vec![]; - let mut source_addr = None; - let mut selected_source_port = if source_port == 0 { - None - } else { - Some(source_port) - }; - for server in stun_servers.iter() { - let resp = TcpStunClient::new(*server, selected_source_port.unwrap_or(0)) - .bind_request() - .await; - if let Ok(resp) = resp { - if selected_source_port.is_none() { - selected_source_port = Some(resp.local_addr.port()); - } - source_addr.get_or_insert(resp.local_addr); - bind_resps.push(resp); - if bind_resps.len() >= 3 { - break; - } - } - } - - let Some(source_addr) = source_addr else { - return Err(Error::NotFound); - }; - Ok(StunNatTypeDetectResult::new( - StunTransport::Tcp, - source_addr, - bind_resps, - )) - } -} - -#[async_trait::async_trait] -#[auto_impl::auto_impl(&, Arc, Box)] -pub trait StunInfoCollectorTrait: Send + Sync { - fn get_stun_info(&self) -> StunInfo; - async fn get_udp_port_mapping(&self, local_port: u16) -> Result; - async fn get_udp_port_mapping_with_socket( - &self, - udp: Arc, - ) -> Result; - async fn get_tcp_port_mapping(&self, local_port: u16) -> Result; -} - -pub struct StunInfoCollector { - stun_servers: Arc>>, - tcp_stun_servers: Arc>>, - stun_servers_v6: Arc>>, - udp_nat_test_result: Arc>>, - tcp_nat_test_result: Arc>>, - public_ipv6: Arc>>, - nat_test_result_time: Arc>>, - redetect_notify: Arc, - tasks: std::sync::Mutex>, - started: AtomicBool, -} - -#[async_trait::async_trait] -impl StunInfoCollectorTrait for StunInfoCollector { - fn get_stun_info(&self) -> StunInfo { - self.start_stun_routine(); - - let udp_result = self.udp_nat_test_result.read().unwrap().clone(); - let tcp_result = self.tcp_nat_test_result.read().unwrap().clone(); - if udp_result.is_none() && tcp_result.is_none() { - return Default::default(); - } - - let mut public_ip = BTreeSet::::new(); - if let Some(result) = &udp_result { - public_ip.extend(result.public_ips().into_iter().map(|x| x.to_string())); - } - if let Some(result) = &tcp_result { - public_ip.extend(result.public_ips().into_iter().map(|x| x.to_string())); - } - if let Some(v6) = self.public_ipv6.load() { - public_ip.insert(v6.to_string()); - } - - StunInfo { - udp_nat_type: udp_result - .as_ref() - .map(|x| x.nat_type() as i32) - .unwrap_or(NatType::Unknown as i32), - tcp_nat_type: tcp_result - .as_ref() - .map(|x| x.nat_type() as i32) - .unwrap_or(NatType::Unknown as i32), - last_update_time: self.nat_test_result_time.load().timestamp(), - public_ip: public_ip.into_iter().collect(), - min_port: udp_result - .as_ref() - .map(|x| x.min_port() as u32) - .or_else(|| tcp_result.as_ref().map(|x| x.min_port() as u32)) - .unwrap_or(0), - max_port: udp_result - .as_ref() - .map(|x| x.max_port() as u32) - .or_else(|| tcp_result.as_ref().map(|x| x.max_port() as u32)) - .unwrap_or(0), - } - } - - async fn get_udp_port_mapping(&self, local_port: u16) -> Result { - let udp = Arc::new(UdpSocket::bind(format!("0.0.0.0:{}", local_port)).await?); - self.get_udp_port_mapping_with_socket(udp).await - } - - async fn get_udp_port_mapping_with_socket( - &self, - udp: Arc, - ) -> Result { - self.start_stun_routine(); - - let mut stun_servers = self - .udp_nat_test_result - .read() - .unwrap() - .clone() - .map(|x| x.collect_available_stun_server()) - .unwrap_or_default(); - - if stun_servers.is_empty() { - let mut host_resolver = - HostResolverIter::new(self.stun_servers.read().unwrap().clone(), 2, false); - while let Some(addr) = host_resolver.next().await { - stun_servers.push(addr); - if stun_servers.len() >= 2 { - break; - } - } - } - - if stun_servers.is_empty() { - return Err(Error::NotFound); - } - - let mut client_builder = StunClientBuilder::new(udp.clone()); - - for server in stun_servers.iter() { - let Ok(ret) = client_builder - .new_stun_client(*server) - .bind_request(false, false) - .await - else { - tracing::warn!(?server, "stun bind request failed"); - continue; - }; - if let Some(mapped_addr) = ret.mapped_socket_addr { - // make sure udp socket is available after return ok. - client_builder.stop().await; - return Ok(mapped_addr); - } - } - - Err(Error::NotFound) - } - - async fn get_tcp_port_mapping(&self, local_port: u16) -> Result { - self.start_stun_routine(); - - let mut stun_servers = self - .tcp_nat_test_result - .read() - .unwrap() - .clone() - .map(|x| x.collect_available_stun_server()) - .unwrap_or_default(); - - if stun_servers.is_empty() { - let mut host_resolver = - HostResolverIter::new(self.tcp_stun_servers.read().unwrap().clone(), 2, false); - while let Some(addr) = host_resolver.next().await { - stun_servers.push(addr); - if stun_servers.len() >= 2 { - break; - } - } - } - - if stun_servers.is_empty() { - return Err(Error::NotFound); - } - - for server in stun_servers.iter() { - let Ok(ret) = TcpStunClient::new(*server, local_port).bind_request().await else { - tracing::warn!(?server, "tcp stun bind request failed"); - continue; - }; - - if let Some(mapped_addr) = ret.mapped_socket_addr { - return Ok(mapped_addr); - } - } - - Err(Error::NotFound) - } -} - -impl StunInfoCollector { - pub fn new( - udp_stun_servers: Vec, - tcp_stun_servers: Vec, - stun_servers_v6: Vec, - ) -> Self { - Self { - stun_servers: Arc::new(RwLock::new(udp_stun_servers)), - tcp_stun_servers: Arc::new(RwLock::new(tcp_stun_servers)), - stun_servers_v6: Arc::new(RwLock::new(stun_servers_v6)), - udp_nat_test_result: Arc::new(RwLock::new(None)), - tcp_nat_test_result: Arc::new(RwLock::new(None)), - public_ipv6: Arc::new(AtomicCell::new(None)), - nat_test_result_time: Arc::new(AtomicCell::new(Local::now())), - redetect_notify: Arc::new(tokio::sync::Notify::new()), - tasks: std::sync::Mutex::new(JoinSet::new()), - started: AtomicBool::new(false), - } - } - - pub fn new_with_default_servers() -> Self { - Self::new( - Self::get_default_servers(), - Self::get_default_tcp_servers(), - Self::get_default_servers_v6(), - ) - } - - pub fn set_stun_servers(&self, stun_servers: Vec) { - let mut g = self.stun_servers.write().unwrap(); - *g = stun_servers; - } - - pub fn set_stun_servers_v6(&self, stun_servers_v6: Vec) { - let mut g = self.stun_servers_v6.write().unwrap(); - *g = stun_servers_v6; - } - - pub fn set_tcp_stun_servers(&self, stun_servers: Vec) { - let mut g = self.tcp_stun_servers.write().unwrap(); - *g = stun_servers; - } - - pub fn get_default_servers() -> Vec { - if cfg!(test) { - Vec::new() - } else { - // NOTICE: we may need to choose stun server based on geolocation - // stun server cross nation may return an external ip address with high latency and loss rate - DEFAULT_UDP_STUN_SERVERS - .iter() - .map(ToString::to_string) - .collect() - } - } - - pub fn get_default_tcp_servers() -> Vec { - // if test, return empty vector - if cfg!(test) { - Vec::new() - } else { - DEFAULT_TCP_STUN_SERVERS - .iter() - .map(ToString::to_string) - .collect() - } - } - - pub fn get_default_servers_v6() -> Vec { - if cfg!(test) { - Vec::new() - } else { - DEFAULT_UDP_V6_STUN_SERVERS - .iter() - .map(ToString::to_string) - .collect() - } - } - - async fn get_public_ipv6(servers: &[String]) -> Option { - let mut ips = HostResolverIter::new(servers.to_vec(), 10, true); - while let Some(ip) = ips.next().await { - let Ok(udp_socket) = UdpSocket::bind("[::]:0".to_string()).await else { - break; - }; - let udp = Arc::new(udp_socket); - let ret = StunClientBuilder::new(udp.clone()) - .new_stun_client(ip) - .bind_request(false, false) - .await; - tracing::debug!(?ret, "finish ipv6 udp nat type detect"); - if let Ok(Some(IpAddr::V6(v6))) = ret.map(|x| x.mapped_socket_addr.map(|x| x.ip())) { - return Some(v6); - } - } - None - } - - fn start_stun_routine(&self) { - if self.started.load(std::sync::atomic::Ordering::Relaxed) { - return; - } - self.started - .store(true, std::sync::atomic::Ordering::Relaxed); - - let stun_servers = self.stun_servers.clone(); - let udp_nat_test_result = self.udp_nat_test_result.clone(); - let nat_test_time = self.nat_test_result_time.clone(); - let redetect_notify = self.redetect_notify.clone(); - self.tasks.lock().unwrap().spawn(async move { - loop { - let udp_servers = stun_servers.read().unwrap().clone(); - let udp_servers: Vec = udp_servers - .iter() - .take(2) - .chain(udp_servers.iter().skip(2).choose(&mut rand::thread_rng())) - .map(|x| x.to_string()) - .collect(); - - let udp_detector = UdpNatTypeDetector::new(udp_servers, 1); - let mut udp_ret = udp_detector.detect_nat_type(0).await; - tracing::debug!(?udp_ret, "finish udp nat type detect"); - - let mut nat_type = NatType::Unknown; - if let Ok(resp) = &udp_ret { - tracing::debug!(?resp, "got udp nat type detect result"); - nat_type = resp.nat_type(); - } - - // if nat type is symmtric, detect with another port to gather more info - if nat_type == NatType::Symmetric { - let old_resp = udp_ret.as_mut().unwrap(); - tracing::debug!(?old_resp, "start get extra bind result"); - let available_stun_servers = old_resp.collect_available_stun_server(); - for server in available_stun_servers.iter() { - let ret = udp_detector - .get_extra_bind_result(0, *server) - .await - .with_context(|| "get extra bind result failed"); - tracing::debug!(?ret, "finish udp nat type detect with another port"); - if let Ok(resp) = ret { - old_resp.extra_bind_test = Some(resp); - break; - } - } - } - - let mut sleep_sec = 10; - if let Ok(resp) = &udp_ret { - nat_test_time.store(Local::now()); - *udp_nat_test_result.write().unwrap() = Some(resp.clone()); - if nat_type != NatType::Unknown - && (nat_type != NatType::Symmetric || resp.extra_bind_test.is_some()) - { - sleep_sec = 600 - } - } - - tokio::select! { - _ = redetect_notify.notified() => {} - _ = tokio::time::sleep(Duration::from_secs(sleep_sec)) => {} - } - } - }); - - let tcp_stun_servers = self.tcp_stun_servers.clone(); - let tcp_nat_test_result = self.tcp_nat_test_result.clone(); - let nat_test_time = self.nat_test_result_time.clone(); - let redetect_notify = self.redetect_notify.clone(); - self.tasks.lock().unwrap().spawn(async move { - loop { - let tcp_servers = tcp_stun_servers.read().unwrap().clone(); - let tcp_servers: Vec = tcp_servers - .iter() - .take(2) - .chain(tcp_servers.iter().skip(2).choose(&mut rand::thread_rng())) - .map(|x| x.to_string()) - .collect(); - - let tcp_detector = TcpNatTypeDetector::new(tcp_servers, 1); - let tcp_ret = tcp_detector.detect_nat_type(0).await; - tracing::debug!(?tcp_ret, "finish tcp nat type detect"); - - let mut sleep_sec = 10; - if let Ok(resp) = &tcp_ret { - nat_test_time.store(Local::now()); - *tcp_nat_test_result.write().unwrap() = Some(resp.clone()); - if resp.nat_type() != NatType::Unknown { - sleep_sec = 600; - } - } - - tokio::select! { - _ = redetect_notify.notified() => {} - _ = tokio::time::sleep(Duration::from_secs(sleep_sec)) => {} - } - } - }); - - // for ipv6 - let stun_servers = self.stun_servers_v6.clone(); - let stored_ipv6 = self.public_ipv6.clone(); - let redetect_notify = self.redetect_notify.clone(); - self.tasks.lock().unwrap().spawn(async move { - loop { - let servers = stun_servers.read().unwrap().clone(); - if let Some(x) = Self::get_public_ipv6(&servers).await { - stored_ipv6.store(Some(x)) - } - - let sleep_sec = if stored_ipv6.load().is_none() { - 60 - } else { - 360 - }; - - tokio::select! { - _ = redetect_notify.notified() => {} - _ = tokio::time::sleep(Duration::from_secs(sleep_sec)) => {} - } - } - }); - } - - pub fn update_stun_info(&self) { - self.redetect_notify.notify_waiters(); - } -} - -pub struct MockStunInfoCollector { - pub udp_nat_type: NatType, -} - -#[async_trait::async_trait] -impl StunInfoCollectorTrait for MockStunInfoCollector { - fn get_stun_info(&self) -> StunInfo { - StunInfo { - udp_nat_type: self.udp_nat_type as i32, - tcp_nat_type: NatType::Unknown as i32, - last_update_time: Local::now().timestamp(), - min_port: 100, - max_port: 200, - public_ip: vec!["127.0.0.1".to_string(), "::1".to_string()], - } - } - - async fn get_udp_port_mapping(&self, mut port: u16) -> Result { - if port == 0 { - port = 40144; - } - Ok(format!("127.0.0.1:{}", port).parse().unwrap()) - } - - async fn get_udp_port_mapping_with_socket( - &self, - udp: Arc, - ) -> Result { - self.get_udp_port_mapping(udp.local_addr()?.port()).await - } - - async fn get_tcp_port_mapping(&self, mut port: u16) -> Result { - if port == 0 { - port = 40144; - } - Ok(format!("127.0.0.1:{}", port).parse().unwrap()) +pub fn default_udp_v6_stun_servers() -> Vec { + if cfg!(test) { + Vec::new() + } else { + StunInfoCollector::get_default_servers_v6() } } +#[cfg(test)] +pub struct MockStunInfoCollector { + pub udp_nat_type: NatType, +} + +#[cfg(test)] +#[async_trait] +impl StunInfoProvider for MockStunInfoCollector { + fn get_stun_info(&self) -> StunInfo { + StunInfo { + udp_nat_type: self.udp_nat_type as i32, + tcp_nat_type: NatType::Unknown as i32, + last_update_time: unix_timestamp(), + min_port: 100, + max_port: 200, + public_ip: vec!["127.0.0.1".to_owned(), "::1".to_owned()], + } + } + + async fn get_udp_port_mapping(&self, mut port: u16) -> anyhow::Result { + if port == 0 { + port = 40144; + } + Ok(SocketAddr::from(([127, 0, 0, 1], port))) + } + + async fn get_tcp_port_mapping(&self, mut port: u16) -> anyhow::Result { + if port == 0 { + port = 40144; + } + Ok(SocketAddr::from(([127, 0, 0, 1], port))) + } + + fn update_stun_info(&self) {} +} + +#[cfg(test)] +#[async_trait] +impl StunSocketMapper for MockStunInfoCollector { + async fn get_udp_port_mapping_with_socket( + &self, + socket: Arc, + ) -> anyhow::Result { + use easytier_core::socket::udp::VirtualUdpSocket as _; + self.get_udp_port_mapping(socket.local_addr()?.port()).await + } +} + +pub fn runtime_stun_info_collector(socket_context: SocketContext) -> StunInfoCollector { + let runtime = native_host_runtime(); + StunInfoCollector::new( + runtime.clone(), + runtime, + socket_context, + default_udp_stun_servers(), + default_tcp_stun_servers(), + default_udp_v6_stun_servers(), + ) +} + +#[cfg(test)] +fn unix_timestamp() -> i64 { + std::time::SystemTime::now() + .duration_since(std::time::UNIX_EPOCH) + .unwrap_or_default() + .as_secs() as i64 +} + #[cfg(test)] mod tests { - use crate::tunnel::{TunnelListener, udp::UdpTunnelListener}; - use tokio::time::{sleep, timeout}; + use std::time::Duration; + + use bytecodec::{DecodeExt as _, EncodeExt as _}; + use easytier_core::{ + connectivity::stun::{TcpNatTypeDetector, UdpNatTypeDetector}, + packet::stun::Attribute, + socket::SocketListener, + socket::{NetNamespace, udp::VirtualUdpSocketFactory}, + }; + use stun_codec::rfc5389::{attributes::XorMappedAddress, methods::BINDING}; + use stun_codec::{Message, MessageClass, MessageDecoder, MessageEncoder}; + use tokio::{ + io::{AsyncReadExt as _, AsyncWriteExt as _}, + task::JoinSet, + }; use tokio_util::task::AbortOnDropHandle; + use crate::{ + host_runtime::native_host_runtime, proto::rpc::standalone::runtime_udp_tunnel_listener, + }; + use super::*; #[tokio::test] - async fn test_udp_nat_type_detector() { - let collector = StunInfoCollector::new( - DEFAULT_UDP_STUN_SERVERS - .iter() - .map(ToString::to_string) - .collect(), - vec![], - vec![], - ); - collector.update_stun_info(); - loop { - let ret = collector.get_stun_info(); - if ret.udp_nat_type != NatType::Unknown as i32 { - println!("{:#?}", ret); - break; - } - tokio::time::sleep(Duration::from_secs(1)).await; - } + async fn runtime_collector_starts_with_native_runtime() { + let context = SocketContext::default() + .with_socket_mark(Some(0)) + .with_netns(Some(NetNamespace::new("instance-a"))); + let collector = runtime_stun_info_collector(context); - let port_mapping = collector.get_udp_port_mapping(3000).await; - println!("{:#?}", port_mapping); + assert_eq!( + StunInfoProvider::get_stun_info(&collector), + StunInfo::default() + ); + } + + #[test] + fn native_runtime_supports_collector_socket() { + fn assert_factory>() {} + assert_factory::(); } #[tokio::test] - async fn test_internal_stun_server() { - let mut udp_server1 = UdpTunnelListener::new("udp://0.0.0.0:55555".parse().unwrap()); - let mut udp_server2 = UdpTunnelListener::new("udp://0.0.0.0:55535".parse().unwrap()); - + async fn native_udp_runtime_drives_core_stun_detector() { + let mut first = runtime_udp_tunnel_listener( + "udp://127.0.0.1:0".parse().unwrap(), + "127.0.0.1:0".parse().unwrap(), + ); + let mut second = runtime_udp_tunnel_listener( + "udp://127.0.0.1:0".parse().unwrap(), + "127.0.0.1:0".parse().unwrap(), + ); + first.listen().await.unwrap(); + second.listen().await.unwrap(); + let servers = vec![ + SocketAddr::from(([127, 0, 0, 1], first.local_url().port().unwrap())).to_string(), + SocketAddr::from(([127, 0, 0, 1], second.local_url().port().unwrap())).to_string(), + ]; let mut tasks = JoinSet::new(); tasks.spawn(async move { - udp_server1.listen().await.unwrap(); loop { - udp_server1.accept().await.unwrap(); + first.accept().await.unwrap(); } }); tasks.spawn(async move { - udp_server2.listen().await.unwrap(); loop { - udp_server2.accept().await.unwrap(); + second.accept().await.unwrap(); } }); - let stun_servers = vec!["127.0.0.1:55555".to_string(), "127.0.0.1:55535".to_string()]; - let detector = UdpNatTypeDetector::new(stun_servers, 1); - let ret = detector.detect_nat_type(0).await; - println!("{:#?}, {:?}", ret, ret.as_ref().unwrap().nat_type()); - assert_eq!(ret.unwrap().nat_type(), NatType::Restricted); + let runtime = native_host_runtime(); + let detector = UdpNatTypeDetector::new( + runtime.clone(), + runtime, + SocketContext::default(), + servers, + 1, + ); + let result = detector.detect_nat_type(0).await.unwrap(); + + assert_eq!(result.nat_type(), NatType::Restricted); + tasks.abort_all(); } #[tokio::test] - async fn test_txt_public_stun_server() { - let stun_servers = vec!["txt:stun.easytier.cn".to_string()]; - let detector = UdpNatTypeDetector::new(stun_servers, 1); - timeout(Duration::from_secs(30), async { - loop { - let ret = detector.detect_nat_type(0).await; - println!("{:#?}, {:?}", ret, ret.as_ref().map(|x| x.nat_type())); - if let Ok(resp) = ret - && !resp.stun_resps.is_empty() - { - return; - } - sleep(Duration::from_secs(1)).await; - } - }) - .await - .expect("stun server should be available"); - } - - #[tokio::test] - #[ignore] - async fn test_public_tcp_stun_server_fitauto_ru() { - let stun_servers = vec![ - "stun.fitauto.ru".to_string(), - "stun.hot-chilli.net".to_string(), - ]; - let detector = TcpNatTypeDetector::new(stun_servers, 3); - let ret = detector.detect_nat_type(0).await; - println!("{:#?}, {:?}", ret, ret.as_ref().map(|x| x.nat_type())); - if let Ok(resp) = ret { - assert!(!resp.stun_resps.is_empty()); - } - } - - #[tokio::test] - async fn test_internal_tcp_stun_server_reuse_same_local_port() { - use stun_codec::rfc5389::attributes::XorMappedAddress; - use tokio::net::TcpListener; - - async fn spawn_tcp_stun_server() -> (SocketAddr, AbortOnDropHandle<()>) { - let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); - let server_addr = listener.local_addr().unwrap(); - + async fn native_tcp_runtime_drives_core_stun_detector() { + async fn spawn_server() -> (SocketAddr, AbortOnDropHandle<()>) { + let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); + let address = listener.local_addr().unwrap(); let task = tokio::spawn(async move { let (mut stream, peer_addr) = listener.accept().await.unwrap(); - - let req = TcpStunClient::tcp_read_stun_message(&mut stream, Duration::from_secs(2)) - .await + let mut header = [0_u8; 20]; + stream.read_exact(&mut header).await.unwrap(); + let payload_len = u16::from_be_bytes([header[2], header[3]]) as usize; + let mut bytes = vec![0_u8; 20 + payload_len]; + bytes[..20].copy_from_slice(&header); + stream.read_exact(&mut bytes[20..]).await.unwrap(); + let request = MessageDecoder::::new() + .decode_from_bytes(&bytes) + .unwrap() .unwrap(); - let mut resp_msg = Message::::new( + let mut response = Message::::new( MessageClass::SuccessResponse, BINDING, - req.transaction_id(), + request.transaction_id(), ); - resp_msg.add_attribute(Attribute::XorMappedAddress(XorMappedAddress::new( + response.add_attribute(Attribute::XorMappedAddress(XorMappedAddress::new( peer_addr, ))); - - let mut encoder = MessageEncoder::new(); - let rsp_buf = encoder.encode_into_bytes(resp_msg).unwrap(); - stream.write_all(rsp_buf.as_slice()).await.unwrap(); + let bytes = MessageEncoder::new().encode_into_bytes(response).unwrap(); + stream.write_all(&bytes).await.unwrap(); }); - - (server_addr, AbortOnDropHandle::new(task)) + (address, AbortOnDropHandle::new(task)) } - let (server1, _t1) = spawn_tcp_stun_server().await; - let (server2, _t2) = spawn_tcp_stun_server().await; - - let stun_servers = vec![server1.to_string(), server2.to_string()]; - let detector = TcpNatTypeDetector::new(stun_servers, 1); - - let ret = detector.detect_nat_type(0).await.unwrap(); - assert!(ret.stun_resps.len() >= 2); - - let local_ports = ret - .stun_resps - .iter() - .map(|x| x.local_addr.port()) - .collect::>(); - assert_eq!(local_ports.len(), 1); - - let mapped_ports = ret - .stun_resps - .iter() - .map(|x| x.mapped_socket_addr.unwrap().port()) - .collect::>(); - assert_eq!(mapped_ports.len(), 1); - assert_eq!( - local_ports.into_iter().next(), - mapped_ports.into_iter().next() + let (first, _first_task) = spawn_server().await; + let (second, _second_task) = spawn_server().await; + let runtime = native_host_runtime(); + let detector = TcpNatTypeDetector::new( + runtime.clone(), + runtime, + SocketContext::default(), + vec![first.to_string(), second.to_string()], + 1, ); - } - #[tokio::test] - async fn test_stun_info_collector_tcp_port_mapping() { - use stun_codec::rfc5389::attributes::XorMappedAddress; - use tokio::net::TcpListener; + let result = tokio::time::timeout(Duration::from_secs(5), detector.detect_nat_type(0)) + .await + .unwrap() + .unwrap(); - let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); - let server_addr = listener.local_addr().unwrap(); - - let _t = AbortOnDropHandle::new(tokio::spawn(async move { - for _ in 0..8 { - let Ok((mut stream, peer_addr)) = listener.accept().await else { - break; - }; - - let req = TcpStunClient::tcp_read_stun_message(&mut stream, Duration::from_secs(2)) - .await - .unwrap(); - let mut resp_msg = Message::::new( - MessageClass::SuccessResponse, - BINDING, - req.transaction_id(), - ); - resp_msg.add_attribute(Attribute::XorMappedAddress(XorMappedAddress::new( - peer_addr, - ))); - - let mut encoder = MessageEncoder::new(); - let rsp_buf = encoder.encode_into_bytes(resp_msg).unwrap(); - stream.write_all(rsp_buf.as_slice()).await.unwrap(); - } - })); - - let collector = StunInfoCollector::new(vec![], vec![server_addr.to_string()], vec![]); - collector.set_tcp_stun_servers(vec![server_addr.to_string()]); - let mapped = collector.get_tcp_port_mapping(0).await.unwrap(); - assert_eq!(mapped.ip(), IpAddr::V4(Ipv4Addr::LOCALHOST)); - assert!(mapped.port() > 0); - } - - #[tokio::test] - async fn test_v4_stun() { - let mut udp_server = UdpTunnelListener::new("udp://0.0.0.0:55355".parse().unwrap()); - let mut tasks = JoinSet::new(); - tasks.spawn(async move { - udp_server.listen().await.unwrap(); - loop { - udp_server.accept().await.unwrap(); - } - }); - let stun_servers = vec!["127.0.0.1:55355".to_string()]; - - let detector = UdpNatTypeDetector::new(stun_servers, 1); - let ret = detector.detect_nat_type(0).await; - println!("{:#?}, {:?}", ret, ret.as_ref().unwrap().nat_type()); - assert_eq!(ret.unwrap().nat_type(), NatType::Restricted); - } - - #[tokio::test] - async fn test_v6_stun() { - let mut udp_server = UdpTunnelListener::new("udp://[::]:55355".parse().unwrap()); - let mut tasks = JoinSet::new(); - tasks.spawn(async move { - udp_server.listen().await.unwrap(); - loop { - udp_server.accept().await.unwrap(); - } - }); - let stun_servers = vec!["::1:55355".to_string()]; - let ret = StunInfoCollector::get_public_ipv6(&stun_servers).await; - println!("{:#?}", ret); + assert_eq!(result.nat_type(), NatType::OpenInternet); + assert_eq!(result.usable_stun_resp_count(), 2); } } diff --git a/easytier/src/common/tracing_rolling_appender/mod.rs b/easytier/src/common/tracing_rolling_appender/mod.rs index 2074586b..8f7231f9 100644 --- a/easytier/src/common/tracing_rolling_appender/mod.rs +++ b/easytier/src/common/tracing_rolling_appender/mod.rs @@ -185,35 +185,18 @@ pub struct FileAppenderWrapper { appender: std::sync::Arc>, } -impl tracing_subscriber::fmt::MakeWriter<'_> for FileAppenderWrapper { - type Writer = FileAppenderWriter; - - fn make_writer(&self) -> Self::Writer { - FileAppenderWriter { - appender: self.appender.clone(), - } - } -} - impl FileAppenderWrapper { pub fn new(appender: RollingFileAppenderBase) -> Self { Self { appender: std::sync::Arc::new(parking_lot::Mutex::new(appender)), } } -} -#[derive(Debug, Clone)] -pub struct FileAppenderWriter { - appender: std::sync::Arc>, -} - -impl std::io::Write for FileAppenderWriter { - fn write(&mut self, buf: &[u8]) -> std::io::Result { - self.appender.lock().write(buf) + pub fn write_all(&self, buf: &[u8]) -> std::io::Result<()> { + self.appender.lock().write_all(buf) } - fn flush(&mut self) -> std::io::Result<()> { + pub fn flush(&self) -> std::io::Result<()> { self.appender.lock().flush() } } diff --git a/easytier/src/common/upnp.rs b/easytier/src/common/upnp.rs index d86bda8c..ff0aa71e 100644 --- a/easytier/src/common/upnp.rs +++ b/easytier/src/common/upnp.rs @@ -1,14 +1,15 @@ use std::{ fmt, net::{Ipv4Addr, SocketAddr, SocketAddrV4}, - sync::Arc, time::Duration, }; -#[cfg(test)] -use std::sync::atomic::{AtomicUsize, Ordering}; - use anyhow::{Context, anyhow, bail}; +use async_trait::async_trait; +use easytier_core::connectivity::hole_punch::port_mapping::{ + ActiveUdpPortMapping as CoreActiveUdpPortMapping, UdpPortMappingAttemptError, + UdpPortMappingBackend, UdpPortMappingLifecycle, +}; use igd_next::{ AddAnyPortError, PortMappingProtocol, SearchOptions, aio::{ @@ -19,52 +20,22 @@ use igd_next::{ use natpmp::{ Protocol as NatPmpProtocol, Response as NatPmpResponse, new_tokio_natpmp, new_tokio_natpmp_with, }; -use tokio::{net::UdpSocket, sync::oneshot}; -use super::{ - global_ctx::{ArcGlobalCtx, GlobalCtxEvent}, - stun::StunInfoCollectorTrait as _, -}; -use crate::tunnel::build_url_from_socket_addr; +use super::netns::NetNS; const UPNP_SEARCH_TIMEOUT: Duration = Duration::from_secs(1); const UPNP_SEARCH_RESPONSE_TIMEOUT: Duration = Duration::from_millis(300); const NAT_PMP_RESPONSE_TIMEOUT: Duration = Duration::from_secs(1); const UPNP_LEASE_DURATION_SECS: u32 = 300; -const UPNP_RENEW_INTERVAL: Duration = Duration::from_secs(240); const UPNP_DESCRIPTION: &str = "EasyTier udp hole punch"; -const PORT_MAPPING_BACKEND_NAT_PMP: &str = "nat-pmp"; -const PORT_MAPPING_BACKEND_IGD: &str = "igd"; type TokioGateway = Gateway; -#[cfg(test)] -static UDP_PORT_MAPPING_ATTEMPTS: AtomicUsize = AtomicUsize::new(0); - -#[cfg(test)] -pub(crate) fn reset_udp_port_mapping_attempts_for_test() { - UDP_PORT_MAPPING_ATTEMPTS.store(0, Ordering::Relaxed); -} - -#[cfg(test)] -pub(crate) fn udp_port_mapping_attempts_for_test() -> usize { - UDP_PORT_MAPPING_ATTEMPTS.load(Ordering::Relaxed) -} - enum PortMappingBackend { NatPmp { gateway: Ipv4Addr }, Igd { gateway: TokioGateway }, } -impl PortMappingBackend { - fn name(&self) -> &'static str { - match self { - Self::NatPmp { .. } => PORT_MAPPING_BACKEND_NAT_PMP, - Self::Igd { .. } => PORT_MAPPING_BACKEND_IGD, - } - } -} - struct ActiveUdpPortMapping { backend: PortMappingBackend, local_listener: url::Url, @@ -72,7 +43,25 @@ struct ActiveUdpPortMapping { gateway_external_port: u16, } +impl fmt::Debug for ActiveUdpPortMapping { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.debug_struct("ActiveUdpPortMapping") + .field("backend", &self.core_backend()) + .field("local_listener", &self.local_listener) + .field("local_addr", &self.local_addr) + .field("gateway_external_port", &self.gateway_external_port) + .finish() + } +} + impl ActiveUdpPortMapping { + fn core_backend(&self) -> UdpPortMappingBackend { + match self.backend { + PortMappingBackend::Igd { .. } => UdpPortMappingBackend::Igd, + PortMappingBackend::NatPmp { .. } => UdpPortMappingBackend::NatPmp, + } + } + async fn discover_nat_pmp_gateway( local_listener: &url::Url, ) -> anyhow::Result<(Ipv4Addr, SocketAddr)> { @@ -104,10 +93,10 @@ impl ActiveUdpPortMapping { } async fn discover_igd_gateway( - global_ctx: &ArcGlobalCtx, + net_ns: &NetNS, local_listener: &url::Url, ) -> anyhow::Result<(TokioGateway, SocketAddr)> { - let _g = global_ctx.net_ns.guard(); + let _g = net_ns.guard(); let gateway = search_gateway(SearchOptions { timeout: Some(UPNP_SEARCH_TIMEOUT), single_search_timeout: Some(UPNP_SEARCH_RESPONSE_TIMEOUT), @@ -142,10 +131,6 @@ impl ActiveUdpPortMapping { }) } - fn backend_name(&self) -> &'static str { - self.backend.name() - } - async fn renew(&self) -> anyhow::Result<()> { match &self.backend { PortMappingBackend::NatPmp { gateway } => { @@ -188,236 +173,98 @@ impl ActiveUdpPortMapping { } } -pub struct UdpPortMappingLease { - backend: &'static str, - gateway_external_port: u16, - stop_tx: Option>, -} - -impl UdpPortMappingLease { - pub fn backend(&self) -> &'static str { - self.backend +#[async_trait] +impl CoreActiveUdpPortMapping for ActiveUdpPortMapping { + fn backend(&self) -> UdpPortMappingBackend { + self.core_backend() } - pub fn gateway_external_port(&self) -> u16 { + fn local_addr(&self) -> SocketAddr { + self.local_addr + } + + fn gateway_external_port(&self) -> u16 { self.gateway_external_port } -} -impl fmt::Debug for UdpPortMappingLease { - fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { - f.debug_struct("UdpPortMappingLease") - .field("backend", &self.backend) - .field("gateway_external_port", &self.gateway_external_port) - .finish() + async fn renew(&self) -> anyhow::Result<()> { + ActiveUdpPortMapping::renew(self).await + } + + async fn remove(&self) -> anyhow::Result<()> { + ActiveUdpPortMapping::remove(self).await } } -impl Drop for UdpPortMappingLease { - fn drop(&mut self) { - if let Some(stop_tx) = self.stop_tx.take() { - let _ = stop_tx.send(()); +pub(crate) async fn establish_udp_port_mapping( + net_ns: NetNS, + backend: UdpPortMappingBackend, + local_listener: url::Url, +) -> Result, UdpPortMappingAttemptError> { + let mapping = match backend { + UdpPortMappingBackend::Igd => { + let (gateway, local_addr) = + discover_igd_gateway_in_netns(net_ns.clone(), local_listener.clone()) + .await + .map_err(UdpPortMappingAttemptError::discovery)?; + establish_igd_mapping_in_netns(net_ns, local_listener, gateway, local_addr) + .await + .map_err(UdpPortMappingAttemptError::establishment)? } - } -} - -pub async fn resolve_udp_public_addr( - global_ctx: ArcGlobalCtx, - local_listener: &url::Url, - socket: Arc, -) -> anyhow::Result<(SocketAddr, Option)> { - let port_mapping = match try_start_udp_port_mapping(&global_ctx, local_listener).await { - Ok(mapping) => mapping, - Err(err) => { - tracing::warn!( - ?err, - %local_listener, - "failed to establish udp port mapping, fallback to stun-only public addr resolution" - ); - None + UdpPortMappingBackend::NatPmp => { + let (gateway, local_addr) = + discover_nat_pmp_gateway_in_netns(net_ns.clone(), local_listener.clone()) + .await + .map_err(UdpPortMappingAttemptError::discovery)?; + establish_nat_pmp_mapping_in_netns(net_ns, local_listener, gateway, local_addr) + .await + .map_err(UdpPortMappingAttemptError::establishment)? } }; - - let mapped_addr = global_ctx - .get_stun_info_collector() - .get_udp_port_mapping_with_socket(socket) - .await - .map_err(anyhow::Error::from) - .with_context(|| format!("resolve udp public addr for {local_listener}"))?; - - if let Some(port_mapping) = port_mapping.as_ref() { - let mapped_listener = build_url_from_socket_addr(&mapped_addr.to_string(), "udp"); - global_ctx.issue_event(GlobalCtxEvent::ListenerPortMappingEstablished { - local_listener: local_listener.clone(), - mapped_listener, - backend: port_mapping.backend().to_string(), - }); - tracing::info!( - %local_listener, - backend = port_mapping.backend(), - gateway_external_port = port_mapping.gateway_external_port(), - stun_mapped_addr = %mapped_addr, - "udp public addr resolved after port mapping" - ); - } else { - tracing::debug!( - %local_listener, - stun_mapped_addr = %mapped_addr, - "udp public addr resolved without port mapping" - ); - } - - Ok((mapped_addr, port_mapping)) + Ok(Box::new(mapping)) } -async fn try_start_udp_port_mapping( - global_ctx: &ArcGlobalCtx, - local_listener: &url::Url, -) -> anyhow::Result> { - if global_ctx.get_flags().disable_upnp || !should_map_udp_listener(local_listener) { - return Ok(None); - } - - #[cfg(test)] - UDP_PORT_MAPPING_ATTEMPTS.fetch_add(1, Ordering::Relaxed); - - let mapping = discover_udp_port_mapping(global_ctx.clone(), local_listener.clone()).await?; - tracing::info!( - %local_listener, - backend = mapping.backend_name(), - local_addr = %mapping.local_addr, - gateway_external_port = mapping.gateway_external_port, - "udp port mapping established" - ); - - let backend = mapping.backend_name(); - let gateway_external_port = mapping.gateway_external_port; - let runtime_global_ctx = global_ctx.clone(); - let runtime_local_listener = local_listener.clone(); - let (stop_tx, stop_rx) = oneshot::channel(); - if should_run_port_mapping_in_dedicated_thread(&runtime_global_ctx) { +pub(crate) fn spawn_udp_port_mapping_lifecycle( + net_ns: NetNS, + local_listener: url::Url, + lifecycle: UdpPortMappingLifecycle, +) { + if should_run_port_mapping_in_dedicated_thread(&net_ns) { tokio::task::spawn_blocking(move || { - let _g = runtime_global_ctx.net_ns.guard(); + let _g = net_ns.guard(); match tokio::runtime::Builder::new_current_thread() .enable_all() .build() { - Ok(runtime) => { - runtime.block_on(run_udp_port_mapping_task( - runtime_local_listener, - mapping, - stop_rx, - )); - } - Err(err) => { - tracing::error!( - ?err, - %runtime_local_listener, - "failed to build runtime for udp port mapping renew task" - ); - } + Ok(runtime) => runtime.block_on(lifecycle), + Err(err) => tracing::error!( + ?err, + %local_listener, + "failed to build runtime for udp port mapping renew task" + ), } }); } else { - tokio::spawn(run_udp_port_mapping_task( - runtime_local_listener, - mapping, - stop_rx, - )); - } - - Ok(Some(UdpPortMappingLease { - backend, - gateway_external_port, - stop_tx: Some(stop_tx), - })) -} - -async fn discover_udp_port_mapping( - global_ctx: ArcGlobalCtx, - local_listener: url::Url, -) -> anyhow::Result { - match discover_igd_gateway_in_netns(global_ctx.clone(), local_listener.clone()).await { - Ok((gateway, local_addr)) => match establish_igd_mapping_in_netns( - global_ctx.clone(), - local_listener.clone(), - gateway, - local_addr, - ) - .await - { - Ok(mapping) => Ok(mapping), - Err(igd_err) => { - tracing::debug!( - ?igd_err, - %local_listener, - "igd udp port mapping failed, retry with nat-pmp" - ); - match discover_nat_pmp_gateway_in_netns(global_ctx.clone(), local_listener.clone()) - .await - { - Ok((gateway, local_addr)) => establish_nat_pmp_mapping_in_netns( - global_ctx, - local_listener.clone(), - gateway, - local_addr, - ) - .await - .map_err(|nat_pmp_err| { - anyhow!( - "udp port mapping failed for {local_listener}: igd error: {igd_err}; nat-pmp error: {nat_pmp_err}" - ) - }), - Err(nat_pmp_err) => Err(anyhow!( - "udp port mapping failed for {local_listener}: igd error: {igd_err}; nat-pmp discovery error: {nat_pmp_err}" - )), - } - } - }, - Err(igd_err) => { - tracing::debug!( - ?igd_err, - %local_listener, - "igd gateway discovery failed, retry with nat-pmp" - ); - match discover_nat_pmp_gateway_in_netns(global_ctx.clone(), local_listener.clone()).await - { - Ok((gateway, local_addr)) => establish_nat_pmp_mapping_in_netns( - global_ctx, - local_listener.clone(), - gateway, - local_addr, - ) - .await - .map_err(|nat_pmp_err| { - anyhow!( - "udp port mapping failed for {local_listener}: igd discovery error: {igd_err}; nat-pmp error: {nat_pmp_err}" - ) - }), - Err(nat_pmp_err) => Err(anyhow!( - "udp port mapping failed for {local_listener}: igd discovery error: {igd_err}; nat-pmp discovery error: {nat_pmp_err}" - )), - } - } + tokio::spawn(lifecycle); } } async fn discover_igd_gateway_in_netns( - global_ctx: ArcGlobalCtx, + net_ns: NetNS, local_listener: url::Url, ) -> anyhow::Result<(TokioGateway, SocketAddr)> { - if !should_run_port_mapping_in_dedicated_thread(&global_ctx) { - return ActiveUdpPortMapping::discover_igd_gateway(&global_ctx, &local_listener).await; + if !should_run_port_mapping_in_dedicated_thread(&net_ns) { + return ActiveUdpPortMapping::discover_igd_gateway(&net_ns, &local_listener).await; } tokio::task::spawn_blocking(move || { - let _g = global_ctx.net_ns.guard(); + let _g = net_ns.guard(); tokio::runtime::Builder::new_current_thread() .enable_all() .build() .context("build runtime for igd gateway discovery")? .block_on(ActiveUdpPortMapping::discover_igd_gateway( - &global_ctx, + &net_ns, &local_listener, )) }) @@ -426,17 +273,17 @@ async fn discover_igd_gateway_in_netns( } async fn establish_igd_mapping_in_netns( - global_ctx: ArcGlobalCtx, + net_ns: NetNS, local_listener: url::Url, gateway: TokioGateway, local_addr: SocketAddr, ) -> anyhow::Result { - if !should_run_port_mapping_in_dedicated_thread(&global_ctx) { + if !should_run_port_mapping_in_dedicated_thread(&net_ns) { return ActiveUdpPortMapping::establish_via_igd(&local_listener, gateway, local_addr).await; } tokio::task::spawn_blocking(move || { - let _g = global_ctx.net_ns.guard(); + let _g = net_ns.guard(); tokio::runtime::Builder::new_current_thread() .enable_all() .build() @@ -452,15 +299,15 @@ async fn establish_igd_mapping_in_netns( } async fn discover_nat_pmp_gateway_in_netns( - global_ctx: ArcGlobalCtx, + net_ns: NetNS, local_listener: url::Url, ) -> anyhow::Result<(Ipv4Addr, SocketAddr)> { - if !should_run_port_mapping_in_dedicated_thread(&global_ctx) { + if !should_run_port_mapping_in_dedicated_thread(&net_ns) { return ActiveUdpPortMapping::discover_nat_pmp_gateway(&local_listener).await; } tokio::task::spawn_blocking(move || { - let _g = global_ctx.net_ns.guard(); + let _g = net_ns.guard(); tokio::runtime::Builder::new_current_thread() .enable_all() .build() @@ -474,18 +321,18 @@ async fn discover_nat_pmp_gateway_in_netns( } async fn establish_nat_pmp_mapping_in_netns( - global_ctx: ArcGlobalCtx, + net_ns: NetNS, local_listener: url::Url, gateway: Ipv4Addr, local_addr: SocketAddr, ) -> anyhow::Result { - if !should_run_port_mapping_in_dedicated_thread(&global_ctx) { + if !should_run_port_mapping_in_dedicated_thread(&net_ns) { return ActiveUdpPortMapping::establish_via_nat_pmp(&local_listener, gateway, local_addr) .await; } tokio::task::spawn_blocking(move || { - let _g = global_ctx.net_ns.guard(); + let _g = net_ns.guard(); tokio::runtime::Builder::new_current_thread() .enable_all() .build() @@ -500,41 +347,8 @@ async fn establish_nat_pmp_mapping_in_netns( .context("join nat-pmp mapping establishment task")? } -async fn run_udp_port_mapping_task( - local_listener: url::Url, - mapping: ActiveUdpPortMapping, - mut stop_rx: oneshot::Receiver<()>, -) { - loop { - tokio::select! { - _ = tokio::time::sleep(UPNP_RENEW_INTERVAL) => { - if let Err(err) = mapping.renew().await { - tracing::warn!( - ?err, - %local_listener, - backend = mapping.backend_name(), - gateway_external_port = mapping.gateway_external_port, - "failed to renew udp port mapping" - ); - } - } - _ = &mut stop_rx => break, - } - } - - if let Err(err) = mapping.remove().await { - tracing::debug!( - ?err, - %local_listener, - backend = mapping.backend_name(), - gateway_external_port = mapping.gateway_external_port, - "failed to remove udp port mapping" - ); - } -} - -fn should_run_port_mapping_in_dedicated_thread(global_ctx: &ArcGlobalCtx) -> bool { - global_ctx.net_ns.name().is_some() +fn should_run_port_mapping_in_dedicated_thread(net_ns: &NetNS) -> bool { + net_ns.name().is_some() } async fn add_udp_mapping_port_igd( @@ -687,22 +501,6 @@ async fn remove_udp_mapping_nat_pmp( .with_context(|| format!("remove udp port mapping {local_listener}")) } -fn should_map_udp_listener(local_listener: &url::Url) -> bool { - if local_listener.scheme() != "udp" { - return false; - } - - let Some(host) = listener_ipv4_host(local_listener) else { - return false; - }; - - if host.is_loopback() || host.is_broadcast() { - return false; - } - - host.is_unspecified() || host.is_private() || host.is_link_local() -} - fn listener_ipv4_host(local_listener: &url::Url) -> Option { local_listener.host_str()?.parse().ok() } @@ -762,25 +560,3 @@ async fn remove_udp_mapping_igd( .await .with_context(|| format!("remove udp port mapping {local_listener}")) } - -#[cfg(test)] -mod tests { - #[test] - fn udp_mapping_requires_private_or_unspecified_ipv4_listener() { - assert!(super::should_map_udp_listener( - &"udp://0.0.0.0:11010".parse().unwrap() - )); - assert!(super::should_map_udp_listener( - &"udp://192.168.1.10:11010".parse().unwrap() - )); - assert!(!super::should_map_udp_listener( - &"udp://127.0.0.1:11010".parse().unwrap() - )); - assert!(!super::should_map_udp_listener( - &"udp://8.8.8.8:11010".parse().unwrap() - )); - assert!(!super::should_map_udp_listener( - &"tcp://0.0.0.0:11010".parse().unwrap() - )); - } -} diff --git a/easytier/src/connector/direct.rs b/easytier/src/connector/direct.rs deleted file mode 100644 index 1557d50a..00000000 --- a/easytier/src/connector/direct.rs +++ /dev/null @@ -1,1115 +0,0 @@ -// try connect peers directly, with either its public ip or lan ip - -use std::{ - collections::HashSet, - net::{IpAddr, Ipv4Addr, Ipv6Addr, SocketAddr}, - str::FromStr, - sync::{ - Arc, - atomic::{AtomicBool, Ordering}, - }, - time::Duration, -}; - -use quanta::Instant; - -use crate::{ - common::{ - PeerId, dns::socket_addrs, error::Error, global_ctx::ArcGlobalCtx, - stun::StunInfoCollectorTrait, - }, - connector::udp_hole_punch::handle_rpc_result, - peers::{ - peer_conn::PeerConnId, - peer_manager::PeerManager, - peer_rpc::PeerRpcManager, - peer_rpc_service::DirectConnectorManagerRpcServer, - peer_task::{PeerTaskLauncher, PeerTaskManager}, - }, - proto::{ - peer_rpc::{ - DirectConnectorRpc, DirectConnectorRpcClientFactory, DirectConnectorRpcServer, - GetIpListRequest, GetIpListResponse, SendUdpHolePunchPacketRequest, - }, - rpc_types::controller::BaseController, - }, - tunnel::{IpVersion, matches_protocol, udp::UdpTunnelConnector}, - use_global_var, -}; - -use super::{ - create_connector_by_url, should_background_p2p_with_peer, should_try_p2p_with_peer, - udp_hole_punch, -}; -use crate::tunnel::{FromUrl, IpScheme, TunnelScheme, matches_scheme}; -use anyhow::Context; -use rand::Rng; -use socket2::Protocol; -use tokio::{net::UdpSocket, task::JoinSet, time::timeout}; -use url::Host; - -pub const DIRECT_CONNECTOR_SERVICE_ID: u32 = 1; -pub const DIRECT_CONNECTOR_BLACKLIST_TIMEOUT_SEC: u64 = 300; -const MAX_IPV6_HOLE_PUNCH_CONNECTOR_ADDRS: usize = 16; - -static TESTING: AtomicBool = AtomicBool::new(false); - -fn mapped_listener_port(url: &url::Url) -> Option { - url.port().or_else(|| { - TunnelScheme::try_from(url) - .ok() - .and_then(|scheme| IpScheme::try_from(scheme).ok()) - .map(IpScheme::default_port) - }) -} - -async fn resolve_mapped_listener_addrs(listener: &url::Url) -> Result, Error> { - socket_addrs(listener, || mapped_listener_port(listener)).await -} - -fn is_usable_public_ipv6_candidate(ip: &Ipv6Addr, global_ctx: &ArcGlobalCtx) -> bool { - is_usable_public_ipv6_candidate_with_mode(ip, global_ctx, TESTING.load(Ordering::Relaxed)) -} - -fn is_usable_public_ipv6_candidate_with_mode( - ip: &Ipv6Addr, - global_ctx: &ArcGlobalCtx, - testing: bool, -) -> bool { - !global_ctx.is_ip_easytier_managed_ipv6(ip) - && (testing - || (!ip.is_loopback() - && !ip.is_unspecified() - && !ip.is_unique_local() - && !ip.is_unicast_link_local() - && !ip.is_multicast())) -} - -fn push_ipv6_hole_punch_candidate( - candidates: &mut Vec, - ip: Ipv6Addr, - global_ctx: &ArcGlobalCtx, - limit: usize, -) { - if candidates.len() >= limit - || !is_usable_public_ipv6_candidate(&ip, global_ctx) - || candidates.contains(&ip) - { - return; - } - candidates.push(ip); -} - -async fn collect_ipv6_hole_punch_candidates(global_ctx: &ArcGlobalCtx) -> Vec { - let mut candidates = Vec::new(); - for ip in global_ctx - .get_stun_info_collector() - .get_stun_info() - .public_ip - .iter() - .filter_map(|ip| ip.parse::().ok()) - { - push_ipv6_hole_punch_candidate( - &mut candidates, - ip, - global_ctx, - MAX_IPV6_HOLE_PUNCH_CONNECTOR_ADDRS, - ); - } - - let ip_list = global_ctx.get_ip_collector().collect_ip_addrs().await; - for ip in ip_list - .interface_ipv6s - .iter() - .chain(ip_list.public_ipv6.iter()) - .map(|ip| Ipv6Addr::from(*ip)) - { - push_ipv6_hole_punch_candidate( - &mut candidates, - ip, - global_ctx, - MAX_IPV6_HOLE_PUNCH_CONNECTOR_ADDRS, - ); - } - - candidates -} - -#[async_trait::async_trait] -pub trait PeerManagerForDirectConnector { - async fn list_peers(&self) -> Vec; - fn get_peer_rpc_mgr(&self) -> Arc; -} - -#[async_trait::async_trait] -impl PeerManagerForDirectConnector for PeerManager { - async fn list_peers(&self) -> Vec { - let mut ret = vec![]; - let allow_public_server = use_global_var!(DIRECT_CONNECT_TO_PUBLIC_SERVER); - let flags = self.get_global_ctx().get_flags(); - let lazy_p2p = flags.lazy_p2p; - let now = Instant::now(); - - let routes = self.list_routes().await; - for route in routes.iter() { - let static_allowed = should_background_p2p_with_peer( - route.feature_flag.as_ref(), - allow_public_server, - lazy_p2p, - flags.disable_p2p, - flags.need_p2p, - ); - let dynamic_allowed = should_try_p2p_with_peer( - route.feature_flag.as_ref(), - allow_public_server, - flags.disable_p2p, - flags.need_p2p, - ) && self.has_recent_traffic(route.peer_id, now); - if static_allowed || dynamic_allowed { - ret.push(route.peer_id); - } - } - - ret - } - - fn get_peer_rpc_mgr(&self) -> Arc { - self.get_peer_rpc_mgr() - } -} - -#[derive(Hash, Eq, PartialEq, Clone)] -struct DstBlackListItem(PeerId, String); - -#[derive(Hash, Eq, PartialEq, Clone)] -struct DstListenerUrlBlackListItem(PeerId, String); - -struct DirectConnectorManagerData { - global_ctx: ArcGlobalCtx, - peer_manager: Arc, - dst_listener_blacklist: timedmap::TimedMap, - peer_black_list: timedmap::TimedMap, -} - -impl DirectConnectorManagerData { - pub fn new(global_ctx: ArcGlobalCtx, peer_manager: Arc) -> Self { - Self { - global_ctx, - peer_manager, - dst_listener_blacklist: timedmap::TimedMap::new(), - peer_black_list: timedmap::TimedMap::new(), - } - } - - async fn remote_send_udp_hole_punch_packet( - &self, - dst_peer_id: PeerId, - connector_addrs: Vec, - preferred_src_ipv6: Option, - remote_url: &url::Url, - ) -> Result<(), Error> { - if !matches_scheme!(remote_url, TunnelScheme::Ip(IpScheme::Udp)) { - return Err(anyhow::anyhow!( - "udp hole punch packet only applies to udp listener: {}", - remote_url - ) - .into()); - } - - let global_ctx = self.peer_manager.get_global_ctx(); - let listener_port = mapped_listener_port(remote_url).ok_or(anyhow::anyhow!( - "failed to parse port from remote url: {}", - remote_url - ))?; - - let rpc_stub = self - .peer_manager - .get_peer_rpc_mgr() - .rpc_client() - .scoped_client::>( - self.peer_manager.my_peer_id(), - dst_peer_id, - global_ctx.get_network_name(), - ); - - rpc_stub - .send_udp_hole_punch_packet( - BaseController::default(), - SendUdpHolePunchPacketRequest { - connector_addr: connector_addrs.first().copied().map(Into::into), - listener_port: listener_port as u32, - preferred_src_ipv6: preferred_src_ipv6.map(Into::into), - connector_addrs: connector_addrs.into_iter().map(Into::into).collect(), - }, - ) - .await - .with_context(|| { - format!( - "do rpc, send udp hole punch packet to peer {} at {} with preferred source {:?}", - dst_peer_id, remote_url, preferred_src_ipv6 - ) - })?; - - Ok(()) - } - - async fn connect_to_public_ipv6( - &self, - dst_peer_id: PeerId, - remote_url: &url::Url, - ) -> Result<(PeerId, PeerConnId), Error> { - let local_socket = Arc::new( - UdpSocket::bind("[::]:0") - .await - .with_context(|| format!("failed to bind local socket for {}", remote_url))?, - ); - let connector_ips = collect_ipv6_hole_punch_candidates(&self.global_ctx).await; - - // ask remote to send v6 hole punch packet - // and no matter what the result is, continue to connect - if !connector_ips.is_empty() { - let local_port = local_socket.local_addr()?.port(); - let connector_addrs = connector_ips - .into_iter() - .map(|ip| SocketAddr::new(IpAddr::V6(ip), local_port)) - .collect::>(); - let preferred_src_ipv6 = match remote_url.host() { - Some(Host::Ipv6(ip)) => Some(ip), - _ => None, - }; - tracing::debug!( - ?connector_addrs, - ?preferred_src_ipv6, - ?remote_url, - "request remote IPv6 hole-punch packets" - ); - if let Err(err) = self - .remote_send_udp_hole_punch_packet( - dst_peer_id, - connector_addrs, - preferred_src_ipv6, - remote_url, - ) - .await - { - tracing::debug!( - ?err, - ?remote_url, - "remote IPv6 hole-punch packet request failed" - ); - } - } else { - tracing::debug!( - ?remote_url, - "skip remote IPv6 hole-punch packet; no non-EasyTier public IPv6 in STUN info" - ); - } - - let udp_connector = UdpTunnelConnector::new(remote_url.clone()); - let remote_addr = SocketAddr::from_url(remote_url.clone(), IpVersion::V6).await?; - let ret = udp_connector - .try_connect_with_socket(local_socket, remote_addr) - .await?; - - // NOTICE: must add as directly connected tunnel - self.peer_manager - .add_client_tunnel_with_peer_id_hint(ret, true, Some(dst_peer_id)) - .await - } - - async fn connect_to_public_ipv4( - &self, - dst_peer_id: PeerId, - remote_url: &url::Url, - ) -> Result<(PeerId, PeerConnId), Error> { - let local_socket = { - let _g = self.global_ctx.net_ns.guard(); - Arc::new( - UdpSocket::bind("0.0.0.0:0") - .await - .with_context(|| format!("failed to bind local socket for {}", remote_url))?, - ) - }; - let connector_addr = self - .peer_manager - .get_global_ctx() - .get_stun_info_collector() - .get_udp_port_mapping_with_socket(local_socket.clone()) - .await - .with_context(|| format!("failed to get udp port mapping for {}", remote_url))?; - - let _ = self - .remote_send_udp_hole_punch_packet(dst_peer_id, vec![connector_addr], None, remote_url) - .await; - - let udp_connector = UdpTunnelConnector::new(remote_url.clone()); - let remote_addr = SocketAddr::from_url(remote_url.clone(), IpVersion::V4).await?; - let ret = udp_connector - .try_connect_with_socket(local_socket, remote_addr) - .await?; - - self.peer_manager - .add_client_tunnel_with_peer_id_hint(ret, true, Some(dst_peer_id)) - .await - } - - async fn do_try_connect_to_ip(&self, dst_peer_id: PeerId, addr: String) -> Result<(), Error> { - let connector = create_connector_by_url(&addr, &self.global_ctx, IpVersion::Both).await?; - let remote_url = connector.remote_url(); - let (peer_id, conn_id) = if matches_scheme!(remote_url, TunnelScheme::Ip(IpScheme::Udp)) { - match remote_url.host() { - Some(Host::Ipv6(_)) => { - self.connect_to_public_ipv6(dst_peer_id, &remote_url) - .await? - } - Some(Host::Ipv4(ip)) if is_public_ipv4(ip) => { - match self.connect_to_public_ipv4(dst_peer_id, &remote_url).await { - Ok(ret) => ret, - Err(err) => { - tracing::debug!( - ?err, - %remote_url, - "udp public ipv4 listener punch failed, falling back to direct connect" - ); - timeout( - std::time::Duration::from_secs(3), - self.peer_manager.try_direct_connect_with_peer_id_hint( - connector, - Some(dst_peer_id), - ), - ) - .await?? - } - } - } - _ => { - timeout( - std::time::Duration::from_secs(3), - self.peer_manager - .try_direct_connect_with_peer_id_hint(connector, Some(dst_peer_id)), - ) - .await?? - } - } - } else { - timeout( - std::time::Duration::from_secs(3), - self.peer_manager - .try_direct_connect_with_peer_id_hint(connector, Some(dst_peer_id)), - ) - .await?? - }; - - if peer_id != dst_peer_id && !TESTING.load(Ordering::Relaxed) { - tracing::info!( - "connect to ip succ: {}, but peer id mismatch, expect: {}, actual: {}", - addr, - dst_peer_id, - peer_id - ); - self.peer_manager.close_peer_conn(peer_id, &conn_id).await?; - return Err(Error::InvalidUrl(addr)); - } - - Ok(()) - } - - #[tracing::instrument(skip(self))] - async fn try_connect_to_ip( - self: Arc, - dst_peer_id: PeerId, - addr: String, - ) -> Result<(), Error> { - let mut rand_gen = rand::rngs::OsRng; - let backoff_ms = [1000, 2000, 4000]; - let mut backoff_idx = 0; - - tracing::debug!(?dst_peer_id, ?addr, "try_connect_to_ip start"); - - self.dst_listener_blacklist.cleanup(); - - if self - .dst_listener_blacklist - .contains(&DstListenerUrlBlackListItem(dst_peer_id, addr.clone())) - { - return Err(Error::UrlInBlacklist); - } - - loop { - if self.peer_manager.has_directly_connected_conn(dst_peer_id) { - return Ok(()); - } - - tracing::debug!(?dst_peer_id, ?addr, "try_connect_to_ip start one round"); - let ret = self.do_try_connect_to_ip(dst_peer_id, addr.clone()).await; - tracing::debug!(?ret, ?dst_peer_id, ?addr, "try_connect_to_ip return"); - if ret.is_ok() { - return Ok(()); - } - - if self.peer_manager.has_directly_connected_conn(dst_peer_id) { - return Ok(()); - } - - if backoff_idx < backoff_ms.len() { - let delta = backoff_ms[backoff_idx] >> 1; - assert!(delta > 0); - assert!(delta < backoff_ms[backoff_idx]); - - tokio::time::sleep(Duration::from_millis( - (backoff_ms[backoff_idx] + rand_gen.gen_range(-delta..delta)) as u64, - )) - .await; - - backoff_idx += 1; - continue; - } else { - self.dst_listener_blacklist.insert( - DstListenerUrlBlackListItem(dst_peer_id, addr), - (), - std::time::Duration::from_secs(DIRECT_CONNECTOR_BLACKLIST_TIMEOUT_SEC), - ); - return ret; - } - } - } - - async fn spawn_direct_connect_task( - self: &Arc, - dst_peer_id: PeerId, - ip_list: &GetIpListResponse, - listener: &url::Url, - tasks: &mut JoinSet>, - ) { - let Ok(mut addrs) = resolve_mapped_listener_addrs(listener).await else { - tracing::error!(?listener, "failed to parse socket address from listener"); - return; - }; - let listener_host = addrs.pop(); - tracing::info!(?listener_host, ?listener, "try direct connect to peer"); - - let is_udp = matches_protocol!(listener, Protocol::UDP); - // Snapshot running listeners once; used for cheap port pre-checks before the - // expensive should_deny_proxy call (which binds a socket per IP) in the - // unspecified-address expansion loops below. - let local_listeners = self.global_ctx.get_running_listeners(); - let port_has_local_listener = |port: u16| -> bool { - local_listeners - .iter() - .any(|l| l.port() == Some(port) && matches_protocol!(l, Protocol::UDP) == is_udp) - }; - - match listener_host { - Some(SocketAddr::V4(s_addr)) => { - if s_addr.ip().is_unspecified() { - // Only pay the should_deny_proxy cost (bind per IP) when a local - // listener actually uses this port+protocol; otherwise the check - // can never return true. - let check_self = port_has_local_listener(s_addr.port()); - ip_list - .interface_ipv4s - .iter() - .chain(ip_list.public_ipv4.iter()) - .for_each(|ip| { - let sock_addr = SocketAddr::new( - IpAddr::V4(std::net::Ipv4Addr::from(ip.addr)), - s_addr.port(), - ); - if check_self && self.global_ctx.should_deny_proxy(&sock_addr, is_udp) { - tracing::debug!( - ?ip, - ?listener, - "skip self-connection (0.0.0.0 expansion)" - ); - return; - } - let mut addr = (*listener).clone(); - if addr.set_host(Some(ip.to_string().as_str())).is_ok() { - tasks.spawn(Self::try_connect_to_ip( - self.clone(), - dst_peer_id, - addr.to_string(), - )); - } else { - tracing::error!( - ?ip, - ?listener, - ?dst_peer_id, - "failed to set host for interface ipv4" - ); - } - }); - } else if !s_addr.ip().is_loopback() || TESTING.load(Ordering::Relaxed) { - if self - .global_ctx - .should_deny_proxy(&SocketAddr::from(s_addr), is_udp) - { - tracing::debug!(?listener, "skip self-connection (specific IPv4)"); - } else { - tasks.spawn(Self::try_connect_to_ip( - self.clone(), - dst_peer_id, - listener.to_string(), - )); - } - } - } - Some(SocketAddr::V6(s_addr)) => { - if s_addr.ip().is_unspecified() { - // for ipv6, only try public ip - // Same port pre-check as IPv4: avoid binding per IP when no local - // listener uses this port+protocol. - let check_self = port_has_local_listener(s_addr.port()); - ip_list - .interface_ipv6s - .iter() - .chain(ip_list.public_ipv6.iter()) - .filter_map(|x| Ipv6Addr::from_str(&x.to_string()).ok()) - .filter(|x| is_usable_public_ipv6_candidate(x, &self.global_ctx)) - .collect::>() - .iter() - .for_each(|ip| { - let sock_addr = SocketAddr::new(IpAddr::V6(*ip), s_addr.port()); - if check_self && self.global_ctx.should_deny_proxy(&sock_addr, is_udp) { - tracing::debug!( - ?ip, - ?listener, - "skip self-connection (:: expansion)" - ); - return; - } - let mut addr = (*listener).clone(); - if addr.set_host(Some(format!("[{}]", ip).as_str())).is_ok() { - tasks.spawn(Self::try_connect_to_ip( - self.clone(), - dst_peer_id, - addr.to_string(), - )); - } else { - tracing::error!( - ?ip, - ?listener, - ?dst_peer_id, - "failed to set host for public ipv6" - ); - } - }); - } else if self.global_ctx.is_ip_easytier_managed_ipv6(s_addr.ip()) { - tracing::debug!( - ?listener, - "skip EasyTier-managed IPv6 as direct-connect target" - ); - } else if !s_addr.ip().is_loopback() || TESTING.load(Ordering::Relaxed) { - if self - .global_ctx - .should_deny_proxy(&SocketAddr::from(s_addr), is_udp) - { - tracing::debug!(?listener, "skip self-connection (specific IPv6)"); - } else { - tasks.spawn(Self::try_connect_to_ip( - self.clone(), - dst_peer_id, - listener.to_string(), - )); - } - } - } - p => { - tracing::error!(?p, ?listener, "failed to parse ip version from listener"); - } - } - } - - #[tracing::instrument(skip(self))] - async fn do_try_direct_connect_internal( - self: &Arc, - dst_peer_id: PeerId, - ip_list: GetIpListResponse, - ) -> Result<(), Error> { - let enable_ipv6 = self.global_ctx.get_flags().enable_ipv6; - let available_listeners = ip_list - .listeners - .clone() - .into_iter() - .map(Into::::into) - .filter_map(|l| if l.scheme() != "ring" { Some(l) } else { None }) - .filter(|l| mapped_listener_port(l).is_some() && l.host().is_some()) - .filter(|l| enable_ipv6 || !matches!(l.host().unwrap().to_owned(), Host::Ipv6(_))) - .collect::>(); - - tracing::debug!(?available_listeners, "got available listeners"); - - if available_listeners.is_empty() { - return Err(anyhow::anyhow!("peer {} have no valid listener", dst_peer_id).into()); - } - - let default_protocol = self.global_ctx.get_flags().default_protocol; - // sort available listeners, default protocol has the highest priority, udp is second, others just random - // highest priority is in the last - let mut available_listeners = available_listeners; - available_listeners.sort_by_key(|l| { - let scheme = l.scheme(); - if scheme == default_protocol { - 3 - } else if scheme == "udp" { - 2 - } else { - 1 - } - }); - - while !available_listeners.is_empty() { - let mut tasks = JoinSet::new(); - let mut listener_list = vec![]; - - let cur_scheme = available_listeners.last().unwrap().scheme().to_owned(); - while let Some(listener) = available_listeners.last() { - if listener.scheme() != cur_scheme { - break; - } - - tracing::debug!("try direct connect to peer with listener: {}", listener); - self.spawn_direct_connect_task(dst_peer_id, &ip_list, listener, &mut tasks) - .await; - - listener_list.push(listener.clone().to_string()); - available_listeners.pop(); - } - - let ret = tasks.join_all().await; - tracing::debug!( - ?ret, - ?dst_peer_id, - ?cur_scheme, - ?listener_list, - "all tasks finished for current scheme" - ); - - if self.peer_manager.has_directly_connected_conn(dst_peer_id) { - tracing::info!( - "direct connect to peer {} success, has direct conn", - dst_peer_id - ); - return Ok(()); - } - } - - Ok(()) - } - - #[tracing::instrument(skip(self))] - async fn do_try_direct_connect( - self: Arc, - dst_peer_id: PeerId, - ) -> Result<(), Error> { - let mut backoff = - udp_hole_punch::BackOff::new(vec![1000, 2000, 2000, 5000, 5000, 10000, 30000, 60000]); - let mut attempt = 0; - loop { - if self.peer_black_list.contains(&dst_peer_id) { - return Err(anyhow::anyhow!("peer {} is blacklisted", dst_peer_id).into()); - } - - if attempt > 0 { - tokio::time::sleep(Duration::from_millis(backoff.next_backoff())).await; - } - attempt += 1; - - let peer_manager = self.peer_manager.clone(); - tracing::debug!("try direct connect to peer: {}", dst_peer_id); - - let rpc_stub = peer_manager - .get_peer_rpc_mgr() - .rpc_client() - .scoped_client::>( - peer_manager.my_peer_id(), - dst_peer_id, - self.global_ctx.get_network_name(), - ); - - let ip_list = rpc_stub - .get_ip_list(BaseController::default(), GetIpListRequest {}) - .await; - let ip_list = handle_rpc_result(ip_list, dst_peer_id, &self.peer_black_list) - .with_context(|| format!("get ip list from peer {}", dst_peer_id))?; - - tracing::info!(ip_list = ?ip_list, dst_peer_id = ?dst_peer_id, "got ip list"); - - let ret = self - .do_try_direct_connect_internal(dst_peer_id, ip_list) - .await; - tracing::info!(?ret, ?dst_peer_id, "do_try_direct_connect return"); - - if peer_manager.has_directly_connected_conn(dst_peer_id) { - tracing::info!( - "direct connect to peer {} success, has direct conn", - dst_peer_id - ); - return Ok(()); - } - } - } -} - -fn is_public_ipv4(ip: Ipv4Addr) -> bool { - !ip.is_private() - && !ip.is_loopback() - && !ip.is_link_local() - && !ip.is_broadcast() - && !ip.is_unspecified() -} - -impl std::fmt::Debug for DirectConnectorManagerData { - fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { - f.debug_struct("DirectConnectorManagerData") - .field("peer_manager", &self.peer_manager) - .finish() - } -} - -pub struct DirectConnectorManager { - global_ctx: ArcGlobalCtx, - data: Arc, - client: PeerTaskManager, - tasks: JoinSet<()>, -} - -#[derive(Clone)] -struct DirectConnectorLauncher(Arc); - -#[async_trait::async_trait] -impl PeerTaskLauncher for DirectConnectorLauncher { - type Data = Arc; - type CollectPeerItem = PeerId; - type TaskRet = (); - - fn new_data(&self, _peer_mgr: Arc) -> Self::Data { - self.0.clone() - } - - async fn collect_peers_need_task(&self, data: &Self::Data) -> Vec { - data.peer_black_list.cleanup(); - let my_peer_id = data.peer_manager.my_peer_id(); - data.peer_manager - .list_peers() - .await - .into_iter() - .filter(|peer_id| { - *peer_id != my_peer_id - && !data.peer_manager.has_directly_connected_conn(*peer_id) - && !data.peer_black_list.contains(peer_id) - }) - .collect() - } - - async fn launch_task( - &self, - data: &Self::Data, - item: Self::CollectPeerItem, - ) -> tokio::task::JoinHandle> { - let data = data.clone(); - tokio::spawn(async move { data.do_try_direct_connect(item).await.map_err(Into::into) }) - } - - async fn all_task_done(&self, _data: &Self::Data) {} - - fn loop_interval_ms(&self) -> u64 { - 5000 - } -} - -impl DirectConnectorManager { - pub fn new(global_ctx: ArcGlobalCtx, peer_manager: Arc) -> Self { - let data = Arc::new(DirectConnectorManagerData::new( - global_ctx.clone(), - peer_manager.clone(), - )); - let client = PeerTaskManager::new_with_external_signal( - DirectConnectorLauncher(data.clone()), - peer_manager.clone(), - Some(peer_manager.p2p_demand_notify()), - ); - Self { - global_ctx, - data, - client, - tasks: JoinSet::new(), - } - } - - pub fn run(&mut self) { - self.run_as_server(); - self.run_as_client(); - } - - pub fn run_as_server(&mut self) { - self.data - .peer_manager - .get_peer_rpc_mgr() - .rpc_server() - .registry() - .register( - DirectConnectorRpcServer::new(DirectConnectorManagerRpcServer::new( - self.global_ctx.clone(), - )), - &self.data.global_ctx.get_network_name(), - ); - } - - pub fn run_as_client(&mut self) { - self.client.start(); - } - - #[cfg(test)] - pub(crate) async fn try_direct_connect_with_ip_list( - &self, - dst_peer_id: PeerId, - ip_list: GetIpListResponse, - ) -> Result<(), Error> { - self.data - .do_try_direct_connect_internal(dst_peer_id, ip_list) - .await - } -} - -#[cfg(test)] -mod tests { - use std::{collections::BTreeSet, sync::Arc}; - - use crate::{ - common::global_ctx::tests::get_mock_global_ctx, - connector::direct::{ - DirectConnectorManager, DirectConnectorManagerData, DstListenerUrlBlackListItem, - }, - instance::listeners::ListenerManager, - peers::tests::{ - connect_peer_manager, create_mock_peer_manager, wait_route_appear, - wait_route_appear_with_cost, - }, - proto::peer_rpc::GetIpListResponse, - tunnel::{IpScheme, TunnelScheme, matches_scheme}, - }; - - use std::net::{IpAddr, Ipv4Addr, Ipv6Addr, SocketAddr}; - - use super::{TESTING, mapped_listener_port, resolve_mapped_listener_addrs}; - - #[tokio::test] - async fn public_ipv6_candidate_rejects_easytier_managed_addr_even_in_tests() { - let global_ctx = get_mock_global_ctx(); - let managed_ipv6: cidr::Ipv6Inet = "2001:db8::2/128".parse().unwrap(); - global_ctx.set_public_ipv6_routes(BTreeSet::from([managed_ipv6])); - - assert!(!super::is_usable_public_ipv6_candidate_with_mode( - &"2001:db8::2".parse().unwrap(), - &global_ctx, - true, - )); - assert!(super::is_usable_public_ipv6_candidate_with_mode( - &"::1".parse().unwrap(), - &global_ctx, - true, - )); - } - - #[tokio::test] - async fn ipv6_hole_punch_candidates_are_deduped_filtered_and_capped() { - let global_ctx = get_mock_global_ctx(); - let managed_ipv6: cidr::Ipv6Inet = "2001:db8::2/128".parse().unwrap(); - global_ctx.set_public_ipv6_routes(BTreeSet::from([managed_ipv6])); - - let first: Ipv6Addr = "2001:db8::1".parse().unwrap(); - let managed = managed_ipv6.address(); - let second: Ipv6Addr = "2001:db8::3".parse().unwrap(); - let third: Ipv6Addr = "2001:db8::4".parse().unwrap(); - let mut candidates = Vec::new(); - - super::push_ipv6_hole_punch_candidate(&mut candidates, first, &global_ctx, 2); - super::push_ipv6_hole_punch_candidate(&mut candidates, first, &global_ctx, 2); - super::push_ipv6_hole_punch_candidate(&mut candidates, managed, &global_ctx, 2); - super::push_ipv6_hole_punch_candidate(&mut candidates, second, &global_ctx, 2); - super::push_ipv6_hole_punch_candidate(&mut candidates, third, &global_ctx, 2); - - assert_eq!(candidates, vec![first, second]); - } - - #[test] - fn udp_ipv6_url_matches_hole_punch_branch_condition() { - let remote_url: url::Url = "udp://[2001:db8::1]:11010".parse().unwrap(); - let takes_udp_ipv6_hole_punch_branch = - matches_scheme!(remote_url, TunnelScheme::Ip(IpScheme::Udp)) - && matches!(remote_url.host(), Some(url::Host::Ipv6(_))); - - assert!(takes_udp_ipv6_hole_punch_branch); - } - - #[test] - fn mapped_listener_port_uses_ip_scheme_defaults() { - assert_eq!( - mapped_listener_port(&"ws://example.com".parse().unwrap()), - Some(80) - ); - assert_eq!( - mapped_listener_port(&"wss://example.com".parse().unwrap()), - Some(443) - ); - assert_eq!( - mapped_listener_port(&"tcp://127.0.0.1".parse().unwrap()), - Some(11010) - ); - assert_eq!( - mapped_listener_port(&"udp://127.0.0.1".parse().unwrap()), - Some(11010) - ); - } - - #[tokio::test] - async fn resolve_mapped_listener_addrs_uses_default_ports() { - let wss_addrs = resolve_mapped_listener_addrs(&"wss://127.0.0.1".parse().unwrap()) - .await - .unwrap(); - assert_eq!( - wss_addrs, - vec![SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 443)] - ); - - let tcp_addrs = resolve_mapped_listener_addrs(&"tcp://127.0.0.1".parse().unwrap()) - .await - .unwrap(); - assert_eq!( - tcp_addrs, - vec![SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 11010)] - ); - } - - async fn run_direct_connector_mapped_listener_test( - mapped_listener: &str, - target_listener: &str, - ) { - TESTING.store(true, std::sync::atomic::Ordering::Relaxed); - let p_a = create_mock_peer_manager().await; - let p_b = create_mock_peer_manager().await; - let p_c = create_mock_peer_manager().await; - let p_x = create_mock_peer_manager().await; - connect_peer_manager(p_a.clone(), p_b.clone()).await; - connect_peer_manager(p_b.clone(), p_c.clone()).await; - connect_peer_manager(p_c.clone(), p_x.clone()).await; - - wait_route_appear(p_a.clone(), p_c.clone()).await.unwrap(); - wait_route_appear(p_a.clone(), p_x.clone()).await.unwrap(); - - let mut f = p_a.get_global_ctx().get_flags(); - f.bind_device = false; - p_a.get_global_ctx().set_flags(f); - - p_c.get_global_ctx() - .config - .set_mapped_listeners(Some(vec![mapped_listener.parse().unwrap()])); - - p_x.get_global_ctx() - .config - .set_listeners(vec![target_listener.parse().unwrap()]); - let mut lis_x = ListenerManager::new(p_x.get_global_ctx(), p_x.clone()); - lis_x.prepare_listeners().await.unwrap(); - lis_x.run().await.unwrap(); - - tokio::time::sleep(std::time::Duration::from_secs(1)).await; - let mut dm_a = DirectConnectorManager::new(p_a.get_global_ctx(), p_a.clone()); - let mut dm_c = DirectConnectorManager::new(p_c.get_global_ctx(), p_c.clone()); - dm_a.run_as_client(); - dm_c.run_as_server(); - // p_c's mapped listener is p_x's listener, so p_a should connect to p_x directly - - wait_route_appear_with_cost(p_a.clone(), p_x.my_peer_id(), Some(1)) - .await - .unwrap(); - } - - #[tokio::test] - async fn direct_connector_mapped_listener() { - run_direct_connector_mapped_listener_test("tcp://127.0.0.1:11334", "tcp://0.0.0.0:11334") - .await; - } - - #[rstest::rstest] - #[tokio::test] - async fn direct_connector_basic_test( - #[values("tcp", "udp", "wg")] proto: &str, - #[values("true", "false")] ipv6: bool, - ) { - TESTING.store(true, std::sync::atomic::Ordering::Relaxed); - - let p_a = create_mock_peer_manager().await; - let p_b = create_mock_peer_manager().await; - let p_c = create_mock_peer_manager().await; - connect_peer_manager(p_a.clone(), p_b.clone()).await; - connect_peer_manager(p_b.clone(), p_c.clone()).await; - - wait_route_appear(p_a.clone(), p_c.clone()).await.unwrap(); - - p_c.get_global_ctx() - .get_ip_collector() - .collect_ip_addrs() - .await; - - tokio::time::sleep(std::time::Duration::from_secs(4)).await; - - let mut dm_a = DirectConnectorManager::new(p_a.get_global_ctx(), p_a.clone()); - let mut dm_c = DirectConnectorManager::new(p_c.get_global_ctx(), p_c.clone()); - - dm_a.run_as_client(); - dm_c.run_as_server(); - - let port = if proto == "wg" { 11040 } else { 11041 }; - if !ipv6 { - p_c.get_global_ctx().config.set_listeners(vec![ - format!("{}://0.0.0.0:{}", proto, port).parse().unwrap(), - ]); - } else { - p_c.get_global_ctx() - .config - .set_listeners(vec![format!("{}://[::]:{}", proto, port).parse().unwrap()]); - } - let mut f = p_c.get_global_ctx().config.get_flags(); - f.enable_ipv6 = ipv6; - p_c.get_global_ctx().set_flags(f); - let mut lis_c = ListenerManager::new(p_c.get_global_ctx(), p_c.clone()); - lis_c.prepare_listeners().await.unwrap(); - - lis_c.run().await.unwrap(); - - wait_route_appear_with_cost(p_a.clone(), p_c.my_peer_id(), Some(1)) - .await - .unwrap(); - } - - #[tokio::test] - async fn direct_connector_scheme_blacklist() { - TESTING.store(true, std::sync::atomic::Ordering::Relaxed); - let p_a = create_mock_peer_manager().await; - let data = Arc::new(DirectConnectorManagerData::new( - p_a.get_global_ctx(), - p_a.clone(), - )); - let mut ip_list = GetIpListResponse::default(); - ip_list - .listeners - .push("tcp://127.0.0.1:10222".parse().unwrap()); - - ip_list - .interface_ipv4s - .push("127.0.0.1".parse::().unwrap().into()); - - data.do_try_direct_connect_internal(1, ip_list.clone()) - .await - .unwrap(); - - assert!( - data.dst_listener_blacklist - .contains(&DstListenerUrlBlackListItem( - 1, - "tcp://127.0.0.1:10222".parse().unwrap() - )) - ); - } -} diff --git a/easytier/src/connector/dns_connector.rs b/easytier/src/connector/dns_connector.rs deleted file mode 100644 index c41a5a53..00000000 --- a/easytier/src/connector/dns_connector.rs +++ /dev/null @@ -1,260 +0,0 @@ -use std::{net::SocketAddr, sync::Arc}; - -use super::{create_connector_by_url, http_connector::TunnelWithInfo}; -use crate::{ - common::{ - dns::{RESOLVER, resolve_txt_record}, - error::Error, - global_ctx::ArcGlobalCtx, - log, - }, - proto::common::TunnelInfo, - tunnel::{IpScheme, IpVersion, Tunnel, TunnelConnector, TunnelError, TunnelScheme}, -}; -use anyhow::Context; -use dashmap::DashSet; -use hickory_resolver::proto::rr::rdata::SRV; -use rand::{Rng as _, seq::SliceRandom}; -use strum::VariantArray; - -fn weighted_choice(options: &[(T, u64)]) -> Option<&T> { - let total_weight = options.iter().map(|(_, weight)| *weight).sum(); - let mut rng = rand::thread_rng(); - let rand_value = rng.gen_range(0..total_weight); - let mut accumulated_weight = 0; - - for (item, weight) in options { - accumulated_weight += *weight; - if rand_value < accumulated_weight { - return Some(item); - } - } - - None -} - -#[derive(Debug)] -pub struct DnsTunnelConnector { - scheme: TunnelScheme, - addr: url::Url, - bind_addrs: Vec, - global_ctx: ArcGlobalCtx, - ip_version: IpVersion, -} - -impl DnsTunnelConnector { - pub fn new(addr: url::Url, global_ctx: ArcGlobalCtx) -> Self { - Self { - scheme: (&addr).try_into().unwrap(), - addr, - bind_addrs: Vec::new(), - global_ctx, - ip_version: IpVersion::Both, - } - } - - #[tracing::instrument(ret, err)] - pub async fn handle_txt_record( - &self, - domain_name: &str, - ) -> Result, Error> { - let txt_data = resolve_txt_record(domain_name) - .await - .with_context(|| format!("resolve txt record failed, domain_name: {}", domain_name))?; - - let candidate_urls = txt_data - .split(" ") - .map(|s| s.to_string()) - .filter_map(|s| url::Url::parse(s.as_str()).ok()) - .collect::>(); - - // shuffle candidate_urls and get the first one - let url = candidate_urls - .choose(&mut rand::thread_rng()) - .with_context(|| { - format!( - "no valid url found, txt_data: {}, expecting an url list splitted by space", - txt_data - ) - })?; - - let connector = - create_connector_by_url(url.as_str(), &self.global_ctx, self.ip_version).await?; - Ok(connector) - } - - fn handle_one_srv_record(record: &SRV, protocol: IpScheme) -> Result<(url::Url, u64), Error> { - // port must be non-zero - if record.port() == 0 { - return Err(anyhow::anyhow!("port must be non-zero").into()); - } - - let connector_dst = record.target().to_utf8(); - let dst_url = format!("{}://{}:{}", protocol, connector_dst, record.port()); - - Ok(( - dst_url.parse().with_context(|| { - format!( - "parse dst_url failed, protocol: {}, connector_dst: {}, port: {}, dst_url: {}", - protocol, - connector_dst, - record.port(), - dst_url - ) - })?, - record.priority() as _, - )) - } - - #[tracing::instrument(ret, err)] - pub async fn handle_srv_record( - &self, - domain_name: &str, - ) -> Result, Error> { - tracing::info!("handle_srv_record: {}", domain_name); - - let srv_domains = IpScheme::VARIANTS - .iter() - .map(|s| (s, format!("_easytier._{}.{}", s, domain_name))) - .collect::>(); - tracing::info!("build srv_domains: {:?}", srv_domains); - let responses = Arc::new(DashSet::new()); - let srv_lookup_tasks = srv_domains - .iter() - .map(|(protocol, srv_domain)| { - let resolver = RESOLVER.clone(); - let responses = responses.clone(); - async move { - let response = resolver.srv_lookup(srv_domain).await.with_context(|| { - format!("srv_lookup failed, srv_domain: {}", srv_domain) - })?; - tracing::info!(?response, ?srv_domain, "srv_lookup response"); - for record in response.iter() { - let parsed_record = Self::handle_one_srv_record(record, **protocol); - tracing::info!(?parsed_record, ?srv_domain, "parsed_record"); - if let Err(e) = &parsed_record { - log::warn!("got invalid srv record {:?}", e); - continue; - } - responses.insert(parsed_record.unwrap()); - } - Ok::<_, Error>(()) - } - }) - .collect::>(); - let _ = futures::future::join_all(srv_lookup_tasks).await; - - let srv_records = responses.iter().map(|r| r.clone()).collect::>(); - if srv_records.is_empty() { - return Err(anyhow::anyhow!("no srv record found").into()); - } - - let url = weighted_choice(srv_records.as_slice()).with_context(|| { - format!( - "failed to choose a srv record, domain_name: {}, srv_records: {:?}", - domain_name, srv_records - ) - })?; - - let connector = - create_connector_by_url(url.as_str(), &self.global_ctx, self.ip_version).await?; - Ok(connector) - } -} - -#[async_trait::async_trait] -impl super::TunnelConnector for DnsTunnelConnector { - async fn connect(&mut self) -> Result, TunnelError> { - let mut conn = match self.scheme { - TunnelScheme::Txt => self - .handle_txt_record( - self.addr - .host_str() - .as_ref() - .ok_or(anyhow::anyhow!("host should not be empty in txt url"))?, - ) - .await - .with_context(|| "get txt record url failed")?, - TunnelScheme::Srv => self - .handle_srv_record( - self.addr - .host_str() - .as_ref() - .ok_or(anyhow::anyhow!("host should not be empty in srv url"))?, - ) - .await - .with_context(|| "get srv record url failed")?, - _ => return Err(anyhow::anyhow!("unsupported dns scheme: {:?}", self.scheme).into()), - }; - let t = conn.connect().await?; - let info = t.info().unwrap_or_default(); - Ok(Box::new(TunnelWithInfo::new( - t, - TunnelInfo { - local_addr: info.local_addr.clone(), - remote_addr: Some(self.addr.clone().into()), - resolved_remote_addr: info - .resolved_remote_addr - .clone() - .or(info.remote_addr.clone()), - tunnel_type: format!("{}-{}", self.addr.scheme(), info.tunnel_type), - }, - ))) - } - - fn remote_url(&self) -> url::Url { - self.addr.clone() - } - - fn set_bind_addrs(&mut self, addrs: Vec) { - self.bind_addrs = addrs; - } - - fn set_ip_version(&mut self, ip_version: IpVersion) { - self.ip_version = ip_version; - } -} - -#[cfg(test)] -mod tests { - use super::*; - use crate::common::global_ctx::tests::get_mock_global_ctx; - - #[tokio::test] - async fn test_txt() { - let url = "txt://txt.easytier.cn"; - let global_ctx = get_mock_global_ctx(); - let mut connector = DnsTunnelConnector::new(url.parse().unwrap(), global_ctx); - connector.set_ip_version(IpVersion::V4); - for _ in 0..5 { - match connector.connect().await { - Ok(ret) => { - println!("{:?}", ret.info()); - return; - } - Err(e) => { - println!("{:?}", e); - } - } - } - } - - #[tokio::test] - async fn test_srv() { - let url = "srv://easytier.cn"; - let global_ctx = get_mock_global_ctx(); - let mut connector = DnsTunnelConnector::new(url.parse().unwrap(), global_ctx); - connector.set_ip_version(IpVersion::V4); - for _ in 0..5 { - match connector.connect().await { - Ok(ret) => { - println!("{:?}", ret.info()); - return; - } - Err(e) => { - println!("{:?}", e); - } - } - } - } -} diff --git a/easytier/src/connector/http_connector.rs b/easytier/src/connector/http_connector.rs deleted file mode 100644 index 154b02a1..00000000 --- a/easytier/src/connector/http_connector.rs +++ /dev/null @@ -1,361 +0,0 @@ -use std::{ - net::SocketAddr, - pin::Pin, - sync::{Arc, RwLock}, -}; - -use anyhow::Context; -use http_req::request::{RedirectPolicy, Request}; -use rand::seq::SliceRandom as _; -use url::Url; - -use crate::{ - VERSION, - common::{error::Error, global_ctx::ArcGlobalCtx}, - tunnel::{IpVersion, Tunnel, TunnelConnector, TunnelError, ZCPacketSink, ZCPacketStream}, -}; - -use crate::proto::common::TunnelInfo; - -use super::create_connector_by_url; - -pub struct TunnelWithInfo { - inner: Box, - info: TunnelInfo, -} - -impl TunnelWithInfo { - pub fn new(inner: Box, info: TunnelInfo) -> Self { - Self { inner, info } - } -} - -impl Tunnel for TunnelWithInfo { - fn split(&self) -> (Pin>, Pin>) { - self.inner.split() - } - - fn info(&self) -> Option { - Some(self.info.clone()) - } -} - -#[derive(Debug, PartialEq, Copy, Clone)] -enum HttpRedirectType { - Unknown, - // redirected url is in the path of new url - RedirectToQuery, - // redirected url is the entire new url - RedirectToUrl, - // redirected url is in the body of response - BodyUrls, -} - -#[derive(Debug)] -pub struct HttpTunnelConnector { - addr: url::Url, - bind_addrs: Vec, - ip_version: IpVersion, - global_ctx: ArcGlobalCtx, - redirect_type: HttpRedirectType, -} - -impl HttpTunnelConnector { - pub fn new(addr: url::Url, global_ctx: ArcGlobalCtx) -> Self { - Self { - addr, - bind_addrs: Vec::new(), - ip_version: IpVersion::Both, - global_ctx, - redirect_type: HttpRedirectType::Unknown, - } - } - - #[tracing::instrument(ret)] - async fn handle_302_redirect( - &mut self, - new_url: url::Url, - url_str: &str, - ) -> Result, Error> { - // the url should be in following format: - // 1: http(s)://easytier.cn/?url=tcp://10.147.22.22:11010 (scheme is http, domain is ignored, path is splitted into proto type and addr) - // 2: http(s)://tcp://10.137.22.22:11010 (connector url is appended to the scheme) - // 3: tcp://10.137.22.22:11010 (scheme is protocol type, the url is used to construct a connector directly) - tracing::info!("redirect to {}", new_url); - let url = url::Url::parse(new_url.as_str()) - .with_context(|| format!("parsing redirect url failed. url: {}", new_url))?; - if url.scheme() == "http" || url.scheme() == "https" { - let mut query = new_url - .query_pairs() - .filter_map(|x| url::Url::parse(&x.1).ok()) - .collect::>(); - query.shuffle(&mut rand::thread_rng()); - if !query.is_empty() { - tracing::info!("try to create connector by url: {}", query[0]); - self.redirect_type = HttpRedirectType::RedirectToQuery; - return create_connector_by_url( - query[0].as_ref(), - &self.global_ctx, - self.ip_version, - ) - .await; - } else if let Some(new_url) = url_str - .strip_prefix(format!("{}://", url.scheme()).as_str()) - .and_then(|x| Url::parse(x).ok()) - { - // stripe the scheme and create connector by url - self.redirect_type = HttpRedirectType::RedirectToUrl; - return create_connector_by_url( - new_url.as_str(), - &self.global_ctx, - self.ip_version, - ) - .await; - } - return Err(Error::InvalidUrl(format!( - "no valid connector url found in url: {}", - url - ))); - } else { - self.redirect_type = HttpRedirectType::RedirectToUrl; - return create_connector_by_url(new_url.as_str(), &self.global_ctx, self.ip_version) - .await; - } - } - - #[tracing::instrument] - async fn handle_200_success( - &mut self, - body: &String, - ) -> Result, Error> { - // resp body should be line of connector urls, like: - // tcp://10.1.1.1:11010 - // udp://10.1.1.1:11010 - let mut lines = body - .lines() - .map(|line| line.trim()) - .filter(|line| !line.is_empty()) - .collect::>(); - - tracing::info!("get {} lines of connector urls", lines.len()); - - // shuffle the lines and pick the usable one - lines.shuffle(&mut rand::thread_rng()); - - for line in lines { - let url = url::Url::parse(line); - if url.is_err() { - tracing::warn!("invalid url: {}, skip it", line); - continue; - } - self.redirect_type = HttpRedirectType::BodyUrls; - return create_connector_by_url(line, &self.global_ctx, self.ip_version).await; - } - - Err(Error::InvalidUrl(format!( - "no valid connector url found, response body: {}", - body - ))) - } - - #[tracing::instrument(ret)] - pub async fn get_redirected_connector( - &mut self, - original_url: &str, - ) -> Result, Error> { - self.redirect_type = HttpRedirectType::Unknown; - tracing::info!("get_redirected_url: {}", original_url); - // Container for body of a response. - let body = Arc::new(RwLock::new(Vec::new())); - - let original_url_clone = original_url.to_string(); - let body_clone = body.clone(); - let network_name = self.global_ctx.network.network_name.clone(); - let user_agent = format!("easytier/{}", VERSION); - let res = tokio::task::spawn_blocking(move || { - let uri = http_req::uri::Uri::try_from(original_url_clone.as_ref()) - .with_context(|| format!("parsing url failed. url: {}", original_url_clone))?; - - tracing::info!( - "sending http request to {}, network_name: {}", - uri, - network_name - ); - - Request::new(&uri) - .header("User-Agent", &user_agent) - .header("X-Network-Name", &network_name) - .redirect_policy(RedirectPolicy::Limit(0)) - .timeout(std::time::Duration::from_secs(20)) - .send(&mut *body_clone.write().unwrap()) - .with_context(|| format!("sending http request failed. url: {}", uri)) - }) - .await - .map_err(|e| Error::InvalidUrl(format!("task join error: {}", e)))??; - - let body = String::from_utf8_lossy(&body.read().unwrap()).to_string(); - - if res.status_code().is_redirect() { - let redirect_url = res - .headers() - .get("Location") - .ok_or_else(|| Error::InvalidUrl("no redirect address found".to_string()))?; - let new_url = url::Url::parse(redirect_url.as_str()) - .with_context(|| format!("parsing redirect url failed. url: {}", redirect_url))?; - return self.handle_302_redirect(new_url, redirect_url).await; - } else if res.status_code().is_success() { - return self.handle_200_success(&body).await; - } else { - return Err(Error::InvalidUrl(format!( - "unexpected response, resp: {:?}, body: {}", - res, body, - ))); - } - } -} - -#[async_trait::async_trait] -impl super::TunnelConnector for HttpTunnelConnector { - async fn connect(&mut self) -> Result, TunnelError> { - let mut conn = self - .get_redirected_connector(self.addr.to_string().as_str()) - .await - .with_context(|| "get redirected url failed")?; - conn.set_ip_version(self.ip_version); - let t = conn.connect().await?; - let info = t.info().unwrap_or_default(); - Ok(Box::new(TunnelWithInfo::new( - t, - TunnelInfo { - local_addr: info.local_addr.clone(), - remote_addr: Some(self.addr.clone().into()), - resolved_remote_addr: info - .resolved_remote_addr - .clone() - .or(info.remote_addr.clone()), - tunnel_type: format!("{}-{}", self.addr.scheme(), info.tunnel_type), - }, - ))) - } - - fn remote_url(&self) -> url::Url { - self.addr.clone() - } - - fn set_bind_addrs(&mut self, addrs: Vec) { - self.bind_addrs = addrs; - } - - fn set_ip_version(&mut self, ip_version: IpVersion) { - self.ip_version = ip_version; - } -} - -#[cfg(test)] -mod tests { - use tokio::{io::AsyncReadExt as _, io::AsyncWriteExt as _, net::TcpListener}; - - use crate::{ - common::global_ctx::tests::get_mock_global_ctx_with_network, - tunnel::{TunnelConnector, TunnelListener, tcp::TcpTunnelListener}, - }; - - use super::*; - - async fn run_http_redirect_server( - port: u16, - test_type: HttpRedirectType, - ) -> Result { - let listener = TcpListener::bind(format!("127.0.0.1:{}", port)).await?; - let (mut stream, _) = listener.accept().await?; - - let mut buf = [0u8; 4096]; - let n = stream.read(&mut buf).await?; - let req = String::from_utf8_lossy(&buf[..n]); - - let mut captured_network_name = String::new(); - for line in req.lines() { - if line.to_lowercase().starts_with("x-network-name:") { - captured_network_name = line - .split_once(':') - .map(|x| x.1) - .unwrap_or_default() - .trim() - .to_string(); - break; - } - } - - match test_type { - HttpRedirectType::RedirectToQuery => { - let resp = "HTTP/1.1 301 Moved Permanently\r\nLocation: http://test.com/?url=tcp://127.0.0.1:25888\r\n\r\n"; - stream.write_all(resp.as_bytes()).await?; - } - HttpRedirectType::RedirectToUrl => { - let resp = - "HTTP/1.1 301 Moved Permanently\r\nLocation: tcp://127.0.0.1:25888\r\n\r\n"; - stream.write_all(resp.as_bytes()).await?; - } - HttpRedirectType::BodyUrls => { - let resp = "HTTP/1.1 200 OK\r\n\r\ntcp://127.0.0.1:25888"; - stream.write_all(resp.as_bytes()).await?; - } - HttpRedirectType::Unknown => { - panic!("unexpected test type"); - } - } - - Ok(captured_network_name) - } - - #[rstest::rstest] - #[serial_test::serial(http_redirect_test)] - #[tokio::test] - async fn http_redirect_test( - // 1. 301 redirect - // 2. 200 success with valid connector urls - #[values( - HttpRedirectType::RedirectToQuery, - HttpRedirectType::RedirectToUrl, - HttpRedirectType::BodyUrls - )] - test_type: HttpRedirectType, - ) { - let network_name = format!("net_{}", rand::random::()); - let http_task = tokio::spawn(run_http_redirect_server(35888, test_type)); - tokio::time::sleep(std::time::Duration::from_millis(10)).await; - let test_url: url::Url = "http://127.0.0.1:35888".parse().unwrap(); - - let identity = crate::common::config::NetworkIdentity { - network_name: network_name.clone(), - ..Default::default() - }; - let global_ctx = get_mock_global_ctx_with_network(Some(identity)); - - let mut flags = global_ctx.config.get_flags(); - flags.bind_device = false; - global_ctx.set_flags(flags); - let mut connector = HttpTunnelConnector::new(test_url.clone(), global_ctx.clone()); - - let mut listener = TcpTunnelListener::new("tcp://0.0.0.0:25888".parse().unwrap()); - listener.listen().await.unwrap(); - - let task = tokio::spawn(async move { - let _conn = listener.accept().await.unwrap(); - }); - - let t = connector.connect().await.unwrap(); - assert_eq!(connector.redirect_type, test_type); - - let captured_name = http_task.await.unwrap().unwrap(); - assert_eq!(captured_name, network_name); - - let info = t.info().unwrap(); - let remote_addr = info.remote_addr.unwrap(); - assert_eq!(remote_addr, test_url.into()); - let resolved_remote_addr = info.resolved_remote_addr.unwrap(); - assert_eq!(resolved_remote_addr.url, "tcp://127.0.0.1:25888"); - - tokio::join!(task).0.unwrap(); - } -} diff --git a/easytier/src/connector/manual.rs b/easytier/src/connector/manual.rs deleted file mode 100644 index c4cce1b9..00000000 --- a/easytier/src/connector/manual.rs +++ /dev/null @@ -1,525 +0,0 @@ -use std::{ - collections::BTreeSet, - future::Future, - sync::{Arc, Weak}, - time::Duration, -}; - -use dashmap::DashSet; -use quanta::Instant; -use tokio::{sync::mpsc, task::JoinSet, time::timeout}; - -use crate::{ - common::{PeerId, dns::socket_addrs, join_joinset_background}, - peers::peer_conn::PeerConnId, - proto::{ - api::instance::{ - Connector, ConnectorManageRpc, ConnectorStatus, ListConnectorRequest, - ListConnectorResponse, - }, - rpc_types::{self, controller::BaseController}, - }, - tunnel::{IpVersion, TunnelConnector, TunnelScheme, matches_scheme}, - utils::weak_upgrade, -}; - -use crate::{ - common::{ - error::Error, - global_ctx::{ArcGlobalCtx, GlobalCtxEvent}, - netns::NetNS, - }, - peers::peer_manager::PeerManager, - use_global_var, -}; - -use super::create_connector_by_url; - -type ConnectorMap = Arc>; - -#[derive(Debug, Clone)] -struct ReconnResult { - dead_url: String, - peer_id: PeerId, - conn_id: PeerConnId, -} - -struct ConnectorManagerData { - connectors: ConnectorMap, - reconnecting: DashSet, - peer_manager: Weak, - alive_conn_urls: Arc>, - // user removed connector urls - removed_conn_urls: Arc>, - net_ns: NetNS, - global_ctx: ArcGlobalCtx, -} - -pub struct ManualConnectorManager { - global_ctx: ArcGlobalCtx, - data: Arc, - tasks: JoinSet<()>, -} - -impl ManualConnectorManager { - pub fn new(global_ctx: ArcGlobalCtx, peer_manager: Arc) -> Self { - let connectors = Arc::new(DashSet::new()); - let tasks = JoinSet::new(); - - let mut ret = Self { - global_ctx: global_ctx.clone(), - data: Arc::new(ConnectorManagerData { - connectors, - reconnecting: DashSet::new(), - peer_manager: Arc::downgrade(&peer_manager), - alive_conn_urls: Arc::new(DashSet::new()), - removed_conn_urls: Arc::new(DashSet::new()), - net_ns: global_ctx.net_ns.clone(), - global_ctx, - }), - tasks, - }; - - ret.tasks - .spawn(Self::conn_mgr_reconn_routine(ret.data.clone())); - - ret - } - - fn reconnect_timeout(dead_url: &url::Url) -> Duration { - let use_long_timeout = matches_scheme!( - dead_url, - TunnelScheme::Http | TunnelScheme::Https | TunnelScheme::Txt | TunnelScheme::Srv - ) || matches!(dead_url.scheme(), "ws" | "wss"); - - Duration::from_secs(if use_long_timeout { 20 } else { 2 }) - } - - fn remaining_budget(started_at: Instant, total_timeout: Duration) -> Option { - let remaining = total_timeout.checked_sub(started_at.elapsed())?; - (!remaining.is_zero()).then_some(remaining) - } - - fn emit_connect_error( - data: &ConnectorManagerData, - dead_url: &url::Url, - ip_version: IpVersion, - error: &Error, - ) { - data.global_ctx.issue_event(GlobalCtxEvent::ConnectError( - dead_url.to_string(), - format!("{:?}", ip_version), - format!("{:#?}", error), - )); - } - - fn reconnect_timeout_error(stage: &str, duration: Duration) -> Error { - Error::AnyhowError(anyhow::anyhow!("{} timeout after {:?}", stage, duration)) - } - - async fn with_reconnect_timeout( - stage: &'static str, - started_at: Instant, - total_timeout: Duration, - fut: F, - ) -> Result - where - F: Future>, - { - let remaining = Self::remaining_budget(started_at, total_timeout) - .ok_or_else(|| Self::reconnect_timeout_error(stage, started_at.elapsed()))?; - timeout(remaining, fut) - .await - .map_err(|_| Self::reconnect_timeout_error(stage, remaining))? - } -} - -impl ManualConnectorManager { - pub fn add_connector(&self, connector: T) - where - T: TunnelConnector + 'static, - { - tracing::info!("add_connector: {}", connector.remote_url()); - self.data.connectors.insert(connector.remote_url()); - } - - pub async fn add_connector_by_url(&self, url: url::Url) -> Result<(), Error> { - self.data.connectors.insert(url); - Ok(()) - } - - pub async fn remove_connector(&self, url: url::Url) -> Result<(), Error> { - tracing::info!("remove_connector: {}", url); - let url = url.into(); - if !self - .list_connectors() - .await - .iter() - .any(|x| x.url.as_ref() == Some(&url)) - { - return Err(Error::NotFound); - } - self.data.removed_conn_urls.insert(url.into()); - Ok(()) - } - - pub async fn clear_connectors(&self) { - self.list_connectors().await.iter().for_each(|x| { - if let Some(url) = &x.url { - self.data.removed_conn_urls.insert(url.clone().into()); - } - }); - } - - pub async fn list_connectors(&self) -> Vec { - let dead_urls: BTreeSet = Self::collect_dead_conns(self.data.clone()) - .await - .into_iter() - .collect(); - - let mut ret = Vec::new(); - - for item in self.data.connectors.iter() { - let conn_url = item.key().clone(); - let mut status = ConnectorStatus::Connected; - if dead_urls.contains(&conn_url) { - status = ConnectorStatus::Disconnected; - } - ret.insert( - 0, - Connector { - url: Some(conn_url.into()), - status: status.into(), - }, - ); - } - - let reconnecting_urls: BTreeSet = - self.data.reconnecting.iter().map(|x| x.clone()).collect(); - - for conn_url in reconnecting_urls { - ret.insert( - 0, - Connector { - url: Some(conn_url.into()), - status: ConnectorStatus::Connecting.into(), - }, - ); - } - - ret - } - - async fn conn_mgr_reconn_routine(data: Arc) { - tracing::warn!("conn_mgr_routine started"); - let mut reconn_interval = tokio::time::interval(std::time::Duration::from_millis( - use_global_var!(MANUAL_CONNECTOR_RECONNECT_INTERVAL_MS), - )); - let (reconn_result_send, mut reconn_result_recv) = mpsc::channel(100); - let tasks = Arc::new(std::sync::Mutex::new(JoinSet::new())); - join_joinset_background(tasks.clone(), "connector_reconnect_tasks".to_string()); - - loop { - tokio::select! { - _ = reconn_interval.tick() => { - let dead_urls = Self::collect_dead_conns(data.clone()).await; - if dead_urls.is_empty() { - continue; - } - for dead_url in dead_urls { - let data_clone = data.clone(); - let sender = reconn_result_send.clone(); - data.connectors.remove(&dead_url).unwrap(); - let insert_succ = data.reconnecting.insert(dead_url.clone()); - assert!(insert_succ); - - tasks.lock().unwrap().spawn(async move { - let reconn_ret = Self::conn_reconnect(data_clone.clone(), dead_url.clone() ).await; - let _ = sender.send(reconn_ret).await; - - data_clone.reconnecting.remove(&dead_url).unwrap(); - data_clone.connectors.insert(dead_url.clone()); - }); - } - tracing::info!("reconn_interval tick, done"); - } - - ret = reconn_result_recv.recv() => { - tracing::warn!("reconn_tasks done, reconn result: {:?}", ret); - } - } - } - } - - fn handle_remove_connector(data: Arc) { - let remove_later = DashSet::new(); - for it in data.removed_conn_urls.iter() { - let url = it.key(); - if data.connectors.remove(url).is_some() { - tracing::warn!("connector: {}, removed", url); - continue; - } else if data.reconnecting.contains(url) { - tracing::warn!("connector: {}, reconnecting, remove later.", url); - remove_later.insert(url.clone()); - continue; - } else { - tracing::warn!("connector: {}, not found", url); - } - } - data.removed_conn_urls.clear(); - for it in remove_later.iter() { - data.removed_conn_urls.insert(it.key().clone()); - } - } - - async fn collect_dead_conns(data: Arc) -> BTreeSet { - Self::handle_remove_connector(data.clone()); - let mut ret = BTreeSet::new(); - let Some(pm) = data.peer_manager.upgrade() else { - tracing::warn!("peer manager is gone, exit"); - return ret; - }; - for url in data.connectors.iter().map(|x| x.key().clone()) { - if !pm.get_peer_map().is_client_url_alive(&url) - && !pm - .get_foreign_network_client() - .get_peer_map() - .is_client_url_alive(&url) - { - ret.insert(url.clone()); - } - } - ret - } - - async fn conn_reconnect_with_ip_version( - data: Arc, - dead_url: url::Url, - ip_version: IpVersion, - started_at: Instant, - total_timeout: Duration, - ) -> Result { - let connector = Self::with_reconnect_timeout( - "resolve", - started_at, - total_timeout, - create_connector_by_url(dead_url.as_str(), &data.global_ctx, ip_version), - ) - .await?; - - data.global_ctx - .issue_event(GlobalCtxEvent::Connecting(connector.remote_url())); - tracing::info!("reconnect try connect... conn: {:?}", connector); - let Some(pm) = data.peer_manager.upgrade() else { - return Err(Error::AnyhowError(anyhow::anyhow!( - "peer manager is gone, cannot reconnect" - ))); - }; - - let tunnel = Self::with_reconnect_timeout( - "connect", - started_at, - total_timeout, - pm.connect_tunnel(connector), - ) - .await?; - - let (peer_id, conn_id) = Self::with_reconnect_timeout( - "handshake", - started_at, - total_timeout, - pm.add_client_tunnel_with_peer_id_hint(tunnel, true, None), - ) - .await?; - - tracing::info!("reconnect succ: {} {} {}", peer_id, conn_id, dead_url); - Ok(ReconnResult { - dead_url: dead_url.to_string(), - peer_id, - conn_id, - }) - } - - async fn conn_reconnect( - data: Arc, - dead_url: url::Url, - ) -> Result { - tracing::info!("reconnect: {}", dead_url); - - let mut ip_versions = vec![]; - if matches_scheme!( - dead_url, - TunnelScheme::Ring | TunnelScheme::Txt | TunnelScheme::Srv - ) { - ip_versions.push(IpVersion::Both); - } else { - let converted_dead_url = - match crate::common::idn::convert_idn_to_ascii(dead_url.clone()) { - Ok(url) => url, - Err(error) => { - let error: Error = error.into(); - Self::emit_connect_error(&data, &dead_url, IpVersion::Both, &error); - return Err(error); - } - }; - let addrs = match Self::with_reconnect_timeout( - "resolve", - Instant::now(), - Self::reconnect_timeout(&dead_url), - socket_addrs(&converted_dead_url, || Some(1000)), - ) - .await - { - Ok(addrs) => addrs, - Err(error) => { - Self::emit_connect_error(&data, &dead_url, IpVersion::Both, &error); - return Err(error); - } - }; - tracing::info!(?addrs, ?dead_url, "get ip from url done"); - let mut has_ipv4 = false; - let mut has_ipv6 = false; - for addr in addrs { - if addr.is_ipv4() { - if !has_ipv4 { - ip_versions.insert(0, IpVersion::V4); - } - has_ipv4 = true; - } else if addr.is_ipv6() { - if !has_ipv6 { - ip_versions.push(IpVersion::V6); - } - has_ipv6 = true; - } - } - } - - let mut reconn_ret = Err(Error::AnyhowError(anyhow::anyhow!( - "cannot get ip from url" - ))); - for ip_version in ip_versions { - let started_at = Instant::now(); - let ret = Self::conn_reconnect_with_ip_version( - data.clone(), - dead_url.clone(), - ip_version, - started_at, - Self::reconnect_timeout(&dead_url), - ) - .await; - tracing::info!("reconnect: {} done, ret: {:?}", dead_url, ret); - - match ret { - Ok(result) => return Ok(result), - Err(error) => { - Self::emit_connect_error(&data, &dead_url, ip_version, &error); - reconn_ret = Err(error); - } - } - } - - reconn_ret - } -} - -#[derive(Clone)] -pub struct ConnectorManagerRpcService(pub Weak); - -#[async_trait::async_trait] -impl ConnectorManageRpc for ConnectorManagerRpcService { - type Controller = BaseController; - - async fn list_connector( - &self, - _: BaseController, - _request: ListConnectorRequest, - ) -> Result { - let mut ret = ListConnectorResponse::default(); - let connectors = weak_upgrade(&self.0)?.list_connectors().await; - ret.connectors = connectors; - Ok(ret) - } -} - -#[cfg(test)] -mod tests { - use crate::{ - peers::tests::create_mock_peer_manager, - set_global_var, - tunnel::{Tunnel, TunnelError}, - }; - - use super::*; - - #[tokio::test] - async fn reconnect_timeout_reports_exhausted_budget_for_stage() { - let started_at = Instant::now() - Duration::from_millis(50); - let err = ManualConnectorManager::with_reconnect_timeout( - "resolve", - started_at, - Duration::from_millis(1), - async { Ok::<(), Error>(()) }, - ) - .await - .unwrap_err(); - - let message = err.to_string(); - assert!(message.contains("resolve timeout after")); - } - - #[tokio::test] - async fn reconnect_timeout_reports_stage_timeout_with_remaining_budget() { - let err = ManualConnectorManager::with_reconnect_timeout( - "handshake", - Instant::now(), - Duration::from_millis(10), - async { - tokio::time::sleep(Duration::from_millis(50)).await; - Ok::<(), Error>(()) - }, - ) - .await - .unwrap_err(); - - let message = err.to_string(); - assert!(message.contains("handshake timeout after")); - } - - #[tokio::test] - async fn reconnect_timeout_preserves_success_within_budget() { - let result = ManualConnectorManager::with_reconnect_timeout( - "connect", - Instant::now(), - Duration::from_millis(50), - async { Ok::<_, Error>(123_u32) }, - ) - .await - .unwrap(); - - assert_eq!(result, 123); - } - - #[tokio::test] - async fn test_reconnect_with_connecting_addr() { - set_global_var!(MANUAL_CONNECTOR_RECONNECT_INTERVAL_MS, 1); - - let peer_mgr = create_mock_peer_manager().await; - let mgr = ManualConnectorManager::new(peer_mgr.get_global_ctx(), peer_mgr); - - struct MockConnector {} - #[async_trait::async_trait] - impl TunnelConnector for MockConnector { - fn remote_url(&self) -> url::Url { - url::Url::parse("tcp://aa.com").unwrap() - } - async fn connect(&mut self) -> Result, TunnelError> { - tokio::time::sleep(std::time::Duration::from_millis(10)).await; - Err(TunnelError::InvalidPacket("fake error".into())) - } - } - - mgr.add_connector(MockConnector {}); - - tokio::time::sleep(std::time::Duration::from_secs(5)).await; - } -} diff --git a/easytier/src/connector/mod.rs b/easytier/src/connector/mod.rs deleted file mode 100644 index a5e72b89..00000000 --- a/easytier/src/connector/mod.rs +++ /dev/null @@ -1,462 +0,0 @@ -use std::net::{IpAddr, Ipv6Addr, SocketAddr, SocketAddrV4, SocketAddrV6}; - -use crate::{ - common::{dns::socket_addrs, error::Error, global_ctx::ArcGlobalCtx, idn}, - connector::dns_connector::DnsTunnelConnector, - proto::common::PeerFeatureFlag, - tunnel::{ - self, IpScheme, IpVersion, TunnelConnector, TunnelError, TunnelScheme, - ring::RingTunnelConnector, tcp::TcpTunnelConnector, udp::UdpTunnelConnector, - }, - utils::BoxExt, -}; -use http_connector::HttpTunnelConnector; -use rand::seq::SliceRandom; - -pub mod direct; -pub mod manual; -pub mod tcp_hole_punch; -pub mod udp_hole_punch; - -pub mod dns_connector; -pub mod http_connector; - -pub(crate) fn should_try_p2p_with_peer( - feature_flag: Option<&PeerFeatureFlag>, - allow_public_server: bool, - local_disable_p2p: bool, - local_need_p2p: bool, -) -> bool { - feature_flag - .map(|flag| { - (allow_public_server || !flag.is_public_server) - && (!local_disable_p2p || flag.need_p2p) - && (!flag.disable_p2p || local_need_p2p) - }) - .unwrap_or(!local_disable_p2p) -} - -pub(crate) fn should_background_p2p_with_peer( - feature_flag: Option<&PeerFeatureFlag>, - allow_public_server: bool, - lazy_p2p: bool, - local_disable_p2p: bool, - local_need_p2p: bool, -) -> bool { - should_try_p2p_with_peer( - feature_flag, - allow_public_server, - local_disable_p2p, - local_need_p2p, - ) && (!lazy_p2p || feature_flag.map(|flag| flag.need_p2p).unwrap_or(false)) -} - -async fn set_bind_addr_for_peer_connector( - connector: &mut (impl TunnelConnector + ?Sized), - is_ipv4: bool, - global_ctx: &ArcGlobalCtx, -) { - if cfg!(any( - target_os = "android", - any( - target_os = "ios", - all(target_os = "macos", feature = "macos-ne") - ), - target_env = "ohos" - )) { - return; - } - - let ips = global_ctx.get_ip_collector().collect_ip_addrs().await; - if is_ipv4 { - let mut bind_addrs = vec![]; - for ipv4 in ips.interface_ipv4s { - let socket_addr = SocketAddrV4::new(ipv4.into(), 0).into(); - bind_addrs.push(socket_addr); - } - connector.set_bind_addrs(bind_addrs); - } else { - let mut bind_addrs = vec![]; - for ipv6 in ips.interface_ipv6s.iter().chain(ips.public_ipv6.iter()) { - let ipv6 = std::net::Ipv6Addr::from(*ipv6); - if global_ctx.is_ip_easytier_managed_ipv6(&ipv6) { - continue; - } - let socket_addr = SocketAddrV6::new(ipv6, 0, 0, 0).into(); - bind_addrs.push(socket_addr); - } - connector.set_bind_addrs(bind_addrs); - } - let _ = connector; -} - -struct ResolvedConnectorAddr { - addr: SocketAddr, - ip_version: IpVersion, -} - -fn connector_default_port(url: &url::Url) -> Option { - url.try_into() - .ok() - .and_then(|s: TunnelScheme| s.try_into().ok()) - .map(IpScheme::default_port) -} - -fn addr_matches_ip_version(addr: &SocketAddr, ip_version: IpVersion) -> bool { - match ip_version { - IpVersion::V4 => addr.is_ipv4(), - IpVersion::V6 => addr.is_ipv6(), - IpVersion::Both => true, - } -} - -fn infer_effective_ip_version(addrs: &[SocketAddr], requested_ip_version: IpVersion) -> IpVersion { - match requested_ip_version { - IpVersion::Both if addrs.iter().all(SocketAddr::is_ipv4) => IpVersion::V4, - IpVersion::Both if addrs.iter().all(SocketAddr::is_ipv6) => IpVersion::V6, - _ => requested_ip_version, - } -} - -async fn easytier_managed_ipv6_source_for_dst( - global_ctx: &ArcGlobalCtx, - dst_addr: SocketAddrV6, -) -> Result, Error> { - let socket = { - let _g = global_ctx.net_ns.guard(); - tokio::net::UdpSocket::bind("[::]:0").await? - }; - socket.connect(SocketAddr::V6(dst_addr)).await?; - - let IpAddr::V6(local_ip) = socket.local_addr()?.ip() else { - return Ok(None); - }; - - Ok(global_ctx - .is_ip_easytier_managed_ipv6(&local_ip) - .then_some(local_ip)) -} - -async fn ipv6_connector_reject_reason( - url: &url::Url, - global_ctx: &ArcGlobalCtx, - v6_addr: SocketAddrV6, - skip_source_validation_errors: bool, -) -> Result, Error> { - if global_ctx.is_ip_easytier_managed_ipv6(v6_addr.ip()) { - return Ok(Some(format!( - "{} resolves to EasyTier-managed IPv6 {}", - url, - v6_addr.ip() - ))); - } - - match easytier_managed_ipv6_source_for_dst(global_ctx, v6_addr).await { - Ok(Some(local_ip)) => Ok(Some(format!( - "{} would use EasyTier-managed IPv6 {} as local source for {}", - url, local_ip, v6_addr - ))), - Ok(None) => Ok(None), - Err(err) if skip_source_validation_errors => Ok(Some(format!( - "{} IPv6 candidate {} could not be validated: {}", - url, v6_addr, err - ))), - Err(err) => Err(err), - } -} - -async fn resolve_connector_socket_addr( - url: &url::Url, - global_ctx: &ArcGlobalCtx, - ip_version: IpVersion, -) -> Result { - let addrs = socket_addrs(url, || connector_default_port(url)) - .await - .map_err(|e| { - TunnelError::InvalidAddr(format!( - "failed to resolve socket addr, url: {}, error: {}", - url, e - )) - })?; - - let mut usable_addrs = Vec::new(); - let mut rejected_ipv6_reason = None; - let skip_source_validation_errors = ip_version == IpVersion::Both; - for addr in addrs - .into_iter() - .filter(|addr| addr_matches_ip_version(addr, ip_version)) - { - if let SocketAddr::V6(v6_addr) = addr - && let Some(reason) = ipv6_connector_reject_reason( - url, - global_ctx, - v6_addr, - skip_source_validation_errors, - ) - .await? - { - rejected_ipv6_reason = Some(reason); - continue; - } - - usable_addrs.push(addr); - } - - if usable_addrs.is_empty() { - if let Some(reason) = rejected_ipv6_reason { - return Err(Error::InvalidUrl(format!( - "{}, refusing overlay-backed underlay connection", - reason - ))); - } - - return Err(Error::TunnelError(TunnelError::NoDnsRecordFound( - ip_version, - ))); - } - - let effective_ip_version = infer_effective_ip_version(&usable_addrs, ip_version); - - let addr = usable_addrs - .choose(&mut rand::thread_rng()) - .copied() - .ok_or_else(|| Error::TunnelError(TunnelError::NoDnsRecordFound(ip_version)))?; - - Ok(ResolvedConnectorAddr { - addr, - ip_version: effective_ip_version, - }) -} - -pub async fn create_connector_by_url( - url: &str, - global_ctx: &ArcGlobalCtx, - ip_version: IpVersion, -) -> Result, Error> { - let url = url::Url::parse(url).map_err(|_| Error::InvalidUrl(url.to_owned()))?; - let url = idn::convert_idn_to_ascii(url)?; - let scheme = (&url) - .try_into() - .map_err(|_| TunnelError::InvalidProtocol(url.scheme().to_owned()))?; - let mut effective_connector_ip_version = ip_version; - let mut connector: Box = match scheme { - TunnelScheme::Ip(scheme) => { - let resolved_addr = resolve_connector_socket_addr(&url, global_ctx, ip_version).await?; - effective_connector_ip_version = resolved_addr.ip_version; - let mut connector: Box = match scheme { - IpScheme::Tcp => TcpTunnelConnector::new(url).boxed(), - IpScheme::Udp => UdpTunnelConnector::new(url).boxed(), - #[cfg(feature = "quic")] - IpScheme::Quic => { - tunnel::quic::QuicTunnelConnector::new(url, global_ctx.clone()).boxed() - } - #[cfg(feature = "wireguard")] - IpScheme::Wg => { - use crate::tunnel::wireguard::{WgConfig, WgTunnelConnector}; - let nid = global_ctx.get_network_identity(); - let wg_config = WgConfig::new_from_network_identity( - &nid.network_name, - &nid.network_secret.unwrap_or_default(), - ); - WgTunnelConnector::new(url, wg_config).boxed() - } - #[cfg(feature = "websocket")] - IpScheme::Ws | IpScheme::Wss => { - tunnel::websocket::WsTunnelConnector::new(url).boxed() - } - #[cfg(feature = "faketcp")] - IpScheme::FakeTcp => tunnel::fake_tcp::FakeTcpTunnelConnector::new(url).boxed(), - }; - connector.set_resolved_addr(resolved_addr.addr); - connector.set_socket_mark(global_ctx.config.get_flags().socket_mark); - if global_ctx.config.get_flags().bind_device { - set_bind_addr_for_peer_connector( - &mut connector, - resolved_addr.addr.is_ipv4(), - global_ctx, - ) - .await; - } - connector - } - #[cfg(unix)] - TunnelScheme::Unix => tunnel::unix::UnixSocketTunnelConnector::new(url).boxed(), - TunnelScheme::Http | TunnelScheme::Https => { - HttpTunnelConnector::new(url, global_ctx.clone()).boxed() - } - TunnelScheme::Ring => RingTunnelConnector::new(url).boxed(), - TunnelScheme::Txt | TunnelScheme::Srv => { - if url.host_str().is_none() { - return Err(Error::InvalidUrl(format!( - "host should not be empty in txt or srv url: {}", - url - ))); - } - DnsTunnelConnector::new(url, global_ctx.clone()).boxed() - } - }; - connector.set_ip_version(effective_connector_ip_version); - - Ok(connector) -} - -#[cfg(test)] -mod tests { - use std::collections::BTreeSet; - - use crate::{ - common::global_ctx::tests::get_mock_global_ctx, proto::common::PeerFeatureFlag, - tunnel::IpVersion, - }; - - use super::{ - create_connector_by_url, should_background_p2p_with_peer, should_try_p2p_with_peer, - }; - - #[tokio::test] - async fn connector_rejects_easytier_managed_ipv6_destination() { - let global_ctx = get_mock_global_ctx(); - let public_route: cidr::Ipv6Inet = "2001:db8::2/128".parse().unwrap(); - global_ctx.set_public_ipv6_routes(BTreeSet::from([public_route])); - - let ret = - create_connector_by_url("tcp://[2001:db8::2]:11010", &global_ctx, IpVersion::V6).await; - - assert!(matches!( - ret, - Err(crate::common::error::Error::InvalidUrl(_)) - )); - } - - #[test] - fn lazy_background_p2p_requires_need_p2p() { - let no_need_p2p = PeerFeatureFlag { - need_p2p: false, - ..Default::default() - }; - let need_p2p = PeerFeatureFlag { - need_p2p: true, - ..Default::default() - }; - - assert!(should_background_p2p_with_peer( - Some(&no_need_p2p), - false, - false, - false, - false - )); - assert!(!should_background_p2p_with_peer( - Some(&no_need_p2p), - false, - true, - false, - false - )); - assert!(should_background_p2p_with_peer( - Some(&need_p2p), - false, - true, - false, - false - )); - } - - #[test] - fn p2p_policy_respects_public_server_setting() { - let public_server = PeerFeatureFlag { - is_public_server: true, - ..Default::default() - }; - - assert!(!should_try_p2p_with_peer( - Some(&public_server), - false, - false, - false - )); - assert!(should_try_p2p_with_peer( - Some(&public_server), - true, - false, - false - )); - assert!(!should_background_p2p_with_peer( - Some(&public_server), - false, - false, - false, - false - )); - assert!(should_background_p2p_with_peer( - Some(&public_server), - true, - false, - false, - false - )); - } - - #[test] - fn disable_p2p_only_allows_need_p2p_exceptions() { - let normal_peer = PeerFeatureFlag::default(); - let need_peer = PeerFeatureFlag { - need_p2p: true, - ..Default::default() - }; - let disable_peer = PeerFeatureFlag { - disable_p2p: true, - ..Default::default() - }; - let disable_need_peer = PeerFeatureFlag { - disable_p2p: true, - need_p2p: true, - ..Default::default() - }; - - assert!(should_try_p2p_with_peer( - Some(&normal_peer), - false, - false, - false - )); - assert!(should_try_p2p_with_peer(None, false, false, false)); - assert!(!should_try_p2p_with_peer(None, false, true, false)); - assert!(!should_try_p2p_with_peer( - Some(&normal_peer), - false, - true, - false - )); - assert!(should_try_p2p_with_peer( - Some(&need_peer), - false, - true, - false - )); - assert!(!should_try_p2p_with_peer( - Some(&disable_peer), - false, - false, - false - )); - assert!(should_try_p2p_with_peer( - Some(&disable_peer), - false, - false, - true - )); - assert!(should_try_p2p_with_peer( - Some(&disable_need_peer), - false, - true, - true - )); - assert!(!should_try_p2p_with_peer( - Some(&disable_need_peer), - false, - true, - false - )); - } -} diff --git a/easytier/src/connector/tcp_hole_punch.rs b/easytier/src/connector/tcp_hole_punch.rs deleted file mode 100644 index 76dfbb8b..00000000 --- a/easytier/src/connector/tcp_hole_punch.rs +++ /dev/null @@ -1,778 +0,0 @@ -use std::{ - net::{IpAddr, Ipv4Addr, Ipv6Addr, SocketAddr}, - sync::Arc, - time::Duration, -}; - -use anyhow::{Context, Error}; -use quanta::Instant; -use rand::Rng as _; -use tokio::task::JoinSet; - -use crate::{ - common::{PeerId, join_joinset_background, stun::StunInfoCollectorTrait}, - connector::udp_hole_punch::BackOff, - peers::{ - peer_manager::PeerManager, - peer_task::{PeerTaskLauncher, PeerTaskManager}, - }, - proto::{ - common::NatType, - peer_rpc::{ - TcpHolePunchRequest, TcpHolePunchResponse, TcpHolePunchRpc, - TcpHolePunchRpcClientFactory, TcpHolePunchRpcServer, - }, - rpc_types::{self, controller::BaseController}, - }, - tunnel::{ - TunnelConnector as _, TunnelListener as _, - tcp::{TcpTunnelConnector, TcpTunnelListener}, - }, -}; - -use crate::connector::{should_background_p2p_with_peer, should_try_p2p_with_peer}; - -pub const BLACKLIST_TIMEOUT_SEC: u64 = 3600; - -fn handle_rpc_result( - ret: Result, - dst_peer_id: PeerId, - blacklist: &timedmap::TimedMap, -) -> Result { - match ret { - Ok(ret) => Ok(ret), - Err(e) => { - if matches!(e, rpc_types::error::Error::InvalidServiceKey(_, _)) { - blacklist.insert(dst_peer_id, (), Duration::from_secs(BLACKLIST_TIMEOUT_SEC)); - } - Err(e) - } - } -} - -fn is_symmetric_tcp_nat(nat_type: NatType) -> bool { - matches!( - nat_type, - NatType::Symmetric | NatType::SymmetricEasyInc | NatType::SymmetricEasyDec - ) -} - -fn bind_addr_for_port(port: u16, is_v6: bool) -> SocketAddr { - if is_v6 { - SocketAddr::new(IpAddr::V6(Ipv6Addr::UNSPECIFIED), port) - } else { - SocketAddr::new(IpAddr::V4(Ipv4Addr::UNSPECIFIED), port) - } -} - -async fn select_local_port(peer_mgr: &Arc, is_v6: bool) -> Result { - let bind_addr = bind_addr_for_port(0, is_v6); - tracing::trace!(?bind_addr, is_v6, "tcp hole punch select local port"); - let _g = peer_mgr.get_global_ctx().net_ns.guard(); - let listener = tokio::net::TcpListener::bind(bind_addr).await?; - let port = listener.local_addr()?.port(); - tracing::debug!(?bind_addr, port, "tcp hole punch selected local port"); - Ok(port) -} - -// tcp support simultaneous connect, so initiator and server can both use connect. -async fn try_connect_to_remote( - peer_mgr: Arc, - a_mapped_addr: SocketAddr, - local_port: u16, - is_client: bool, - max_attempts: u32, -) -> Result<(), Error> { - tracing::info!( - ?a_mapped_addr, - local_port, - "tcp hole punch server start connect loop" - ); - - let mut connector = - TcpTunnelConnector::new(format!("tcp://{}", a_mapped_addr).parse().unwrap()); - connector.set_bind_addrs(vec![bind_addr_for_port( - local_port, - a_mapped_addr.is_ipv6(), - )]); - - let start = tokio::time::Instant::now(); - let mut attempts: u32 = 0; - while start.elapsed() < Duration::from_secs(10) && attempts < max_attempts { - attempts = attempts.wrapping_add(1); - let _g = peer_mgr.get_global_ctx().net_ns.guard(); - if let Ok(Ok(tunnel)) = - tokio::time::timeout(Duration::from_secs(3), connector.connect()).await - { - let add_tunnel_ret = if is_client { - peer_mgr.add_client_tunnel(tunnel, false).await.map(|_| ()) - } else { - peer_mgr.add_tunnel_as_server(tunnel, false).await - }; - if let Err(e) = add_tunnel_ret { - tracing::error!( - ?a_mapped_addr, - local_port, - attempts, - ?e, - "tcp hole punch server connected and added client tunnel failed" - ); - continue; - } else { - tracing::info!( - ?a_mapped_addr, - local_port, - attempts, - is_client, - "tcp hole punch server connected and added tunnel" - ); - return Ok(()); - } - } - tracing::trace!( - ?a_mapped_addr, - local_port, - attempts, - "tcp hole punch server connect attempt failed" - ); - let sleep_ms = rand::thread_rng().gen_range(10..100); - tokio::time::sleep(Duration::from_millis(sleep_ms)).await; - } - - tracing::warn!( - ?a_mapped_addr, - local_port, - attempts, - "tcp hole punch server connect loop timeout" - ); - - Err(anyhow::anyhow!( - "tcp hole punch server connect loop timeout" - )) -} - -struct TcpHolePunchServer { - peer_mgr: Arc, - tasks: Arc>>, -} - -impl TcpHolePunchServer { - fn new(peer_mgr: Arc) -> Arc { - let tasks = Arc::new(std::sync::Mutex::new(JoinSet::new())); - join_joinset_background(tasks.clone(), "tcp hole punch server".to_string()); - Arc::new(Self { peer_mgr, tasks }) - } -} - -#[async_trait::async_trait] -impl TcpHolePunchRpc for TcpHolePunchServer { - type Controller = BaseController; - - #[tracing::instrument(skip(self), fields(a_mapped_addr = ?input.connector_mapped_addr), err)] - async fn exchange_mapped_addr( - &self, - _ctrl: Self::Controller, - input: TcpHolePunchRequest, - ) -> rpc_types::error::Result { - let my_tcp_nat_type = NatType::try_from( - self.peer_mgr - .get_global_ctx() - .get_stun_info_collector() - .get_stun_info() - .tcp_nat_type, - ) - .unwrap_or(NatType::Unknown); - tracing::debug!(?my_tcp_nat_type, "tcp hole punch rpc received"); - if matches!(my_tcp_nat_type, NatType::Unknown) { - tracing::warn!(?my_tcp_nat_type, "tcp hole punch rpc rejected (unknown)"); - return Err(anyhow::anyhow!("tcp nat type unknown not supported").into()); - } - - let a_mapped_addr = input - .connector_mapped_addr - .ok_or(anyhow::anyhow!("connector_mapped_addr is required"))?; - let a_mapped_addr: SocketAddr = a_mapped_addr.into(); - let a_ip = a_mapped_addr.ip(); - if a_ip.is_unspecified() || a_ip.is_multicast() { - tracing::warn!(?a_mapped_addr, "tcp hole punch rpc invalid connector addr"); - return Err(anyhow::anyhow!("connector_mapped_addr is malformed").into()); - } - - let is_v6 = a_mapped_addr.is_ipv6(); - let local_port = select_local_port(&self.peer_mgr, is_v6).await?; - let mapped_addr = self - .peer_mgr - .get_global_ctx() - .get_stun_info_collector() - .get_tcp_port_mapping(local_port) - .await - .with_context(|| "failed to get tcp port mapping")?; - - tracing::info!( - ?a_mapped_addr, - local_port, - ?mapped_addr, - "tcp hole punch rpc responding with listener mapped addr and start connecting" - ); - - let peer_mgr = self.peer_mgr.clone(); - self.tasks.lock().unwrap().spawn(async move { - let _ = try_connect_to_remote(peer_mgr, a_mapped_addr, local_port, true, 5).await; - }); - - Ok(TcpHolePunchResponse { - listener_mapped_addr: Some(mapped_addr.into()), - }) - } -} - -struct TcpHolePunchConnectorData { - peer_mgr: Arc, - blacklist: Arc>, -} - -impl TcpHolePunchConnectorData { - fn new(peer_mgr: Arc) -> Arc { - Arc::new(Self { - peer_mgr, - blacklist: Arc::new(timedmap::TimedMap::new()), - }) - } - - async fn punch_as_initiator(self: Arc, dst_peer_id: PeerId) -> Result<(), Error> { - let mut backoff = BackOff::new(vec![1000, 1000, 4000, 8000]); - - loop { - backoff.sleep_for_next_backoff().await; - if self.do_punch_as_initiator(dst_peer_id).await.is_ok() { - break; - } - - if self.blacklist.contains(&dst_peer_id) { - tracing::warn!( - dst_peer_id, - "tcp hole punch initiator skipped (blacklisted)" - ); - break; - } - } - - Ok(()) - } - - #[tracing::instrument(skip(self), fields(dst_peer_id), err)] - async fn do_punch_as_initiator(&self, dst_peer_id: PeerId) -> Result<(), Error> { - let global_ctx = self.peer_mgr.get_global_ctx(); - let my_tcp_nat_type = NatType::try_from( - global_ctx - .get_stun_info_collector() - .get_stun_info() - .tcp_nat_type, - ) - .unwrap_or(NatType::Unknown); - tracing::debug!(?my_tcp_nat_type, "tcp hole punch initiator start"); - if is_symmetric_tcp_nat(my_tcp_nat_type) || my_tcp_nat_type == NatType::Unknown { - tracing::debug!("tcp hole punch initiator skipped (symmetric)"); - return Ok(()); - } - - let local_port = select_local_port(&self.peer_mgr, false).await?; - let mapped_addr = global_ctx - .get_stun_info_collector() - .get_tcp_port_mapping(local_port) - .await - .with_context(|| "failed to get tcp port mapping")?; - - tracing::info!( - dst_peer_id, - local_port, - ?mapped_addr, - "tcp hole punch initiator got mapped addr, start rpc exchange" - ); - - let rpc_stub = self - .peer_mgr - .get_peer_rpc_mgr() - .rpc_client() - .scoped_client::>( - self.peer_mgr.my_peer_id(), - dst_peer_id, - global_ctx.get_network_name(), - ); - - let resp = rpc_stub - .exchange_mapped_addr( - BaseController { - timeout_ms: 6000, - ..Default::default() - }, - TcpHolePunchRequest { - connector_mapped_addr: Some(mapped_addr.into()), - }, - ) - .await; - let resp = handle_rpc_result(resp, dst_peer_id, &self.blacklist)?; - let remote_mapped_addr = resp - .listener_mapped_addr - .ok_or(anyhow::anyhow!("listener_mapped_addr is required"))?; - let remote_mapped_addr: SocketAddr = remote_mapped_addr.into(); - tracing::info!( - dst_peer_id, - ?remote_mapped_addr, - "tcp hole punch initiator rpc returned" - ); - - if let Ok(()) = try_connect_to_remote( - self.peer_mgr.clone(), - remote_mapped_addr, - local_port, - false, - 1, - ) - .await - { - tracing::info!( - dst_peer_id, - local_port, - ?remote_mapped_addr, - "tcp hole punch initiator connected to remote mapped addr with simultaneous connection" - ); - return Ok(()); - } - - tracing::debug!( - dst_peer_id, - local_port, - ?remote_mapped_addr, - "tcp hole punch initiator sent syn to remote mapped addr" - ); - - let mut listener = - TcpTunnelListener::new(format!("tcp://0.0.0.0:{}", local_port).parse().unwrap()); - { - let _g = self.peer_mgr.get_global_ctx().net_ns.guard(); - listener.listen().await?; - } - tracing::info!( - dst_peer_id, - local_port, - url = %listener.local_url(), - "tcp hole punch initiator listening" - ); - - tokio::time::timeout( - Duration::from_secs(10), - self.accept_loop(&mut listener, dst_peer_id), - ) - .await??; - - tracing::info!( - dst_peer_id, - "tcp hole punch initiator accepted and added server tunnel" - ); - - Ok(()) - } - - async fn accept_loop( - &self, - listener: &mut TcpTunnelListener, - dst_peer_id: PeerId, - ) -> Result<(), Error> { - loop { - match listener.accept().await { - Ok(tunnel) => { - if let Err(e) = self.peer_mgr.add_tunnel_as_server(tunnel, false).await { - tracing::error!("tcp hole punch add tunnel error: {}", e); - continue; - } - - tracing::info!( - dst_peer_id, - "tcp hole punch initiator accepted and added server tunnel" - ); - } - Err(e) => { - tracing::error!("tcp hole punch accept error: {}", e); - } - } - } - } -} - -#[derive(Clone, Debug, Hash, Eq, PartialEq)] -struct TcpPunchTaskInfo { - dst_peer_id: PeerId, -} - -#[derive(Clone)] -struct TcpHolePunchPeerTaskLauncher {} - -#[async_trait::async_trait] -impl PeerTaskLauncher for TcpHolePunchPeerTaskLauncher { - type Data = Arc; - type CollectPeerItem = TcpPunchTaskInfo; - type TaskRet = (); - - fn new_data(&self, peer_mgr: Arc) -> Self::Data { - TcpHolePunchConnectorData::new(peer_mgr) - } - - #[tracing::instrument(skip(self, data))] - async fn collect_peers_need_task(&self, data: &Self::Data) -> Vec { - let global_ctx = data.peer_mgr.get_global_ctx(); - let flags = global_ctx.get_flags(); - let lazy_p2p = flags.lazy_p2p; - let my_tcp_nat_type = NatType::try_from( - global_ctx - .get_stun_info_collector() - .get_stun_info() - .tcp_nat_type, - ) - .unwrap_or(NatType::Unknown); - if is_symmetric_tcp_nat(my_tcp_nat_type) || my_tcp_nat_type == NatType::Unknown { - tracing::trace!( - ?my_tcp_nat_type, - "tcp hole punch task collect skipped (symmetric)" - ); - return vec![]; - } - - let my_peer_id = data.peer_mgr.my_peer_id(); - let now = Instant::now(); - - data.blacklist.cleanup(); - - let mut peers_to_connect = Vec::new(); - for route in data.peer_mgr.list_routes().await.iter() { - let static_allowed = should_background_p2p_with_peer( - route.feature_flag.as_ref(), - false, - lazy_p2p, - flags.disable_p2p, - flags.need_p2p, - ); - let dynamic_allowed = should_try_p2p_with_peer( - route.feature_flag.as_ref(), - false, - flags.disable_p2p, - flags.need_p2p, - ) && data.peer_mgr.has_recent_traffic(route.peer_id, now); - if !static_allowed && !dynamic_allowed { - continue; - } - - let peer_id: PeerId = route.peer_id; - if peer_id == my_peer_id { - tracing::trace!(peer_id, "tcp hole punch task collect skip self"); - continue; - } - - if data.blacklist.contains(&peer_id) { - tracing::debug!(peer_id, "tcp hole punch task collect skip blacklisted"); - continue; - } - - if data.peer_mgr.get_peer_map().has_peer(peer_id) { - tracing::trace!(peer_id, "tcp hole punch task collect skip already has peer"); - continue; - } - - let peer_tcp_nat_type = route - .stun_info - .as_ref() - .map(|x| x.tcp_nat_type) - .unwrap_or(0); - let peer_tcp_nat_type = - NatType::try_from(peer_tcp_nat_type).unwrap_or(NatType::Unknown); - if matches!(peer_tcp_nat_type, NatType::Unknown) { - tracing::debug!( - peer_id, - ?peer_tcp_nat_type, - "tcp hole punch task collect skip peer unknown" - ); - continue; - } - - tracing::info!( - peer_id, - my_peer_id, - ?my_tcp_nat_type, - ?peer_tcp_nat_type, - "tcp hole punch task collect add peer" - ); - peers_to_connect.push(TcpPunchTaskInfo { - dst_peer_id: peer_id, - }); - } - - peers_to_connect - } - - async fn launch_task( - &self, - data: &Self::Data, - item: Self::CollectPeerItem, - ) -> tokio::task::JoinHandle> { - let data = data.clone(); - tokio::spawn(async move { data.punch_as_initiator(item.dst_peer_id).await.map(|_| ()) }) - } - - async fn all_task_done(&self, _data: &Self::Data) {} - - fn loop_interval_ms(&self) -> u64 { - 5000 - } -} - -pub struct TcpHolePunchConnector { - server: Arc, - client: PeerTaskManager, - peer_mgr: Arc, -} - -impl TcpHolePunchConnector { - pub fn new(peer_mgr: Arc) -> Self { - Self { - server: TcpHolePunchServer::new(peer_mgr.clone()), - client: PeerTaskManager::new_with_external_signal( - TcpHolePunchPeerTaskLauncher {}, - peer_mgr.clone(), - Some(peer_mgr.p2p_demand_notify()), - ), - peer_mgr, - } - } - - pub async fn run_as_client(&mut self) -> Result<(), Error> { - tracing::info!("tcp hole punch client start"); - self.client.start(); - Ok(()) - } - - pub async fn run_as_server(&mut self) -> Result<(), Error> { - tracing::info!("tcp hole punch server register rpc"); - self.peer_mgr - .get_peer_rpc_mgr() - .rpc_server() - .registry() - .register( - TcpHolePunchRpcServer::new_arc(self.server.clone()), - &self.peer_mgr.get_global_ctx().get_network_name(), - ); - Ok(()) - } - - pub async fn run(&mut self) -> Result<(), Error> { - let flags = self.peer_mgr.get_global_ctx().get_flags(); - if flags.disable_tcp_hole_punching { - tracing::debug!( - "tcp hole punch disabled by disable_tcp_hole_punching(={});", - flags.disable_tcp_hole_punching - ); - return Ok(()); - } - - self.run_as_client().await?; - self.run_as_server().await?; - Ok(()) - } -} - -#[cfg(test)] -mod tests { - use std::{net::SocketAddr, sync::Arc, time::Duration}; - - use crate::{ - common::{error::Error, stun::StunInfoCollectorTrait}, - connector::tcp_hole_punch::TcpHolePunchConnector, - peers::{ - peer_manager::PeerManager, - peer_task::PeerTaskLauncher, - tests::{connect_peer_manager, create_mock_peer_manager, wait_route_appear}, - }, - proto::common::{NatType, StunInfo}, - tunnel::common::tests::wait_for_condition, - }; - - use super::TcpHolePunchPeerTaskLauncher; - - struct MockStunInfoCollector { - udp_nat_type: NatType, - tcp_nat_type: NatType, - } - - #[async_trait::async_trait] - impl StunInfoCollectorTrait for MockStunInfoCollector { - fn get_stun_info(&self) -> StunInfo { - StunInfo { - udp_nat_type: self.udp_nat_type as i32, - tcp_nat_type: self.tcp_nat_type as i32, - last_update_time: 0, - public_ip: vec!["127.0.0.1".to_string(), "::1".to_string()], - min_port: 100, - max_port: 200, - } - } - - async fn get_udp_port_mapping(&self, mut port: u16) -> Result { - if port == 0 { - port = 40144; - } - Ok(format!("127.0.0.1:{}", port).parse().unwrap()) - } - - async fn get_udp_port_mapping_with_socket( - &self, - udp: std::sync::Arc, - ) -> Result { - self.get_udp_port_mapping(udp.local_addr()?.port()).await - } - - async fn get_tcp_port_mapping(&self, mut port: u16) -> Result { - if port == 0 { - port = 40144; - } - Ok(format!("127.0.0.1:{}", port).parse().unwrap()) - } - } - - fn replace_stun_info_collector(peer_mgr: Arc, tcp_nat_type: NatType) { - let collector = Box::new(MockStunInfoCollector { - udp_nat_type: NatType::Unknown, - tcp_nat_type, - }); - peer_mgr - .get_global_ctx() - .replace_stun_info_collector(collector); - } - - async fn collect_lazy_punch_peers(peer_mgr: Arc) -> Vec { - let launcher = TcpHolePunchPeerTaskLauncher {}; - let data = launcher.new_data(peer_mgr); - launcher - .collect_peers_need_task(&data) - .await - .into_iter() - .map(|task| task.dst_peer_id) - .collect() - } - - #[tokio::test] - async fn tcp_hole_punch_connects() { - let p_a = create_mock_peer_manager().await; - let p_b = create_mock_peer_manager().await; - let p_c = create_mock_peer_manager().await; - - replace_stun_info_collector(p_a.clone(), NatType::PortRestricted); - replace_stun_info_collector(p_b.clone(), NatType::PortRestricted); - replace_stun_info_collector(p_c.clone(), NatType::PortRestricted); - - connect_peer_manager(p_a.clone(), p_b.clone()).await; - connect_peer_manager(p_b.clone(), p_c.clone()).await; - wait_route_appear(p_a.clone(), p_c.clone()).await.unwrap(); - - let mut hole_punching_a = TcpHolePunchConnector::new(p_a.clone()); - let mut hole_punching_c = TcpHolePunchConnector::new(p_c.clone()); - hole_punching_a.run().await.unwrap(); - hole_punching_c.run().await.unwrap(); - - hole_punching_a.client.run_immediately().await; - hole_punching_c.client.run_immediately().await; - - wait_for_condition( - || { - let p_a = p_a.clone(); - let p_c = p_c.clone(); - async move { - let a_has = p_a - .get_peer_map() - .list_peer_conns(p_c.my_peer_id()) - .await - .is_some_and(|c| !c.is_empty()); - let c_has = p_c - .get_peer_map() - .list_peer_conns(p_a.my_peer_id()) - .await - .is_some_and(|c| !c.is_empty()); - a_has || c_has - } - }, - Duration::from_secs(15), - ) - .await; - } - - #[tokio::test] - async fn tcp_hole_punch_skip_symmetric_peer() { - let p_a = create_mock_peer_manager().await; - let p_b = create_mock_peer_manager().await; - let p_c = create_mock_peer_manager().await; - - replace_stun_info_collector(p_a.clone(), NatType::Symmetric); - replace_stun_info_collector(p_b.clone(), NatType::PortRestricted); - replace_stun_info_collector(p_c.clone(), NatType::Symmetric); - - connect_peer_manager(p_a.clone(), p_b.clone()).await; - connect_peer_manager(p_b.clone(), p_c.clone()).await; - wait_route_appear(p_a.clone(), p_c.clone()).await.unwrap(); - - let mut hole_punching_a = TcpHolePunchConnector::new(p_a.clone()); - let mut hole_punching_c = TcpHolePunchConnector::new(p_c.clone()); - hole_punching_a.run().await.unwrap(); - hole_punching_c.run().await.unwrap(); - - hole_punching_a.client.run_immediately().await; - hole_punching_c.client.run_immediately().await; - - tokio::time::sleep(Duration::from_secs(2)).await; - - assert!( - p_a.get_peer_map() - .list_peer_conns(p_c.my_peer_id()) - .await - .map(|c| c.is_empty()) - .unwrap_or(true) - ); - assert!( - p_c.get_peer_map() - .list_peer_conns(p_a.my_peer_id()) - .await - .map(|c| c.is_empty()) - .unwrap_or(true) - ); - } - - #[tokio::test] - async fn lazy_p2p_collects_tcp_hole_punch_tasks_only_after_recent_traffic() { - let p_a = create_mock_peer_manager().await; - let p_b = create_mock_peer_manager().await; - let p_c = create_mock_peer_manager().await; - - replace_stun_info_collector(p_a.clone(), NatType::PortRestricted); - replace_stun_info_collector(p_b.clone(), NatType::PortRestricted); - replace_stun_info_collector(p_c.clone(), NatType::PortRestricted); - - let mut flags = p_a.get_global_ctx().get_flags(); - flags.lazy_p2p = true; - p_a.get_global_ctx().set_flags(flags); - - connect_peer_manager(p_a.clone(), p_b.clone()).await; - connect_peer_manager(p_b.clone(), p_c.clone()).await; - wait_route_appear(p_a.clone(), p_c.clone()).await.unwrap(); - - assert!( - !collect_lazy_punch_peers(p_a.clone()) - .await - .contains(&p_c.my_peer_id()) - ); - - p_a.mark_recent_traffic(p_c.my_peer_id()); - - assert!( - collect_lazy_punch_peers(p_a.clone()) - .await - .contains(&p_c.my_peer_id()) - ); - } -} diff --git a/easytier/src/connector/udp_hole_punch/both_easy_sym.rs b/easytier/src/connector/udp_hole_punch/both_easy_sym.rs deleted file mode 100644 index c46d409a..00000000 --- a/easytier/src/connector/udp_hole_punch/both_easy_sym.rs +++ /dev/null @@ -1,422 +0,0 @@ -use std::{ - net::{IpAddr, SocketAddr, SocketAddrV4}, - sync::Arc, - time::Duration, -}; - -use anyhow::Context; -use quanta::Instant; -use tokio::sync::Mutex; -use tokio_util::task::AbortOnDropHandle; - -use crate::{ - common::{PeerId, stun::StunInfoCollectorTrait}, - connector::udp_hole_punch::common::{ - HOLE_PUNCH_PACKET_BODY_LEN, UdpHolePunchListener, try_connect_with_socket, - }, - connector::udp_hole_punch::handle_rpc_result, - peers::peer_manager::PeerManager, - proto::{ - peer_rpc::{ - SendPunchPacketBothEasySymRequest, SendPunchPacketBothEasySymResponse, - UdpHolePunchRpcClientFactory, - }, - rpc_types::{self, controller::BaseController}, - }, - tunnel::{Tunnel, udp::new_hole_punch_packet}, -}; - -use super::common::{PunchHoleServerCommon, UdpNatType, UdpSocketArray}; - -const UDP_ARRAY_SIZE_FOR_BOTH_EASY_SYM: usize = 25; -const DST_PORT_OFFSET: u16 = 20; -const REMOTE_WAIT_TIME_MS: u64 = 5000; - -pub(crate) struct PunchBothEasySymHoleServer { - common: Arc, - task: Mutex>>, -} - -impl PunchBothEasySymHoleServer { - pub(crate) fn new(common: Arc) -> Self { - Self { - common, - task: Mutex::new(None), - } - } - - // hard sym means public port is random and cannot be predicted - #[tracing::instrument(skip(self), ret, err)] - pub(crate) async fn send_punch_packet_both_easy_sym( - &self, - request: SendPunchPacketBothEasySymRequest, - ) -> Result { - tracing::info!("send_punch_packet_both_easy_sym start"); - let busy_resp = Ok(SendPunchPacketBothEasySymResponse { - is_busy: true, - ..Default::default() - }); - let Ok(mut locked_task) = self.task.try_lock() else { - return busy_resp; - }; - if locked_task.is_some() && !locked_task.as_ref().unwrap().is_finished() { - return busy_resp; - } - - let global_ctx = self.common.get_global_ctx(); - let cur_mapped_addr = global_ctx - .get_stun_info_collector() - .get_udp_port_mapping(0) - .await - .with_context(|| "failed to get udp port mapping")?; - - tracing::info!("send_punch_packet_hard_sym start"); - let socket_count = request.udp_socket_count as usize; - let public_ips = request - .public_ip - .ok_or(anyhow::anyhow!("public_ip is required"))?; - let transaction_id = request.transaction_id; - - let udp_array = - UdpSocketArray::new(socket_count, self.common.get_global_ctx().net_ns.clone()); - udp_array.start().await?; - udp_array.add_intreast_tid(transaction_id); - let peer_mgr = self.common.get_peer_mgr(); - - let punch_packet = - new_hole_punch_packet(transaction_id, HOLE_PUNCH_PACKET_BODY_LEN).into_bytes(); - let mut punched = vec![]; - let common = self.common.clone(); - - let task = tokio::spawn(async move { - let mut listeners = Vec::new(); - let start_time = Instant::now(); - let wait_time_ms = request.wait_time_ms.min(8000); - while start_time.elapsed() < Duration::from_millis(wait_time_ms as u64) { - if let Err(e) = udp_array - .send_with_all( - &punch_packet, - SocketAddr::V4(SocketAddrV4::new( - public_ips.into(), - request.dst_port_num as u16, - )), - ) - .await - { - tracing::error!(?e, "failed to send hole punch packet"); - break; - } - - tokio::time::sleep(Duration::from_millis(100)).await; - - if let Some(s) = udp_array.try_fetch_punched_socket(transaction_id) { - tracing::info!(?s, ?transaction_id, "got punched socket in both easy sym"); - assert!(Arc::strong_count(&s.socket) == 1); - let Some(port) = s.socket.local_addr().ok().map(|addr| addr.port()) else { - tracing::warn!("failed to get local addr from punched socket"); - continue; - }; - let remote_addr = s.remote_addr; - drop(s); - - let listener = - match UdpHolePunchListener::new_ext(peer_mgr.clone(), false, Some(port)) - .await - { - Ok(l) => l, - Err(e) => { - tracing::warn!(?e, "failed to create listener"); - continue; - } - }; - punched.push((listener.get_socket().await, remote_addr)); - listeners.push(listener); - } - - // if any listener is punched, we can break the loop - for l in &listeners { - if l.get_conn_count().await > 0 { - tracing::info!(?l, "got punched listener"); - break; - } - } - - if !punched.is_empty() { - tracing::debug!(?punched, "got punched socket and keep sending punch packet"); - } - - for p in &punched { - let (socket, remote_addr) = p; - let send_remote_ret = socket.send_to(&punch_packet, remote_addr).await; - tracing::debug!( - ?send_remote_ret, - ?socket, - "send hole punch packet to punched remote" - ); - } - } - - for l in listeners { - if l.get_conn_count().await > 0 { - common.add_listener(l).await; - } - } - }); - - *locked_task = Some(AbortOnDropHandle::new(task)); - return Ok(SendPunchPacketBothEasySymResponse { - is_busy: false, - base_mapped_addr: Some(cur_mapped_addr.into()), - }); - } -} - -#[derive(Debug)] -pub(crate) struct PunchBothEasySymHoleClient { - peer_mgr: Arc, - blacklist: Arc>, -} - -impl PunchBothEasySymHoleClient { - pub(crate) fn new( - peer_mgr: Arc, - blacklist: Arc>, - ) -> Self { - Self { - peer_mgr, - blacklist, - } - } - - #[tracing::instrument(ret)] - pub(crate) async fn do_hole_punching( - &self, - dst_peer_id: PeerId, - my_nat_info: UdpNatType, - peer_nat_info: UdpNatType, - is_busy: &mut bool, - ) -> Result>, anyhow::Error> { - // Check if peer is blacklisted - if self.blacklist.contains(&dst_peer_id) { - tracing::debug!(?dst_peer_id, "peer is blacklisted, skipping hole punching"); - return Ok(None); - } - - *is_busy = false; - - let udp_array = UdpSocketArray::new( - UDP_ARRAY_SIZE_FOR_BOTH_EASY_SYM, - self.peer_mgr.get_global_ctx().net_ns.clone(), - ); - udp_array.start().await?; - - let global_ctx = self.peer_mgr.get_global_ctx(); - let cur_mapped_addr = global_ctx - .get_stun_info_collector() - .get_udp_port_mapping(0) - .await - .with_context(|| "failed to get udp port mapping")?; - let my_public_ip = match cur_mapped_addr.ip() { - IpAddr::V4(v4) => v4, - _ => { - anyhow::bail!("ipv6 is not supported"); - } - }; - let me_is_incremental = my_nat_info - .get_inc_of_easy_sym() - .ok_or(anyhow::anyhow!("me_is_incremental is required"))?; - let peer_is_incremental = peer_nat_info - .get_inc_of_easy_sym() - .ok_or(anyhow::anyhow!("peer_is_incremental is required"))?; - - let rpc_stub = self - .peer_mgr - .get_peer_rpc_mgr() - .rpc_client() - .scoped_client::>( - self.peer_mgr.my_peer_id(), - dst_peer_id, - global_ctx.get_network_name(), - ); - - let tid = rand::random(); - udp_array.add_intreast_tid(tid); - - let remote_ret = rpc_stub - .send_punch_packet_both_easy_sym( - BaseController { - timeout_ms: 2000, - ..Default::default() - }, - SendPunchPacketBothEasySymRequest { - transaction_id: tid, - public_ip: Some(my_public_ip.into()), - dst_port_num: if me_is_incremental { - cur_mapped_addr.port().saturating_add(DST_PORT_OFFSET) - } else { - cur_mapped_addr.port().saturating_sub(DST_PORT_OFFSET) - } as u32, - udp_socket_count: UDP_ARRAY_SIZE_FOR_BOTH_EASY_SYM as u32, - wait_time_ms: REMOTE_WAIT_TIME_MS as u32, - }, - ) - .await; - - let remote_ret = handle_rpc_result(remote_ret, dst_peer_id, &self.blacklist)?; - - if remote_ret.is_busy { - *is_busy = true; - anyhow::bail!("remote is busy"); - } - - let mut remote_mapped_addr = remote_ret - .base_mapped_addr - .ok_or(anyhow::anyhow!("remote_mapped_addr is required"))?; - - let now = Instant::now(); - remote_mapped_addr.port = if peer_is_incremental { - remote_mapped_addr - .port - .saturating_add(DST_PORT_OFFSET as u32) - } else { - remote_mapped_addr - .port - .saturating_sub(DST_PORT_OFFSET as u32) - }; - tracing::debug!( - ?remote_mapped_addr, - ?remote_ret, - "start send hole punch packet for both easy sym" - ); - - while now.elapsed().as_millis() < (REMOTE_WAIT_TIME_MS + 1000).into() { - udp_array - .send_with_all( - &new_hole_punch_packet(tid, HOLE_PUNCH_PACKET_BODY_LEN).into_bytes(), - remote_mapped_addr.into(), - ) - .await?; - - tokio::time::sleep(Duration::from_millis(100)).await; - - let Some(socket) = udp_array.try_fetch_punched_socket(tid) else { - tracing::trace!( - ?remote_mapped_addr, - ?tid, - "no punched socket found, send some more hole punch packets" - ); - continue; - }; - - tracing::info!( - ?socket, - ?remote_mapped_addr, - ?tid, - "got punched socket in both easy sym" - ); - - for _ in 0..2 { - match try_connect_with_socket( - global_ctx.clone(), - socket.socket.clone(), - remote_mapped_addr.into(), - ) - .await - { - Ok(tunnel) => { - return Ok(Some(tunnel)); - } - Err(e) => { - tracing::error!(?e, "failed to connect with socket"); - continue; - } - } - } - udp_array.add_new_socket(socket.socket).await?; - } - - Ok(None) - } -} - -#[cfg(test)] -pub mod tests { - use std::{ - sync::{Arc, atomic::AtomicU32}, - time::Duration, - }; - - use tokio::net::UdpSocket; - - use crate::connector::udp_hole_punch::RUN_TESTING; - use crate::{ - connector::udp_hole_punch::{ - UdpHolePunchConnector, tests::create_mock_peer_manager_with_mock_stun, - }, - peers::tests::{connect_peer_manager, wait_route_appear}, - proto::common::NatType, - tunnel::common::tests::wait_for_condition, - }; - - #[rstest::rstest] - #[tokio::test] - #[serial_test::serial(hole_punch)] - async fn hole_punching_easy_sym(#[values("true", "false")] is_inc: bool) { - RUN_TESTING.store(true, std::sync::atomic::Ordering::Relaxed); - - let p_a = create_mock_peer_manager_with_mock_stun(if is_inc { - NatType::SymmetricEasyInc - } else { - NatType::SymmetricEasyDec - }) - .await; - let p_b = create_mock_peer_manager_with_mock_stun(NatType::PortRestricted).await; - let p_c = create_mock_peer_manager_with_mock_stun(if !is_inc { - NatType::SymmetricEasyInc - } else { - NatType::SymmetricEasyDec - }) - .await; - connect_peer_manager(p_a.clone(), p_b.clone()).await; - connect_peer_manager(p_b.clone(), p_c.clone()).await; - wait_route_appear(p_a.clone(), p_c.clone()).await.unwrap(); - - let mut hole_punching_a = UdpHolePunchConnector::new(p_a.clone()); - let mut hole_punching_c = UdpHolePunchConnector::new(p_c.clone()); - - hole_punching_a.run().await.unwrap(); - hole_punching_c.run().await.unwrap(); - - // 144 + DST_PORT_OFFSET = 164 - let udp1 = Arc::new(UdpSocket::bind("0.0.0.0:40164").await.unwrap()); - // 144 - DST_PORT_OFFSET = 124 - let udp2 = Arc::new(UdpSocket::bind("0.0.0.0:40124").await.unwrap()); - let udps = [udp1, udp2]; - - let counter = Arc::new(AtomicU32::new(0)); - - // all these sockets should receive hole punching packet - for udp in udps.iter().map(Arc::clone) { - let counter = counter.clone(); - tokio::spawn(async move { - let mut buf = [0u8; 1024]; - let (len, addr) = udp.recv_from(&mut buf).await.unwrap(); - println!( - "got predictable punch packet, {:?} {:?} {:?}", - len, - addr, - udp.local_addr() - ); - counter.fetch_add(1, std::sync::atomic::Ordering::Relaxed); - }); - } - - hole_punching_a.client.run_immediately().await; - let udp_len = udps.len(); - wait_for_condition( - || async { counter.load(std::sync::atomic::Ordering::Relaxed) == udp_len as u32 }, - Duration::from_secs(30), - ) - .await; - } -} diff --git a/easytier/src/connector/udp_hole_punch/common.rs b/easytier/src/connector/udp_hole_punch/common.rs deleted file mode 100644 index 28d59a27..00000000 --- a/easytier/src/connector/udp_hole_punch/common.rs +++ /dev/null @@ -1,850 +0,0 @@ -use std::{ - net::{IpAddr, Ipv4Addr, SocketAddr, SocketAddrV4}, - sync::Arc, - time::Duration, -}; - -use crossbeam::atomic::AtomicCell; -use dashmap::{DashMap, DashSet}; -use guarden::defer; -use quanta::Instant; -use rand::seq::SliceRandom as _; -use tokio::{net::UdpSocket, sync::Mutex, task::JoinSet}; -use tracing::{Instrument, Level, instrument}; -use zerocopy::FromBytes as _; - -use crate::{ - common::{ - PeerId, error::Error, global_ctx::ArcGlobalCtx, join_joinset_background, netns::NetNS, upnp, - }, - peers::peer_manager::PeerManager, - proto::common::NatType, - tunnel::{ - Tunnel, TunnelConnCounter, TunnelListener as _, - packet_def::{UDP_TUNNEL_HEADER_SIZE, UDPTunnelHeader, UdpPacketType}, - udp::{UdpTunnelConnector, UdpTunnelListener, new_hole_punch_packet}, - }, -}; - -pub(crate) const HOLE_PUNCH_PACKET_BODY_LEN: u16 = 16; -const MAX_PUBLIC_UDP_HOLE_PUNCH_LISTENERS: usize = 4; - -fn generate_shuffled_port_vec() -> Vec { - let mut rng = rand::thread_rng(); - let mut port_vec: Vec = (1..=65535).collect(); - port_vec.shuffle(&mut rng); - port_vec -} - -pub(crate) enum UdpPunchClientMethod { - None, - ConeToCone, - SymToCone, - EasySymToEasySym, -} - -#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] -pub(crate) enum UdpNatType { - Unknown, - Open(NatType), - Cone(NatType), - // bool means if it is incremental - EasySymmetric(NatType, bool), - HardSymmetric(NatType), -} - -impl From for UdpNatType { - fn from(nat_type: NatType) -> Self { - match nat_type { - NatType::Unknown => UdpNatType::Unknown, - NatType::OpenInternet => UdpNatType::Open(nat_type), - NatType::NoPat | NatType::FullCone | NatType::Restricted | NatType::PortRestricted => { - UdpNatType::Cone(nat_type) - } - NatType::Symmetric | NatType::SymUdpFirewall => UdpNatType::HardSymmetric(nat_type), - NatType::SymmetricEasyInc => UdpNatType::EasySymmetric(nat_type, true), - NatType::SymmetricEasyDec => UdpNatType::EasySymmetric(nat_type, false), - } - } -} - -impl From for NatType { - fn from(val: UdpNatType) -> Self { - match val { - UdpNatType::Unknown => NatType::Unknown, - UdpNatType::Open(nat_type) => nat_type, - UdpNatType::Cone(nat_type) => nat_type, - UdpNatType::EasySymmetric(nat_type, _) => nat_type, - UdpNatType::HardSymmetric(nat_type) => nat_type, - } - } -} - -impl UdpNatType { - pub(crate) fn is_open(&self) -> bool { - matches!(self, UdpNatType::Open(_)) - } - - pub(crate) fn is_unknown(&self) -> bool { - matches!(self, UdpNatType::Unknown) - } - - pub(crate) fn is_sym(&self) -> bool { - self.is_hard_sym() || self.is_easy_sym() - } - - pub(crate) fn is_hard_sym(&self) -> bool { - matches!(self, UdpNatType::HardSymmetric(_)) - } - - pub(crate) fn is_easy_sym(&self) -> bool { - matches!(self, UdpNatType::EasySymmetric(_, _)) - } - - pub(crate) fn is_cone(&self) -> bool { - matches!(self, UdpNatType::Cone(_)) - } - - pub(crate) fn get_inc_of_easy_sym(&self) -> Option { - match self { - UdpNatType::EasySymmetric(_, inc) => Some(*inc), - _ => None, - } - } - - pub(crate) fn get_punch_hole_method( - &self, - other: Self, - global_ctx: ArcGlobalCtx, - ) -> UdpPunchClientMethod { - // Check if symmetric NAT hole punching is disabled - let disable_sym_hole_punching = global_ctx.get_flags().disable_sym_hole_punching; - - // If symmetric NAT hole punching is disabled, treat symmetric as cone - if disable_sym_hole_punching && self.is_sym() { - // Convert symmetric to cone type for hole punching logic - if other.is_sym() { - return UdpPunchClientMethod::None; - } else { - return UdpPunchClientMethod::ConeToCone; - } - } - - if other.is_unknown() { - if self.is_sym() { - return UdpPunchClientMethod::SymToCone; - } else { - return UdpPunchClientMethod::ConeToCone; - } - } - - if self.is_unknown() { - if other.is_sym() { - return UdpPunchClientMethod::None; - } else { - return UdpPunchClientMethod::ConeToCone; - } - } - - if self.is_open() || other.is_open() { - // open nat does not need to punch hole - return UdpPunchClientMethod::None; - } - - if self.is_cone() { - if other.is_sym() { - return UdpPunchClientMethod::None; - } else { - return UdpPunchClientMethod::ConeToCone; - } - } else if self.is_easy_sym() { - if other.is_hard_sym() { - return UdpPunchClientMethod::None; - } else if other.is_easy_sym() { - return UdpPunchClientMethod::EasySymToEasySym; - } else { - return UdpPunchClientMethod::SymToCone; - } - } else if self.is_hard_sym() { - if other.is_sym() { - return UdpPunchClientMethod::None; - } else { - return UdpPunchClientMethod::SymToCone; - } - } - - unreachable!("invalid nat type"); - } - - pub(crate) fn can_punch_hole_as_client( - &self, - other: Self, - my_peer_id: PeerId, - dst_peer_id: PeerId, - global_ctx: ArcGlobalCtx, - ) -> bool { - match self.get_punch_hole_method(other, global_ctx) { - UdpPunchClientMethod::None => false, - UdpPunchClientMethod::ConeToCone | UdpPunchClientMethod::SymToCone => true, - UdpPunchClientMethod::EasySymToEasySym => my_peer_id < dst_peer_id, - } - } -} - -#[derive(Debug)] -pub(crate) struct PunchedUdpSocket { - pub(crate) socket: Arc, - pub(crate) tid: u32, - pub(crate) remote_addr: SocketAddr, -} - -// used for symmetric hole punching, binding to multiple ports to increase the chance of success -pub(crate) struct UdpSocketArray { - sockets: Arc>>, - max_socket_count: usize, - net_ns: NetNS, - tasks: Arc>>, - - intreast_tids: Arc>, - tid_to_socket: Arc>>, -} - -impl UdpSocketArray { - pub fn new(max_socket_count: usize, net_ns: NetNS) -> Self { - let tasks = Arc::new(std::sync::Mutex::new(JoinSet::new())); - join_joinset_background(tasks.clone(), "UdpSocketArray".to_owned()); - - Self { - sockets: Arc::new(DashMap::new()), - max_socket_count, - net_ns, - tasks, - - intreast_tids: Arc::new(DashSet::new()), - tid_to_socket: Arc::new(DashMap::new()), - } - } - - pub fn started(&self) -> bool { - !self.sockets.is_empty() - } - - pub async fn add_new_socket(&self, socket: Arc) -> Result<(), anyhow::Error> { - let socket_map = self.sockets.clone(); - let local_addr = socket.local_addr()?; - let intreast_tids = self.intreast_tids.clone(); - let tid_to_socket = self.tid_to_socket.clone(); - socket_map.insert(local_addr, socket.clone()); - self.tasks.lock().unwrap().spawn( - async move { - defer!(socket_map.remove(&local_addr);); - let mut buf = [0u8; UDP_TUNNEL_HEADER_SIZE + HOLE_PUNCH_PACKET_BODY_LEN as usize]; - tracing::trace!(?local_addr, "udp socket added"); - loop { - let Ok((len, addr)) = socket.recv_from(&mut buf).await else { - break; - }; - - tracing::debug!(?len, ?addr, "got raw packet"); - - if len != UDP_TUNNEL_HEADER_SIZE + HOLE_PUNCH_PACKET_BODY_LEN as usize { - continue; - } - - let Some(p) = UDPTunnelHeader::ref_from_prefix(&buf) else { - continue; - }; - - let tid = p.conn_id.get(); - let valid = p.msg_type == UdpPacketType::HolePunch as u8 - && p.len.get() == HOLE_PUNCH_PACKET_BODY_LEN; - tracing::debug!(?p, ?addr, ?tid, ?valid, ?p, "got udp hole punch packet"); - - if !valid { - continue; - } - - if intreast_tids.contains(&tid) { - tracing::info!(?addr, ?tid, "got hole punching packet with intreast tid"); - tid_to_socket - .entry(tid) - .or_default() - .push(PunchedUdpSocket { - socket: socket.clone(), - tid, - remote_addr: addr, - }); - break; - } - } - tracing::debug!(?local_addr, "udp socket recv loop end"); - } - .instrument(tracing::info_span!("udp array socket recv loop")), - ); - Ok(()) - } - - #[instrument(err)] - pub async fn start(&self) -> Result<(), anyhow::Error> { - tracing::info!("starting udp socket array"); - - while self.sockets.len() < self.max_socket_count { - let socket = { - let _g = self.net_ns.guard(); - Arc::new(UdpSocket::bind("0.0.0.0:0").await?) - }; - - self.add_new_socket(socket).await?; - } - - Ok(()) - } - - #[instrument(err)] - pub async fn send_with_all(&self, data: &[u8], addr: SocketAddr) -> Result<(), anyhow::Error> { - tracing::info!(?addr, "sending hole punching packet"); - - let sockets = self - .sockets - .iter() - .map(|s| s.value().clone()) - .collect::>(); - - for socket in sockets.iter() { - for _ in 0..3 { - socket.send_to(data, addr).await?; - } - } - - Ok(()) - } - - #[instrument(ret(level = Level::DEBUG))] - pub fn try_fetch_punched_socket(&self, tid: u32) -> Option { - tracing::debug!(?tid, "try fetch punched socket"); - self.tid_to_socket.get_mut(&tid)?.value_mut().pop() - } - - pub fn add_intreast_tid(&self, tid: u32) { - self.intreast_tids.insert(tid); - } - - pub fn remove_intreast_tid(&self, tid: u32) { - self.intreast_tids.remove(&tid); - self.tid_to_socket.remove(&tid); - } -} - -impl std::fmt::Debug for UdpSocketArray { - fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { - f.debug_struct("UdpSocketArray") - .field("sockets", &self.sockets.len()) - .field("max_socket_count", &self.max_socket_count) - .field("started", &self.started()) - .field("intreast_tids", &self.intreast_tids.len()) - .field("tid_to_socket", &self.tid_to_socket.len()) - .finish() - } -} - -#[derive(Debug)] -pub(crate) struct UdpHolePunchListener { - socket: Arc, - tasks: JoinSet<()>, - running: Arc>, - mapped_addr: SocketAddr, - has_port_mapping_lease: bool, - _port_mapping_lease: Option, - conn_counter: Arc>, - - listen_time: Instant, - last_select_time: AtomicCell, - last_active_time: Arc>, -} - -impl UdpHolePunchListener { - #[instrument(err)] - pub async fn new(peer_mgr: Arc) -> Result { - Self::new_ext(peer_mgr, true, None).await - } - - #[instrument(err)] - pub async fn new_ext( - peer_mgr: Arc, - with_mapped_addr: bool, - port: Option, - ) -> Result { - let socket = { - let _g = peer_mgr.get_global_ctx().net_ns.guard(); - Arc::new(UdpSocket::bind((Ipv4Addr::UNSPECIFIED, port.unwrap_or(0))).await?) - }; - let local_port = socket.local_addr()?.port(); - let listen_url: url::Url = format!("udp://0.0.0.0:{local_port}").parse().unwrap(); - - let (mapped_addr, port_mapping_lease) = if with_mapped_addr { - upnp::resolve_udp_public_addr(peer_mgr.get_global_ctx(), &listen_url, socket.clone()) - .await? - } else { - ( - SocketAddr::V4(SocketAddrV4::new(Ipv4Addr::UNSPECIFIED, local_port)), - None, - ) - }; - - let mut listener = UdpTunnelListener::new_with_socket(listen_url, socket.clone()); - - { - let _g = peer_mgr.get_global_ctx().net_ns.guard(); - listener.listen().await?; - } - let socket = listener.get_socket().unwrap(); - - let running = Arc::new(AtomicCell::new(true)); - let running_clone = running.clone(); - - let conn_counter = listener.get_conn_counter(); - let mut tasks = JoinSet::new(); - - tasks.spawn(async move { - while let Ok(conn) = listener.accept().await { - tracing::warn!(?conn, "udp hole punching listener got peer connection"); - let peer_mgr = peer_mgr.clone(); - tokio::spawn(async move { - if let Err(e) = peer_mgr.add_tunnel_as_server(conn, false).await { - tracing::error!( - ?e, - "failed to add tunnel as server in hole punch listener" - ); - } - }); - } - - running_clone.store(false); - }); - - let last_active_time = Arc::new(AtomicCell::new(Instant::now())); - let conn_counter_clone = conn_counter.clone(); - let last_active_time_clone = last_active_time.clone(); - tasks.spawn(async move { - loop { - tokio::time::sleep(std::time::Duration::from_secs(5)).await; - if conn_counter_clone.get().unwrap_or(0) != 0 { - last_active_time_clone.store(Instant::now()); - } - } - }); - - tracing::warn!(?mapped_addr, ?socket, "udp hole punching listener started"); - - Ok(Self { - tasks, - socket, - running, - mapped_addr, - has_port_mapping_lease: port_mapping_lease.is_some(), - _port_mapping_lease: port_mapping_lease, - conn_counter, - - listen_time: Instant::now(), - last_select_time: AtomicCell::new(Instant::now()), - last_active_time, - }) - } - - pub async fn get_socket(&self) -> Arc { - self.last_select_time.store(Instant::now()); - self.socket.clone() - } - - pub async fn get_conn_count(&self) -> usize { - self.conn_counter.get().unwrap_or(0) as usize - } -} - -pub(crate) struct PunchHoleServerCommon { - peer_mgr: Arc, - - listeners: Arc>>, - tasks: Arc>>, -} - -impl PunchHoleServerCommon { - pub(crate) fn new(peer_mgr: Arc) -> Self { - let tasks = Arc::new(std::sync::Mutex::new(JoinSet::new())); - join_joinset_background(tasks.clone(), "PunchHoleServerCommon".to_owned()); - - let listeners = Arc::new(Mutex::new(Vec::::new())); - - let l = listeners.clone(); - tasks.lock().unwrap().spawn(async move { - loop { - tokio::time::sleep(Duration::from_secs(5)).await; - { - // remove listener that is not active for 40 seconds but keep listeners that are selected less than 30 seconds - l.lock().await.retain(|listener| { - listener.last_active_time.load().elapsed().as_secs() < 40 - || listener.last_select_time.load().elapsed().as_secs() < 30 - }); - } - } - }); - - Self { - peer_mgr, - - listeners, - tasks, - } - } - - pub(crate) async fn add_listener(&self, listener: UdpHolePunchListener) { - self.listeners.lock().await.push(listener); - } - - pub(crate) async fn find_listener(&self, addr: &SocketAddr) -> Option> { - let all_listener_sockets = self.listeners.lock().await; - - let listener = all_listener_sockets - .iter() - .find(|listener| listener.mapped_addr == *addr && listener.running.load())?; - - Some(listener.get_socket().await) - } - - pub(crate) async fn my_udp_nat_type(&self) -> i32 { - self.peer_mgr - .get_global_ctx() - .get_stun_info_collector() - .get_stun_info() - .udp_nat_type - } - - #[async_recursion::async_recursion] - pub(crate) async fn select_listener( - &self, - use_new_listener: bool, - prefer_port_mapping: bool, - ) -> Option<(Arc, SocketAddr)> { - let (listener_count, has_reusable_listener, has_port_mapping_listener) = { - let locked = self.listeners.lock().await; - ( - locked.len(), - locked.iter().any(can_reuse_public_listener), - locked.iter().any(can_reuse_port_mapping_listener), - ) - }; - let should_create = should_create_public_listener( - listener_count, - has_reusable_listener, - has_port_mapping_listener, - use_new_listener, - prefer_port_mapping, - ); - - if should_create { - tracing::warn!( - max_listeners = MAX_PUBLIC_UDP_HOLE_PUNCH_LISTENERS, - "creating udp hole punching listener" - ); - match UdpHolePunchListener::new(self.peer_mgr.clone()).await { - Ok(listener) => self.listeners.lock().await.push(listener), - Err(err) => { - tracing::warn!(?err, "failed to create udp hole punching listener"); - } - } - } - - let mut locked = self.listeners.lock().await; - let listener_count = locked.len(); - let listener_idx = if prefer_port_mapping { - select_reusable_port_mapping_listener_idx(locked.as_slice()) - .or_else(|| { - if should_create && locked.last().is_some_and(can_reuse_public_listener) { - Some(locked.len() - 1) - } else { - None - } - }) - .or_else(|| select_reusable_public_listener_idx(locked.as_slice())) - } else if should_create { - locked.len().checked_sub(1) - } else { - select_reusable_public_listener_idx(locked.as_slice()) - }; - - let Some(listener_idx) = listener_idx else { - tracing::warn!( - ?use_new_listener, - ?prefer_port_mapping, - listener_count, - max_listeners = MAX_PUBLIC_UDP_HOLE_PUNCH_LISTENERS, - "no available udp hole punching listener with mapped address" - ); - if should_retry_public_listener_selection( - use_new_listener, - listener_count, - prefer_port_mapping, - has_port_mapping_listener, - ) { - drop(locked); - return self.select_listener(true, prefer_port_mapping).await; - } - return None; - }; - - let listener = &mut locked[listener_idx]; - if !can_reuse_public_listener(listener) { - tracing::warn!( - ?use_new_listener, - ?prefer_port_mapping, - listener_count, - max_listeners = MAX_PUBLIC_UDP_HOLE_PUNCH_LISTENERS, - "selected udp hole punching listener is not reusable" - ); - return None; - } - - Some((listener.get_socket().await, listener.mapped_addr)) - } - - pub(crate) fn get_joinset(&self) -> Arc>> { - self.tasks.clone() - } - - pub(crate) fn get_global_ctx(&self) -> ArcGlobalCtx { - self.peer_mgr.get_global_ctx() - } - - pub(crate) fn get_peer_mgr(&self) -> Arc { - self.peer_mgr.clone() - } -} - -fn can_reuse_public_listener(listener: &UdpHolePunchListener) -> bool { - listener.running.load() && !listener.mapped_addr.ip().is_unspecified() -} - -fn can_reuse_port_mapping_listener(listener: &UdpHolePunchListener) -> bool { - can_reuse_public_listener(listener) && listener.has_port_mapping_lease -} - -fn select_reusable_public_listener_idx(listeners: &[UdpHolePunchListener]) -> Option { - // Reuse the listener that was active most recently. - listeners - .iter() - .enumerate() - .filter(|(_, listener)| can_reuse_public_listener(listener)) - .max_by_key(|(_, listener)| listener.last_active_time.load()) - .map(|(idx, _)| idx) -} - -fn select_reusable_port_mapping_listener_idx(listeners: &[UdpHolePunchListener]) -> Option { - listeners - .iter() - .enumerate() - .filter(|(_, listener)| can_reuse_port_mapping_listener(listener)) - .max_by_key(|(_, listener)| listener.last_active_time.load()) - .map(|(idx, _)| idx) -} - -fn should_create_public_listener( - current_listener_count: usize, - has_reusable_listener: bool, - has_port_mapping_listener: bool, - force_new_listener: bool, - prefer_port_mapping: bool, -) -> bool { - if current_listener_count >= MAX_PUBLIC_UDP_HOLE_PUNCH_LISTENERS { - return false; - } - - if current_listener_count == 0 { - return true; - } - - if force_new_listener { - return true; - } - - if prefer_port_mapping && !has_port_mapping_listener { - return true; - } - - !has_reusable_listener -} - -fn should_retry_public_listener_selection( - force_new_listener: bool, - current_listener_count: usize, - prefer_port_mapping: bool, - has_port_mapping_listener: bool, -) -> bool { - if prefer_port_mapping && has_port_mapping_listener { - return false; - } - - !force_new_listener && current_listener_count < MAX_PUBLIC_UDP_HOLE_PUNCH_LISTENERS -} - -#[tracing::instrument(err, ret(level=Level::DEBUG))] -pub(crate) async fn send_symmetric_hole_punch_packet( - ports: &[u16], - udp: Arc, - transaction_id: u32, - public_ips: &Vec, - port_start_idx: usize, - max_packets: usize, -) -> Result { - tracing::debug!("sending hard symmetric hole punching packet"); - let mut sent_packets = 0; - let mut cur_port_idx = port_start_idx; - while sent_packets < max_packets { - let port = ports[cur_port_idx % ports.len()]; - for pub_ip in public_ips { - let addr = SocketAddr::V4(SocketAddrV4::new(*pub_ip, port)); - for _ in 0..3 { - let packet = new_hole_punch_packet(transaction_id, HOLE_PUNCH_PACKET_BODY_LEN); - udp.send_to(&packet.into_bytes(), addr).await?; - } - sent_packets += 1; - } - cur_port_idx = cur_port_idx.wrapping_add(1); - tokio::time::sleep(Duration::from_millis(1)).await; - } - Ok(cur_port_idx % ports.len()) -} - -async fn check_udp_socket_local_addr( - global_ctx: ArcGlobalCtx, - remote_mapped_addr: SocketAddr, -) -> Result<(), Error> { - let socket = UdpSocket::bind("0.0.0.0:0").await?; - socket.connect(remote_mapped_addr).await?; - if let Ok(local_addr) = socket.local_addr() - && let Some(err) = easytier_managed_local_addr_error(&global_ctx, local_addr) - { - return Err(anyhow::anyhow!(err).into()); - } - - Ok(()) -} - -fn easytier_managed_local_addr_error( - global_ctx: &ArcGlobalCtx, - local_addr: SocketAddr, -) -> Option<&'static str> { - // local_addr should not be equal to an EasyTier-managed virtual/public address. - match local_addr.ip() { - IpAddr::V4(ip) if global_ctx.get_ipv4().map(|ip| ip.address()) == Some(ip) => { - Some("local address is virtual ipv4") - } - IpAddr::V6(ip) if global_ctx.is_ip_easytier_managed_ipv6(&ip) => { - Some("local address is easytier-managed ipv6") - } - _ => None, - } -} - -pub(crate) async fn try_connect_with_socket( - global_ctx: ArcGlobalCtx, - socket: Arc, - remote_mapped_addr: SocketAddr, -) -> Result, Error> { - let connector = UdpTunnelConnector::new( - format!( - "udp://{}:{}", - remote_mapped_addr.ip(), - remote_mapped_addr.port() - ) - .parse() - .unwrap(), - ); - - check_udp_socket_local_addr(global_ctx, remote_mapped_addr).await?; - - connector - .try_connect_with_socket(socket, remote_mapped_addr) - .await - .map_err(Error::from) -} - -#[cfg(test)] -mod tests { - use std::{collections::BTreeSet, net::SocketAddr}; - - use crate::common::global_ctx::tests::get_mock_global_ctx; - - use super::{ - MAX_PUBLIC_UDP_HOLE_PUNCH_LISTENERS, easytier_managed_local_addr_error, - should_create_public_listener, should_retry_public_listener_selection, - }; - - #[tokio::test] - async fn local_addr_check_rejects_easytier_public_ipv6_route() { - let global_ctx = get_mock_global_ctx(); - let public_route: cidr::Ipv6Inet = "2001:db8::4/128".parse().unwrap(); - global_ctx.set_public_ipv6_routes(BTreeSet::from([public_route])); - - let local_addr: SocketAddr = "[2001:db8::4]:1234".parse().unwrap(); - - assert_eq!( - easytier_managed_local_addr_error(&global_ctx, local_addr), - Some("local address is easytier-managed ipv6") - ); - } - - #[test] - fn listener_selection_prefers_reuse_before_cap() { - assert!(!should_create_public_listener(1, true, true, false, false)); - assert!(!should_create_public_listener( - MAX_PUBLIC_UDP_HOLE_PUNCH_LISTENERS, - true, - true, - false, - false - )); - } - - #[test] - fn listener_selection_creates_when_empty_or_no_reusable_listener() { - assert!(should_create_public_listener(0, false, false, false, false)); - assert!(should_create_public_listener(1, false, false, false, false)); - } - - #[test] - fn listener_selection_force_new_respects_cap() { - assert!(should_create_public_listener(1, true, true, true, false)); - assert!(!should_create_public_listener( - MAX_PUBLIC_UDP_HOLE_PUNCH_LISTENERS, - true, - true, - true, - false - )); - } - - #[test] - fn listener_selection_prefers_port_mapping_until_available() { - assert!(should_create_public_listener(1, true, false, false, true)); - assert!(!should_create_public_listener(1, true, true, false, true)); - } - - #[test] - fn listener_selection_retry_respects_cap() { - assert!(should_retry_public_listener_selection( - false, 1, false, false - )); - assert!(!should_retry_public_listener_selection( - false, - MAX_PUBLIC_UDP_HOLE_PUNCH_LISTENERS, - false, - false - )); - assert!(!should_retry_public_listener_selection( - true, 1, false, false - )); - assert!(!should_retry_public_listener_selection( - false, 1, true, true - )); - } -} diff --git a/easytier/src/connector/udp_hole_punch/cone.rs b/easytier/src/connector/udp_hole_punch/cone.rs deleted file mode 100644 index bfeded52..00000000 --- a/easytier/src/connector/udp_hole_punch/cone.rs +++ /dev/null @@ -1,308 +0,0 @@ -use std::{sync::Arc, time::Duration}; - -use anyhow::Context; -use quanta::Instant; -use tokio::net::UdpSocket; -use tokio_util::task::AbortOnDropHandle; - -use crate::{ - common::{PeerId, upnp}, - connector::udp_hole_punch::common::{ - HOLE_PUNCH_PACKET_BODY_LEN, UdpSocketArray, try_connect_with_socket, - }, - connector::udp_hole_punch::handle_rpc_result, - peers::peer_manager::PeerManager, - proto::{ - common::Void, - peer_rpc::{ - SelectPunchListenerRequest, SendPunchPacketConeRequest, UdpHolePunchRpcClientFactory, - }, - rpc_types::{self, controller::BaseController}, - }, - tunnel::{Tunnel, udp::new_hole_punch_packet}, -}; - -use super::common::PunchHoleServerCommon; - -pub(crate) struct PunchConeHoleServer { - common: Arc, -} - -impl PunchConeHoleServer { - pub(crate) fn new(common: Arc) -> Self { - Self { common } - } - - #[tracing::instrument(skip(self), ret, err)] - pub(crate) async fn send_punch_packet_cone( - &self, - _: BaseController, - request: SendPunchPacketConeRequest, - ) -> Result { - let listener_addr = request.listener_mapped_addr.ok_or(anyhow::anyhow!( - "send_punch_packet_for_cone request missing listener_mapped_addr" - ))?; - let listener_addr = std::net::SocketAddr::from(listener_addr); - let listener = self - .common - .find_listener(&listener_addr) - .await - .ok_or(anyhow::anyhow!( - "send_punch_packet_for_cone failed to find listener" - ))?; - - let dest_addr = request.dest_addr.ok_or(anyhow::anyhow!( - "send_punch_packet_for_cone request missing dest_addr" - ))?; - let dest_addr = std::net::SocketAddr::from(dest_addr); - let dest_ip = dest_addr.ip(); - if dest_ip.is_unspecified() || dest_ip.is_multicast() { - return Err(anyhow::anyhow!( - "send_punch_packet_for_cone dest_ip is malformed, {:?}", - request - ) - .into()); - } - - for _ in 0..request.packet_batch_count { - tracing::info!(?request, "sending hole punching packet"); - - for _ in 0..request.packet_count_per_batch { - let udp_packet = - new_hole_punch_packet(request.transaction_id, HOLE_PUNCH_PACKET_BODY_LEN); - if let Err(e) = listener.send_to(&udp_packet.into_bytes(), &dest_addr).await { - tracing::error!(?e, "failed to send hole punch packet to dest addr"); - } - } - tokio::time::sleep(Duration::from_millis(request.packet_interval_ms as u64)).await; - } - - Ok(Void::default()) - } -} - -pub(crate) struct PunchConeHoleClient { - peer_mgr: Arc, - blacklist: Arc>, -} - -impl PunchConeHoleClient { - pub(crate) fn new( - peer_mgr: Arc, - blacklist: Arc>, - ) -> Self { - Self { - peer_mgr, - blacklist, - } - } - - pub(crate) async fn do_hole_punching( - &self, - dst_peer_id: PeerId, - ) -> Result>, anyhow::Error> { - // Check if peer is blacklisted - if self.blacklist.contains(&dst_peer_id) { - tracing::debug!(?dst_peer_id, "peer is blacklisted, skipping hole punching"); - return Ok(None); - } - - tracing::info!(?dst_peer_id, "start hole punching"); - let tid = rand::random(); - - let global_ctx = self.peer_mgr.get_global_ctx(); - let udp_array = UdpSocketArray::new(1, global_ctx.net_ns.clone()); - - let rpc_stub = self - .peer_mgr - .get_peer_rpc_mgr() - .rpc_client() - .scoped_client::>( - self.peer_mgr.my_peer_id(), - dst_peer_id, - global_ctx.get_network_name(), - ); - - let resp = rpc_stub - .select_punch_listener( - BaseController::default(), - SelectPunchListenerRequest { - force_new: false, - prefer_port_mapping: true, - }, - ) - .await; - - let resp = handle_rpc_result(resp, dst_peer_id, &self.blacklist)?; - - let remote_mapped_addr = resp.listener_mapped_addr.ok_or(anyhow::anyhow!( - "select_punch_listener response missing listener_mapped_addr" - ))?; - - let local_socket = { - let _g = self.peer_mgr.get_global_ctx().net_ns.guard(); - Arc::new(UdpSocket::bind("0.0.0.0:0").await?) - }; - let local_addr = local_socket - .local_addr() - .with_context(|| "failed to get local addr from udp punch socket")?; - let local_listener: url::Url = format!("udp://0.0.0.0:{}", local_addr.port()) - .parse() - .unwrap(); - let (local_mapped_addr, _local_port_mapping_lease) = upnp::resolve_udp_public_addr( - global_ctx.clone(), - &local_listener, - local_socket.clone(), - ) - .await - .with_context(|| "failed to resolve udp public addr for cone hole punch")?; - - tracing::debug!( - ?local_mapped_addr, - ?remote_mapped_addr, - "hole punch got remote listener" - ); - - udp_array.add_new_socket(local_socket).await?; - udp_array.add_intreast_tid(tid); - let send_from_local = || async { - udp_array - .send_with_all( - &new_hole_punch_packet(tid, HOLE_PUNCH_PACKET_BODY_LEN).into_bytes(), - remote_mapped_addr.into(), - ) - .await - .with_context(|| "failed to send hole punch packet from local") - }; - - send_from_local().await?; - - let punch_task = AbortOnDropHandle::new(tokio::spawn(async move { - if let Err(e) = rpc_stub - .send_punch_packet_cone( - BaseController { - timeout_ms: 4000, - ..Default::default() - }, - SendPunchPacketConeRequest { - listener_mapped_addr: Some(remote_mapped_addr), - dest_addr: Some(local_mapped_addr.into()), - transaction_id: tid, - packet_count_per_batch: 2, - packet_batch_count: 5, - packet_interval_ms: 400, - }, - ) - .await - { - tracing::error!(?e, "failed to call remote send punch packet"); - } - })); - - // server: will send some punching resps, total 10 packets. - // client: use the socket to create UdpTunnel with UdpTunnelConnector - // NOTICE: UdpTunnelConnector will ignore the punching resp packet sent by remote. - let mut finish_time: Option = None; - while finish_time.is_none() || finish_time.as_ref().unwrap().elapsed().as_millis() < 1000 { - tokio::time::sleep(Duration::from_millis(200)).await; - - if finish_time.is_none() && punch_task.is_finished() { - finish_time = Some(Instant::now()); - } - - let Some(socket) = udp_array.try_fetch_punched_socket(tid) else { - tracing::debug!("no punched socket found, send some more hole punch packets"); - send_from_local().await?; - continue; - }; - - tracing::debug!(?socket, ?tid, "punched socket found, try connect with it"); - - for _ in 0..2 { - match try_connect_with_socket( - global_ctx.clone(), - socket.socket.clone(), - remote_mapped_addr.into(), - ) - .await - { - Ok(tunnel) => { - tracing::info!(?tunnel, "hole punched"); - return Ok(Some(tunnel)); - } - Err(e) => { - tracing::error!(?e, "failed to connect with socket"); - } - } - } - } - - Ok(None) - } -} - -#[cfg(test)] -pub mod tests { - use std::sync::Arc; - - use crate::{ - common::upnp::{ - reset_udp_port_mapping_attempts_for_test, udp_port_mapping_attempts_for_test, - }, - connector::udp_hole_punch::{ - UdpHolePunchConnector, cone::PunchConeHoleClient, - tests::create_mock_peer_manager_with_mock_stun, - }, - peers::tests::{connect_peer_manager, wait_route_appear, wait_route_appear_with_cost}, - proto::common::NatType, - }; - - #[tokio::test] - async fn hole_punching_cone() { - let p_a = create_mock_peer_manager_with_mock_stun(NatType::Restricted).await; - let p_b = create_mock_peer_manager_with_mock_stun(NatType::PortRestricted).await; - let p_c = create_mock_peer_manager_with_mock_stun(NatType::Restricted).await; - connect_peer_manager(p_a.clone(), p_b.clone()).await; - connect_peer_manager(p_b.clone(), p_c.clone()).await; - - wait_route_appear(p_a.clone(), p_c.clone()).await.unwrap(); - - println!("{:?}", p_a.list_routes().await); - - let mut hole_punching_a = UdpHolePunchConnector::new(p_a.clone()); - let mut hole_punching_c = UdpHolePunchConnector::new(p_c.clone()); - - hole_punching_a.run_as_client().await.unwrap(); - hole_punching_c.run_as_server().await.unwrap(); - - hole_punching_a.client.run_immediately().await; - - wait_route_appear_with_cost(p_a.clone(), p_c.my_peer_id(), Some(1)) - .await - .unwrap(); - println!("{:?}", p_a.list_routes().await); - } - - #[tokio::test] - async fn cone_hole_punch_does_not_create_upnp_mapping_before_listener_rpc_succeeds() { - let p_a = create_mock_peer_manager_with_mock_stun(NatType::Restricted).await; - let p_b = create_mock_peer_manager_with_mock_stun(NatType::PortRestricted).await; - let p_c = create_mock_peer_manager_with_mock_stun(NatType::Restricted).await; - connect_peer_manager(p_a.clone(), p_b.clone()).await; - connect_peer_manager(p_b.clone(), p_c.clone()).await; - wait_route_appear(p_a.clone(), p_c.clone()).await.unwrap(); - - let mut flags = p_a.get_global_ctx().get_flags(); - flags.disable_upnp = false; - p_a.get_global_ctx().set_flags(flags); - - reset_udp_port_mapping_attempts_for_test(); - - let ret = PunchConeHoleClient::new(p_a.clone(), Arc::new(timedmap::TimedMap::new())) - .do_hole_punching(p_c.my_peer_id()) - .await; - - assert!(ret.is_err()); - assert_eq!(udp_port_mapping_attempts_for_test(), 0); - } -} diff --git a/easytier/src/connector/udp_hole_punch/mod.rs b/easytier/src/connector/udp_hole_punch/mod.rs deleted file mode 100644 index 1586df15..00000000 --- a/easytier/src/connector/udp_hole_punch/mod.rs +++ /dev/null @@ -1,702 +0,0 @@ -use std::{ - sync::{Arc, atomic::AtomicBool}, - time::Duration, -}; - -use anyhow::{Context, Error}; -use both_easy_sym::{PunchBothEasySymHoleClient, PunchBothEasySymHoleServer}; -use common::{PunchHoleServerCommon, UdpNatType, UdpPunchClientMethod}; -use cone::{PunchConeHoleClient, PunchConeHoleServer}; -use dashmap::DashMap; -use once_cell::sync::Lazy; -use quanta::Instant; -use sym_to_cone::{PunchSymToConeHoleClient, PunchSymToConeHoleServer}; -use tokio::{sync::Mutex, task::JoinHandle}; - -use crate::{ - common::{PeerId, stun::StunInfoCollectorTrait}, - peers::{ - peer_manager::PeerManager, - peer_task::{PeerTaskLauncher, PeerTaskManager}, - }, - proto::{ - common::{NatType, Void}, - peer_rpc::{ - SelectPunchListenerRequest, SelectPunchListenerResponse, - SendPunchPacketBothEasySymRequest, SendPunchPacketBothEasySymResponse, - SendPunchPacketConeRequest, SendPunchPacketEasySymRequest, - SendPunchPacketHardSymRequest, SendPunchPacketHardSymResponse, UdpHolePunchRpc, - UdpHolePunchRpcServer, - }, - rpc_types::{self, controller::BaseController}, - }, - tunnel::Tunnel, -}; - -use crate::connector::{should_background_p2p_with_peer, should_try_p2p_with_peer}; - -pub(crate) mod both_easy_sym; -pub(crate) mod common; -pub(crate) mod cone; -pub(crate) mod sym_to_cone; - -// sym punch should be serialized -static SYM_PUNCH_LOCK: Lazy>>> = Lazy::new(DashMap::new); -pub static RUN_TESTING: Lazy = Lazy::new(|| AtomicBool::new(false)); - -// Blacklist timeout in seconds -pub const BLACKLIST_TIMEOUT_SEC: u64 = 3600; - -fn get_sym_punch_lock(peer_id: PeerId) -> Arc> { - SYM_PUNCH_LOCK - .entry(peer_id) - .or_insert_with(|| Arc::new(Mutex::new(()))) - .value() - .clone() -} - -struct UdpHolePunchServer { - common: Arc, - cone_server: PunchConeHoleServer, - sym_to_cone_server: PunchSymToConeHoleServer, - both_easy_sym_server: PunchBothEasySymHoleServer, -} - -impl UdpHolePunchServer { - pub fn new(peer_mgr: Arc) -> Arc { - let common = Arc::new(PunchHoleServerCommon::new(peer_mgr)); - let cone_server = PunchConeHoleServer::new(common.clone()); - let sym_to_cone_server = PunchSymToConeHoleServer::new(common.clone()); - let both_easy_sym_server = PunchBothEasySymHoleServer::new(common.clone()); - - Arc::new(Self { - common, - cone_server, - sym_to_cone_server, - both_easy_sym_server, - }) - } -} - -#[async_trait::async_trait] -impl UdpHolePunchRpc for UdpHolePunchServer { - type Controller = BaseController; - - async fn select_punch_listener( - &self, - _ctrl: Self::Controller, - input: SelectPunchListenerRequest, - ) -> rpc_types::error::Result { - let (_, addr) = self - .common - .select_listener(input.force_new, input.prefer_port_mapping) - .await - .ok_or(anyhow::anyhow!("no listener available"))?; - - Ok(SelectPunchListenerResponse { - listener_mapped_addr: Some(addr.into()), - }) - } - - /// send packet to one remote_addr, used by nat1-3 to nat1-3 - async fn send_punch_packet_cone( - &self, - ctrl: Self::Controller, - input: SendPunchPacketConeRequest, - ) -> rpc_types::error::Result { - self.cone_server.send_punch_packet_cone(ctrl, input).await - } - - /// send packet to multiple remote_addr (birthday attack), used by nat4 to nat1-3 - async fn send_punch_packet_hard_sym( - &self, - _ctrl: Self::Controller, - input: SendPunchPacketHardSymRequest, - ) -> rpc_types::error::Result { - let _locked = get_sym_punch_lock(self.common.get_peer_mgr().my_peer_id()) - .try_lock_owned() - .with_context(|| "sym punch lock is busy")?; - self.sym_to_cone_server - .send_punch_packet_hard_sym(input) - .await - } - - async fn send_punch_packet_easy_sym( - &self, - _ctrl: Self::Controller, - input: SendPunchPacketEasySymRequest, - ) -> rpc_types::error::Result { - let _locked = get_sym_punch_lock(self.common.get_peer_mgr().my_peer_id()) - .try_lock_owned() - .with_context(|| "sym punch lock is busy")?; - self.sym_to_cone_server - .send_punch_packet_easy_sym(input) - .await - .map(|_| Void {}) - } - - /// nat4 to nat4 (both predictably) - async fn send_punch_packet_both_easy_sym( - &self, - _ctrl: Self::Controller, - input: SendPunchPacketBothEasySymRequest, - ) -> rpc_types::error::Result { - let _locked = get_sym_punch_lock(self.common.get_peer_mgr().my_peer_id()) - .try_lock_owned() - .with_context(|| "sym punch lock is busy")?; - self.both_easy_sym_server - .send_punch_packet_both_easy_sym(input) - .await - } -} - -#[derive(Debug)] -pub struct BackOff { - backoffs_ms: Vec, - current_idx: usize, -} - -impl BackOff { - pub fn new(backoffs_ms: Vec) -> Self { - Self { - backoffs_ms, - current_idx: 0, - } - } - - pub fn next_backoff(&mut self) -> u64 { - let backoff = self.backoffs_ms[self.current_idx]; - self.current_idx = (self.current_idx + 1).min(self.backoffs_ms.len() - 1); - backoff - } - - pub fn rollback(&mut self) { - self.current_idx = self.current_idx.saturating_sub(1); - } - - pub async fn sleep_for_next_backoff(&mut self) { - let backoff = self.next_backoff(); - if backoff > 0 { - tokio::time::sleep(tokio::time::Duration::from_millis(backoff)).await; - } - } -} - -pub fn handle_rpc_result( - ret: Result, - dst_peer_id: PeerId, - blacklist: &timedmap::TimedMap, -) -> Result { - match ret { - Ok(ret) => Ok(ret), - Err(e) => { - if matches!(e, rpc_types::error::Error::InvalidServiceKey(_, _)) { - blacklist.insert(dst_peer_id, (), Duration::from_secs(BLACKLIST_TIMEOUT_SEC)); - } - Err(e) - } - } -} - -struct UdpHoePunchConnectorData { - cone_client: PunchConeHoleClient, - sym_to_cone_client: PunchSymToConeHoleClient, - both_easy_sym_client: PunchBothEasySymHoleClient, - peer_mgr: Arc, - blacklist: Arc>, -} - -impl UdpHoePunchConnectorData { - pub fn new(peer_mgr: Arc) -> Arc { - let blacklist = Arc::new(timedmap::TimedMap::new()); - let cone_client = PunchConeHoleClient::new(peer_mgr.clone(), blacklist.clone()); - let sym_to_cone_client = PunchSymToConeHoleClient::new(peer_mgr.clone(), blacklist.clone()); - let both_easy_sym_client = - PunchBothEasySymHoleClient::new(peer_mgr.clone(), blacklist.clone()); - - Arc::new(Self { - cone_client, - sym_to_cone_client, - both_easy_sym_client, - peer_mgr, - blacklist, - }) - } - - #[tracing::instrument(skip(self))] - async fn handle_punch_result( - &self, - ret: Result>, Error>, - backoff: Option<&mut BackOff>, - round: Option<&mut u32>, - ) -> bool { - let op = |rollback: bool| { - if rollback { - if let Some(backoff) = backoff { - backoff.rollback(); - } - if let Some(round) = round { - *round = round.saturating_sub(1); - } - } else if let Some(round) = round { - *round += 1; - } - }; - - match ret { - Ok(Some(tunnel)) => { - tracing::info!(?tunnel, "hole punching get tunnel success"); - - if let Err(e) = self.peer_mgr.add_client_tunnel(tunnel, false).await { - tracing::warn!("add client tunnel failed, err: {}", e); - op(true); - false - } else { - true - } - } - Ok(None) => { - tracing::info!("hole punching failed, no punch tunnel"); - op(false); - false - } - Err(e) => { - tracing::info!("hole punching failed, err: {}", e); - op(true); - false - } - } - } - - #[tracing::instrument(skip(self))] - async fn cone_to_cone(self: Arc, task_info: PunchTaskInfo) -> Result<(), Error> { - let mut backoff = BackOff::new(vec![1000, 1000, 2000, 4000, 4000, 8000, 8000, 16000]); - - loop { - backoff.sleep_for_next_backoff().await; - - let ret = self - .cone_client - .do_hole_punching(task_info.dst_peer_id) - .await; - - if self - .handle_punch_result(ret, Some(&mut backoff), None) - .await - { - break; - } - } - - Ok(()) - } - - #[tracing::instrument(skip(self))] - async fn sym_to_cone(self: Arc, task_info: PunchTaskInfo) -> Result<(), Error> { - let mut backoff = - BackOff::new(vec![1000, 1000, 2000, 4000, 4000, 8000, 8000, 16000, 64000]); - let mut round = 0; - let mut port_idx = rand::random(); - - loop { - backoff.sleep_for_next_backoff().await; - - // always try cone first - if !RUN_TESTING.load(std::sync::atomic::Ordering::Relaxed) { - let ret = self - .cone_client - .do_hole_punching(task_info.dst_peer_id) - .await; - if self.handle_punch_result(ret, None, None).await { - break; - } - } - - let ret = { - let _lock = get_sym_punch_lock(self.peer_mgr.my_peer_id()) - .lock_owned() - .await; - self.sym_to_cone_client - .do_hole_punching( - task_info.dst_peer_id, - round, - &mut port_idx, - task_info.my_nat_type, - ) - .await - }; - - if self - .handle_punch_result(ret, Some(&mut backoff), Some(&mut round)) - .await - { - break; - } - } - - Ok(()) - } - - #[tracing::instrument(skip(self))] - async fn both_easy_sym(self: Arc, task_info: PunchTaskInfo) -> Result<(), Error> { - let mut backoff = - BackOff::new(vec![1000, 1000, 2000, 4000, 4000, 8000, 8000, 16000, 64000]); - - loop { - backoff.sleep_for_next_backoff().await; - - // always try cone first - if !RUN_TESTING.load(std::sync::atomic::Ordering::Relaxed) { - let ret = self - .cone_client - .do_hole_punching(task_info.dst_peer_id) - .await; - if self.handle_punch_result(ret, None, None).await { - break; - } - } - - let mut is_busy = false; - - let ret = { - let _lock = get_sym_punch_lock(self.peer_mgr.my_peer_id()) - .lock_owned() - .await; - self.both_easy_sym_client - .do_hole_punching( - task_info.dst_peer_id, - task_info.my_nat_type, - task_info.dst_nat_type, - &mut is_busy, - ) - .await - }; - - if is_busy { - backoff.rollback(); - } else if self - .handle_punch_result(ret, Some(&mut backoff), None) - .await - { - break; - } - } - - Ok(()) - } -} - -#[derive(Clone)] -struct UdpHolePunchPeerTaskLauncher {} - -#[derive(Clone, Debug, Hash, Eq, PartialEq)] -struct PunchTaskInfo { - dst_peer_id: PeerId, - dst_nat_type: UdpNatType, - my_nat_type: UdpNatType, -} - -#[async_trait::async_trait] -impl PeerTaskLauncher for UdpHolePunchPeerTaskLauncher { - type Data = Arc; - type CollectPeerItem = PunchTaskInfo; - type TaskRet = (); - - fn new_data(&self, peer_mgr: Arc) -> Self::Data { - UdpHoePunchConnectorData::new(peer_mgr) - } - - async fn collect_peers_need_task(&self, data: &Self::Data) -> Vec { - let my_nat_type = data - .peer_mgr - .get_global_ctx() - .get_stun_info_collector() - .get_stun_info() - .udp_nat_type; - let my_nat_type: UdpNatType = NatType::try_from(my_nat_type) - .unwrap_or(NatType::Unknown) - .into(); - if !my_nat_type.is_sym() { - data.sym_to_cone_client.clear_udp_array().await; - } - - let mut peers_to_connect: Vec = Vec::new(); - // do not do anything if: - // 1. our nat type is OpenInternet or NoPat, which means we can wait other peers to connect us - // notice that if we are unknown, we treat ourselves as cone - if my_nat_type.is_open() { - return peers_to_connect; - } - - let my_peer_id = data.peer_mgr.my_peer_id(); - let flags = data.peer_mgr.get_global_ctx().get_flags(); - let lazy_p2p = flags.lazy_p2p; - let now = Instant::now(); - - data.blacklist.cleanup(); - - // collect peer list from peer manager and do some filter: - // 1. peers without direct conns; - // 2. peers is full cone (any restricted type); - // 3. peers not in blacklist; - for route in data.peer_mgr.list_routes().await.iter() { - let static_allowed = should_background_p2p_with_peer( - route.feature_flag.as_ref(), - false, - lazy_p2p, - flags.disable_p2p, - flags.need_p2p, - ); - let dynamic_allowed = should_try_p2p_with_peer( - route.feature_flag.as_ref(), - false, - flags.disable_p2p, - flags.need_p2p, - ) && data.peer_mgr.has_recent_traffic(route.peer_id, now); - if !static_allowed && !dynamic_allowed { - continue; - } - - let peer_nat_type = route - .stun_info - .as_ref() - .map(|x| x.udp_nat_type) - .unwrap_or(0); - let Ok(peer_nat_type) = NatType::try_from(peer_nat_type) else { - continue; - }; - let peer_nat_type = peer_nat_type.into(); - - let peer_id: PeerId = route.peer_id; - - // Check if peer is blacklisted - if data.blacklist.contains(&peer_id) { - tracing::debug!(?peer_id, "peer is blacklisted, skipping"); - continue; - } - - if data.peer_mgr.get_peer_map().has_peer(peer_id) { - continue; - } - - let global_ctx = data.peer_mgr.get_global_ctx(); - if !my_nat_type.can_punch_hole_as_client(peer_nat_type, my_peer_id, peer_id, global_ctx) - { - continue; - } - - tracing::info!( - ?peer_id, - ?peer_nat_type, - ?my_nat_type, - "found peer to do hole punching" - ); - - peers_to_connect.push(PunchTaskInfo { - dst_peer_id: peer_id, - dst_nat_type: peer_nat_type, - my_nat_type, - }); - } - - peers_to_connect - } - - async fn launch_task( - &self, - data: &Self::Data, - item: Self::CollectPeerItem, - ) -> JoinHandle> { - let data = data.clone(); - let global_ctx = data.peer_mgr.get_global_ctx(); - let punch_method = item - .my_nat_type - .get_punch_hole_method(item.dst_nat_type, global_ctx); - match punch_method { - UdpPunchClientMethod::ConeToCone => tokio::spawn(data.cone_to_cone(item)), - UdpPunchClientMethod::SymToCone => tokio::spawn(data.sym_to_cone(item)), - UdpPunchClientMethod::EasySymToEasySym => tokio::spawn(data.both_easy_sym(item)), - _ => unreachable!(), - } - } - - async fn all_task_done(&self, data: &Self::Data) { - data.sym_to_cone_client.clear_udp_array().await; - } - - fn loop_interval_ms(&self) -> u64 { - 5000 - } -} - -pub struct UdpHolePunchConnector { - server: Arc, - client: PeerTaskManager, - peer_mgr: Arc, -} - -// Currently support: -// Symmetric -> Full Cone -// Any Type of Full Cone -> Any Type of Full Cone - -// if same level of full cone, node with smaller peer_id will be the initiator -// if different level of full cone, node with more strict level will be the initiator - -impl UdpHolePunchConnector { - pub fn new(peer_mgr: Arc) -> Self { - Self { - server: UdpHolePunchServer::new(peer_mgr.clone()), - client: PeerTaskManager::new_with_external_signal( - UdpHolePunchPeerTaskLauncher {}, - peer_mgr.clone(), - Some(peer_mgr.p2p_demand_notify()), - ), - peer_mgr, - } - } - - pub async fn run_as_client(&mut self) -> Result<(), Error> { - self.client.start(); - Ok(()) - } - - pub async fn run_as_server(&mut self) -> Result<(), Error> { - self.peer_mgr - .get_peer_rpc_mgr() - .rpc_server() - .registry() - .register( - UdpHolePunchRpcServer::new(Arc::downgrade(&self.server)), - &self.peer_mgr.get_global_ctx().get_network_name(), - ); - - Ok(()) - } - - pub async fn run(&mut self) -> Result<(), Error> { - let global_ctx = self.peer_mgr.get_global_ctx(); - - if global_ctx.get_flags().disable_udp_hole_punching { - return Ok(()); - } - - self.run_as_client().await?; - self.run_as_server().await?; - - Ok(()) - } - - #[cfg(test)] - pub async fn run_immediately_for_test(&self) { - self.client.run_immediately().await; - } -} - -#[cfg(test)] -pub mod tests { - - use std::sync::Arc; - use std::time::Duration; - - use crate::common::stun::MockStunInfoCollector; - use crate::peers::{ - peer_manager::PeerManager, - peer_task::PeerTaskLauncher, - tests::{connect_peer_manager, create_mock_peer_manager, wait_route_appear}, - }; - use crate::proto::common::NatType; - use crate::tunnel::common::tests::wait_for_condition; - - use super::{RUN_TESTING, UdpHolePunchConnector, UdpHolePunchPeerTaskLauncher}; - - pub fn replace_stun_info_collector(peer_mgr: Arc, udp_nat_type: NatType) { - let collector = Box::new(MockStunInfoCollector { udp_nat_type }); - peer_mgr - .get_global_ctx() - .replace_stun_info_collector(collector); - } - - pub async fn create_mock_peer_manager_with_mock_stun( - udp_nat_type: NatType, - ) -> Arc { - let p_a = create_mock_peer_manager().await; - let mut flags = p_a.get_global_ctx().get_flags(); - flags.disable_upnp = true; - p_a.get_global_ctx().set_flags(flags); - replace_stun_info_collector(p_a.clone(), udp_nat_type); - p_a - } - - async fn collect_lazy_punch_peers(peer_mgr: Arc) -> Vec { - let launcher = UdpHolePunchPeerTaskLauncher {}; - let data = launcher.new_data(peer_mgr); - launcher - .collect_peers_need_task(&data) - .await - .into_iter() - .map(|task| task.dst_peer_id) - .collect() - } - - #[rstest::rstest] - #[tokio::test] - pub async fn test_hole_punching_blacklist( - #[values(NatType::Symmetric, NatType::PortRestricted, NatType::Unknown)] nat_type: NatType, - ) { - RUN_TESTING.store(true, std::sync::atomic::Ordering::Relaxed); - - let p_a = create_mock_peer_manager_with_mock_stun(nat_type).await; - let p_b = create_mock_peer_manager_with_mock_stun(NatType::PortRestricted).await; - let p_c = create_mock_peer_manager_with_mock_stun(NatType::PortRestricted).await; - connect_peer_manager(p_a.clone(), p_b.clone()).await; - connect_peer_manager(p_b.clone(), p_c.clone()).await; - wait_route_appear(p_a.clone(), p_c.clone()).await.unwrap(); - - let mut hole_punching_a = UdpHolePunchConnector::new(p_a.clone()); - - hole_punching_a.run().await.unwrap(); - - hole_punching_a.client.run_immediately().await; - - wait_for_condition( - || async { - hole_punching_a - .client - .data() - .blacklist - .contains(&p_c.my_peer_id()) - }, - Duration::from_secs(10), - ) - .await; - } - - #[tokio::test] - async fn lazy_p2p_collects_udp_hole_punch_tasks_only_after_recent_traffic() { - let p_a = create_mock_peer_manager_with_mock_stun(NatType::PortRestricted).await; - let p_b = create_mock_peer_manager_with_mock_stun(NatType::PortRestricted).await; - let p_c = create_mock_peer_manager_with_mock_stun(NatType::PortRestricted).await; - - let mut flags = p_a.get_global_ctx().get_flags(); - flags.lazy_p2p = true; - p_a.get_global_ctx().set_flags(flags); - - connect_peer_manager(p_a.clone(), p_b.clone()).await; - connect_peer_manager(p_b.clone(), p_c.clone()).await; - wait_route_appear(p_a.clone(), p_c.clone()).await.unwrap(); - - assert!( - !collect_lazy_punch_peers(p_a.clone()) - .await - .contains(&p_c.my_peer_id()) - ); - - p_a.mark_recent_traffic(p_c.my_peer_id()); - - assert!( - collect_lazy_punch_peers(p_a.clone()) - .await - .contains(&p_c.my_peer_id()) - ); - } -} diff --git a/easytier/src/connector/udp_hole_punch/sym_to_cone.rs b/easytier/src/connector/udp_hole_punch/sym_to_cone.rs deleted file mode 100644 index f08fe765..00000000 --- a/easytier/src/connector/udp_hole_punch/sym_to_cone.rs +++ /dev/null @@ -1,723 +0,0 @@ -use std::{ - net::Ipv4Addr, - ops::{Div, Mul}, - sync::{ - Arc, - atomic::{AtomicBool, Ordering}, - }, - time::Duration, -}; - -use anyhow::Context; -use guarden::defer; -use quanta::Instant; -use rand::{Rng, seq::SliceRandom}; -use tokio::{net::UdpSocket, sync::RwLock}; -use tokio_util::task::AbortOnDropHandle; -use tracing::Level; - -use crate::{ - common::{PeerId, global_ctx::ArcGlobalCtx, stun::StunInfoCollectorTrait}, - connector::udp_hole_punch::{ - common::{ - HOLE_PUNCH_PACKET_BODY_LEN, send_symmetric_hole_punch_packet, try_connect_with_socket, - }, - handle_rpc_result, - }, - peers::peer_manager::PeerManager, - proto::{ - peer_rpc::{ - SelectPunchListenerRequest, SendPunchPacketEasySymRequest, - SendPunchPacketHardSymRequest, SendPunchPacketHardSymResponse, UdpHolePunchRpc, - UdpHolePunchRpcClientFactory, - }, - rpc_types::{self, controller::BaseController}, - }, - tunnel::{Tunnel, udp::new_hole_punch_packet}, -}; - -use super::common::{PunchHoleServerCommon, UdpNatType, UdpSocketArray}; - -const UDP_ARRAY_SIZE_FOR_HARD_SYM: usize = 84; - -pub(crate) struct PunchSymToConeHoleServer { - common: Arc, - - shuffled_port_vec: Arc>, -} - -impl PunchSymToConeHoleServer { - pub(crate) fn new(common: Arc) -> Self { - let mut shuffled_port_vec: Vec = (1..=65535).collect(); - shuffled_port_vec.shuffle(&mut rand::thread_rng()); - - Self { - common, - shuffled_port_vec: Arc::new(shuffled_port_vec), - } - } - - // hard sym means public port is random and cannot be predicted - #[tracing::instrument(skip(self), ret)] - pub(crate) async fn send_punch_packet_easy_sym( - &self, - request: SendPunchPacketEasySymRequest, - ) -> Result<(), rpc_types::error::Error> { - tracing::info!("send_punch_packet_easy_sym start"); - - let listener_addr = request.listener_mapped_addr.ok_or(anyhow::anyhow!( - "send_punch_packet_easy_sym request missing listener_addr" - ))?; - let listener_addr = std::net::SocketAddr::from(listener_addr); - let listener = self - .common - .find_listener(&listener_addr) - .await - .ok_or(anyhow::anyhow!( - "send_punch_packet_easy_sym failed to find listener" - ))?; - - let public_ips = request - .public_ips - .into_iter() - .map(std::net::Ipv4Addr::from) - .collect::>(); - if public_ips.is_empty() { - tracing::warn!("send_punch_packet_easy_sym got zero len public ip"); - return Err( - anyhow::anyhow!("send_punch_packet_easy_sym got zero len public ip").into(), - ); - } - - let transaction_id = request.transaction_id; - let base_port_num = request.base_port_num; - let max_port_num = request.max_port_num.max(1); - let is_incremental = request.is_incremental; - - let port_start = if is_incremental { - base_port_num.saturating_add(1) - } else { - base_port_num.saturating_sub(max_port_num) - }; - - let port_end = if is_incremental { - base_port_num.saturating_add(max_port_num) - } else { - base_port_num.saturating_sub(1) - }; - - if port_end <= port_start { - return Err(anyhow::anyhow!("send_punch_packet_easy_sym invalid port range").into()); - } - - let ports = (port_start..=port_end) - .map(|x| x as u16) - .collect::>(); - tracing::debug!( - ?ports, - ?public_ips, - "send_punch_packet_easy_sym send to ports" - ); - - for _ in 0..2 { - send_symmetric_hole_punch_packet( - &ports, - listener.clone(), - transaction_id, - &public_ips, - 0, - ports.len(), - ) - .await - .with_context(|| "failed to send symmetric hole punch packet")?; - } - - Ok(()) - } - - // hard sym means public port is random and cannot be predicted - #[tracing::instrument(skip(self))] - pub(crate) async fn send_punch_packet_hard_sym( - &self, - request: SendPunchPacketHardSymRequest, - ) -> Result { - tracing::info!("try_punch_symmetric start"); - - let listener_addr = request.listener_mapped_addr.ok_or(anyhow::anyhow!( - "try_punch_symmetric request missing listener_addr" - ))?; - let listener_addr = std::net::SocketAddr::from(listener_addr); - let listener = self - .common - .find_listener(&listener_addr) - .await - .ok_or(anyhow::anyhow!( - "send_punch_packet_for_cone failed to find listener" - ))?; - - let public_ips = request - .public_ips - .into_iter() - .map(std::net::Ipv4Addr::from) - .collect::>(); - if public_ips.is_empty() { - tracing::warn!("try_punch_symmetric got zero len public ip"); - return Err(anyhow::anyhow!("try_punch_symmetric got zero len public ip").into()); - } - - let transaction_id = request.transaction_id; - let last_port_index = request.port_index as usize; - - let round = std::cmp::max(request.round, 1); - - // send max k1 packets if we are predicting the dst port - let max_k1: u32 = 180; - // send max k2 packets if we are sending to random port - let mut max_k2: u32 = rand::thread_rng().gen_range(600..800); - if round > 2 { - max_k2 = max_k2.mul(2).div(round).max(max_k1); - } - - let mut next_port_index = 0; - for _ in 0..2 { - next_port_index = send_symmetric_hole_punch_packet( - &self.shuffled_port_vec, - listener.clone(), - transaction_id, - &public_ips, - last_port_index, - max_k2 as usize, - ) - .await - .with_context(|| "failed to send symmetric hole punch packet randomly")?; - } - - return Ok(SendPunchPacketHardSymResponse { - next_port_index: next_port_index as u32, - }); - } -} - -pub(crate) struct PunchSymToConeHoleClient { - peer_mgr: Arc, - udp_array: RwLock>>, - try_direct_connect: AtomicBool, - punch_predicablely: AtomicBool, - punch_randomly: AtomicBool, - blacklist: Arc>, -} - -impl PunchSymToConeHoleClient { - pub(crate) fn new( - peer_mgr: Arc, - blacklist: Arc>, - ) -> Self { - Self { - peer_mgr, - udp_array: RwLock::new(None), - try_direct_connect: AtomicBool::new(true), - punch_predicablely: AtomicBool::new(true), - punch_randomly: AtomicBool::new(true), - blacklist, - } - } - - async fn prepare_udp_array(&self) -> Result, anyhow::Error> { - let rlocked = self.udp_array.read().await; - if let Some(udp_array) = rlocked.clone() { - return Ok(udp_array); - } - - drop(rlocked); - let mut wlocked = self.udp_array.write().await; - if let Some(udp_array) = wlocked.clone() { - return Ok(udp_array); - } - - let udp_array = Arc::new(UdpSocketArray::new( - UDP_ARRAY_SIZE_FOR_HARD_SYM, - self.peer_mgr.get_global_ctx().net_ns.clone(), - )); - udp_array.start().await?; - wlocked.replace(udp_array.clone()); - Ok(udp_array) - } - - pub(crate) async fn clear_udp_array(&self) { - let mut wlocked = self.udp_array.write().await; - wlocked.take(); - } - - async fn get_base_port_for_easy_sym(&self, my_nat_info: UdpNatType) -> Option { - let global_ctx = self.peer_mgr.get_global_ctx(); - if my_nat_info.is_easy_sym() { - match global_ctx - .get_stun_info_collector() - .get_udp_port_mapping(0) - .await - { - Ok(addr) => Some(addr.port()), - ret => { - tracing::warn!(?ret, "failed to get udp port mapping for easy sym"); - None - } - } - } else { - None - } - } - - async fn remote_send_hole_punch_packet_predicable< - S: UdpHolePunchRpc, - >( - rpc_stub: S, - base_port_for_easy_sym: Option, - my_nat_info: UdpNatType, - remote_mapped_addr: crate::proto::common::SocketAddr, - public_ips: Vec, - tid: u32, - ) { - let Some(inc) = my_nat_info.get_inc_of_easy_sym() else { - return; - }; - let req = SendPunchPacketEasySymRequest { - listener_mapped_addr: remote_mapped_addr.into(), - public_ips: public_ips.clone().into_iter().map(|x| x.into()).collect(), - transaction_id: tid, - base_port_num: base_port_for_easy_sym.unwrap() as u32, - max_port_num: 50, - is_incremental: inc, - }; - tracing::debug!(?req, "send punch packet for easy sym start"); - let ret = rpc_stub - .send_punch_packet_easy_sym( - BaseController { - timeout_ms: 4000, - trace_id: 0, - ..Default::default() - }, - req, - ) - .await; - tracing::debug!(?ret, "send punch packet for easy sym return"); - } - - async fn remote_send_hole_punch_packet_random< - S: UdpHolePunchRpc, - >( - rpc_stub: S, - remote_mapped_addr: crate::proto::common::SocketAddr, - public_ips: Vec, - tid: u32, - round: u32, - port_index: u32, - ) -> Option { - let req = SendPunchPacketHardSymRequest { - listener_mapped_addr: remote_mapped_addr.into(), - public_ips: public_ips.clone().into_iter().map(|x| x.into()).collect(), - transaction_id: tid, - round, - port_index, - }; - tracing::debug!(?req, "send punch packet for hard sym start"); - match rpc_stub - .send_punch_packet_hard_sym( - BaseController { - timeout_ms: 4000, - trace_id: 0, - ..Default::default() - }, - req, - ) - .await - { - Err(e) => { - tracing::error!(?e, "failed to send punch packet for hard sym"); - None - } - Ok(resp) => Some(resp.next_port_index), - } - } - - async fn get_rpc_stub( - &self, - dst_peer_id: PeerId, - ) -> Box + std::marker::Send + Sync + 'static> - { - self.peer_mgr - .get_peer_rpc_mgr() - .rpc_client() - .scoped_client::>( - self.peer_mgr.my_peer_id(), - dst_peer_id, - self.peer_mgr.get_global_ctx().get_network_name(), - ) - } - - async fn check_hole_punch_result( - global_ctx: ArcGlobalCtx, - udp_array: &Arc, - packet: &[u8], - tid: u32, - remote_mapped_addr: crate::proto::common::SocketAddr, - punch_task: &AbortOnDropHandle, - ) -> Result>, anyhow::Error> { - // no matter what the result is, we should check if we received any hole punching packet - let mut ret_tunnel: Option> = None; - let mut finish_time: Option = None; - while finish_time.is_none() || finish_time.as_ref().unwrap().elapsed().as_millis() < 1000 { - udp_array - .send_with_all(packet, remote_mapped_addr.into()) - .await?; - - tokio::time::sleep(Duration::from_millis(200)).await; - - if finish_time.is_none() && punch_task.is_finished() { - finish_time = Some(Instant::now()); - } - - let Some(socket) = udp_array.try_fetch_punched_socket(tid) else { - tracing::debug!("no punched socket found, wait for more time"); - continue; - }; - - // if hole punched but tunnel creation failed, need to retry entire process. - match try_connect_with_socket( - global_ctx.clone(), - socket.socket.clone(), - remote_mapped_addr.into(), - ) - .await - { - Ok(tunnel) => { - ret_tunnel.replace(tunnel); - break; - } - Err(e) => { - tracing::error!(?e, "failed to connect with socket"); - udp_array.add_new_socket(socket.socket).await?; - continue; - } - } - } - - Ok(ret_tunnel) - } - - #[tracing::instrument(err(level = Level::ERROR), skip(self))] - pub(crate) async fn do_hole_punching( - &self, - dst_peer_id: PeerId, - round: u32, - last_port_idx: &mut usize, - my_nat_info: UdpNatType, - ) -> Result>, anyhow::Error> { - // Check if peer is blacklisted - if self.blacklist.contains(&dst_peer_id) { - tracing::debug!(?dst_peer_id, "peer is blacklisted, skipping hole punching"); - return Ok(None); - } - - let udp_array = self.prepare_udp_array().await?; - let global_ctx = self.peer_mgr.get_global_ctx(); - - let rpc_stub = self - .peer_mgr - .get_peer_rpc_mgr() - .rpc_client() - .scoped_client::>( - self.peer_mgr.my_peer_id(), - dst_peer_id, - global_ctx.get_network_name(), - ); - - let resp = rpc_stub - .select_punch_listener( - BaseController::default(), - SelectPunchListenerRequest { - force_new: false, - prefer_port_mapping: true, - }, - ) - .await; - - let resp = handle_rpc_result(resp, dst_peer_id, &self.blacklist)?; - - let remote_mapped_addr = resp.listener_mapped_addr.ok_or(anyhow::anyhow!( - "select_punch_listener response missing listener_mapped_addr" - ))?; - - // try direct connect first - if self.try_direct_connect.load(Ordering::Relaxed) - && let Ok(tunnel) = try_connect_with_socket( - global_ctx.clone(), - Arc::new(UdpSocket::bind("0.0.0.0:0").await?), - remote_mapped_addr.into(), - ) - .await - { - return Ok(Some(tunnel)); - } - - let stun_info = global_ctx.get_stun_info_collector().get_stun_info(); - let public_ips: Vec = stun_info - .public_ip - .iter() - .filter_map(|x| x.parse().ok()) - .collect(); - if public_ips.is_empty() { - return Err(anyhow::anyhow!("failed to get public ips")); - } - - let tid = rand::thread_rng().r#gen(); - let packet = new_hole_punch_packet(tid, HOLE_PUNCH_PACKET_BODY_LEN).into_bytes(); - udp_array.add_intreast_tid(tid); - defer! { udp_array.remove_intreast_tid(tid);} - - let port_index = *last_port_idx as u32; - let base_port_for_easy_sym = self.get_base_port_for_easy_sym(my_nat_info).await; - udp_array - .send_with_all(&packet, remote_mapped_addr.into()) - .await?; - - if self.punch_predicablely.load(Ordering::Relaxed) && base_port_for_easy_sym.is_some() { - let rpc_stub = self.get_rpc_stub(dst_peer_id).await; - let punch_task = AbortOnDropHandle::new(tokio::spawn( - Self::remote_send_hole_punch_packet_predicable( - rpc_stub, - base_port_for_easy_sym, - my_nat_info, - remote_mapped_addr, - public_ips.clone(), - tid, - ), - )); - let ret_tunnel = Self::check_hole_punch_result( - global_ctx.clone(), - &udp_array, - &packet, - tid, - remote_mapped_addr, - &punch_task, - ) - .await?; - - let task_ret = punch_task.await; - tracing::debug!(?ret_tunnel, ?task_ret, "predictable punch task got result"); - if let Some(tunnel) = ret_tunnel { - return Ok(Some(tunnel)); - } - } - - let rpc_stub = self.get_rpc_stub(dst_peer_id).await; - let punch_task = - AbortOnDropHandle::new(tokio::spawn(Self::remote_send_hole_punch_packet_random( - rpc_stub, - remote_mapped_addr, - public_ips.clone(), - tid, - round, - port_index, - ))); - let ret_tunnel = Self::check_hole_punch_result( - global_ctx, - &udp_array, - &packet, - tid, - remote_mapped_addr, - &punch_task, - ) - .await?; - - let punch_task_result = punch_task.await; - tracing::debug!(?punch_task_result, ?ret_tunnel, "punch task got result"); - - if let Ok(Some(next_port_idx)) = punch_task_result { - *last_port_idx = next_port_idx as usize; - } else { - *last_port_idx = rand::random(); - } - - Ok(ret_tunnel) - } -} - -#[cfg(test)] -pub mod tests { - use std::{ - sync::{Arc, atomic::AtomicU32}, - time::Duration, - }; - - use tokio::net::UdpSocket; - - use crate::{ - connector::udp_hole_punch::{ - RUN_TESTING, UdpHolePunchConnector, tests::create_mock_peer_manager_with_mock_stun, - }, - peers::tests::{connect_peer_manager, wait_route_appear, wait_route_appear_with_cost}, - proto::common::NatType, - tunnel::common::tests::wait_for_condition, - }; - - #[tokio::test] - #[serial_test::serial] - #[serial_test::serial(hole_punch)] - async fn hole_punching_symmetric_only_random() { - RUN_TESTING.store(true, std::sync::atomic::Ordering::Relaxed); - - let p_a = create_mock_peer_manager_with_mock_stun(NatType::Symmetric).await; - let p_b = create_mock_peer_manager_with_mock_stun(NatType::PortRestricted).await; - let p_c = create_mock_peer_manager_with_mock_stun(NatType::PortRestricted).await; - connect_peer_manager(p_a.clone(), p_b.clone()).await; - connect_peer_manager(p_b.clone(), p_c.clone()).await; - wait_route_appear(p_a.clone(), p_c.clone()).await.unwrap(); - - let mut hole_punching_a = UdpHolePunchConnector::new(p_a.clone()); - let mut hole_punching_c = UdpHolePunchConnector::new(p_c.clone()); - - hole_punching_a - .client - .data() - .sym_to_cone_client - .try_direct_connect - .store(false, std::sync::atomic::Ordering::Relaxed); - - hole_punching_a - .client - .data() - .sym_to_cone_client - .punch_predicablely - .store(false, std::sync::atomic::Ordering::Relaxed); - - hole_punching_a.run().await.unwrap(); - hole_punching_c.run().await.unwrap(); - - hole_punching_a.client.run_immediately().await; - - wait_for_condition( - || async { - hole_punching_a - .client - .data() - .sym_to_cone_client - .udp_array - .read() - .await - .is_some() - }, - Duration::from_secs(5), - ) - .await; - - println!("start punching {:?}", p_a.list_routes().await); - - wait_for_condition( - || async { - wait_route_appear_with_cost(p_a.clone(), p_c.my_peer_id(), Some(1)) - .await - .is_ok() - }, - Duration::from_secs(60), - ) - .await; - println!("{:?}", p_a.list_routes().await); - - wait_for_condition( - || async { - hole_punching_a - .client - .data() - .sym_to_cone_client - .udp_array - .read() - .await - .is_none() - }, - Duration::from_secs(10), - ) - .await; - } - - #[rstest::rstest] - #[tokio::test] - #[serial_test::serial(hole_punch)] - async fn hole_punching_symmetric_only_predict(#[values("true", "false")] is_inc: bool) { - use tokio_util::task::AbortOnDropHandle; - - RUN_TESTING.store(true, std::sync::atomic::Ordering::Relaxed); - - let p_a = create_mock_peer_manager_with_mock_stun(if is_inc { - NatType::SymmetricEasyInc - } else { - NatType::SymmetricEasyDec - }) - .await; - let p_b = create_mock_peer_manager_with_mock_stun(NatType::PortRestricted).await; - let p_c = create_mock_peer_manager_with_mock_stun(NatType::PortRestricted).await; - connect_peer_manager(p_a.clone(), p_b.clone()).await; - connect_peer_manager(p_b.clone(), p_c.clone()).await; - wait_route_appear(p_a.clone(), p_c.clone()).await.unwrap(); - - let mut hole_punching_a = UdpHolePunchConnector::new(p_a.clone()); - let mut hole_punching_c = UdpHolePunchConnector::new(p_c.clone()); - - hole_punching_a - .client - .data() - .sym_to_cone_client - .try_direct_connect - .store(false, std::sync::atomic::Ordering::Relaxed); - - hole_punching_a - .client - .data() - .sym_to_cone_client - .punch_randomly - .store(false, std::sync::atomic::Ordering::Relaxed); - - hole_punching_a.run().await.unwrap(); - hole_punching_c.run().await.unwrap(); - - let udps = if is_inc { - let udp1 = Arc::new(UdpSocket::bind("0.0.0.0:40147").await.unwrap()); - let udp2 = Arc::new(UdpSocket::bind("0.0.0.0:40194").await.unwrap()); - vec![udp1, udp2] - } else { - let udp1 = Arc::new(UdpSocket::bind("0.0.0.0:40141").await.unwrap()); - let udp2 = Arc::new(UdpSocket::bind("0.0.0.0:40100").await.unwrap()); - vec![udp1, udp2] - }; - // let udp_dec = Arc::new(UdpSocket::bind("0.0.0.0:40140").await.unwrap()); - // let udp_dec2 = Arc::new(UdpSocket::bind("0.0.0.0:40050").await.unwrap()); - - let counter = Arc::new(AtomicU32::new(0)); - - let mut tasks: Vec> = vec![]; - - // all these sockets should receive hole punching packet - for udp in udps.iter().map(Arc::clone) { - let counter = counter.clone(); - tasks.push(AbortOnDropHandle::new(tokio::spawn(async move { - let mut buf = [0u8; 1024]; - let (len, addr) = udp.recv_from(&mut buf).await.unwrap(); - println!( - "got predictable punch packet, {:?} {:?} {:?}", - len, - addr, - udp.local_addr() - ); - counter.fetch_add(1, std::sync::atomic::Ordering::Relaxed); - }))); - } - - hole_punching_a.client.run_immediately().await; - - let udp_len = udps.len(); - wait_for_condition( - || async { counter.load(std::sync::atomic::Ordering::Relaxed) == udp_len as u32 }, - Duration::from_secs(30), - ) - .await; - } -} diff --git a/easytier/src/core.rs b/easytier/src/core.rs index b2752bc9..f4a325bb 100644 --- a/easytier/src/core.rs +++ b/easytier/src/core.rs @@ -1,19 +1,16 @@ -#![allow(dead_code)] - use crate::{ ShellType, common::{ config::{ ConfigFileControl, ConfigLoader, ConsoleLoggerConfig, EncryptionAlgorithm, FileLoggerConfig, LoggingConfigLoader, NetworkIdentity, PeerConfig, PortForwardConfig, - TomlConfigLoader, VpnPortalConfig, load_config_from_file, parse_mapped_listener_urls, - process_secure_mode_cfg, + TomlConfigLoader, VpnPortalConfig, add_proxy_network_to_config, load_config_from_file, + load_toml_config_from_path, parse_mapped_listener_urls, }, constants::EASYTIER_VERSION, log, }, - instance_manager::NetworkInstanceManager, - launcher::add_proxy_network_to_config, + instance::factory::native_cli_instance_manager, proto::common::{CompressionAlgoPb, SecureModeConfig}, rpc_service::ApiRpcServer, utils::panic::setup_panic_handler, @@ -22,6 +19,7 @@ use crate::{ use anyhow::Context; use cidr::IpCidr; use clap::{CommandFactory, Parser}; +use easytier_core::config::normalize_secure_mode_config; use guarden::defer; use rust_i18n::t; use std::{ @@ -49,6 +47,7 @@ fn set_prof_active(_active: bool) { } } +#[cfg(feature = "jemalloc-prof")] fn get_dump_profile_path(cur_allocated: usize, suffix: &str) -> String { format!( "profile-{}-{}.{}", @@ -303,7 +302,7 @@ struct NetworkOptions { long, env = "ET_ENCRYPTION_ALGORITHM", help = t!("core_clap.encryption_algorithm").to_string(), - value_enum, + value_parser = crate::common::config::parse_encryption_algorithm, )] encryption_algorithm: Option, @@ -1075,7 +1074,7 @@ impl NetworkOptions { local_private_key: Some(credential_secret.clone()), local_public_key: None, }; - cfg.set_secure_mode(Some(process_secure_mode_cfg(c)?)); + cfg.set_secure_mode(Some(normalize_secure_mode_config(c)?)); } else if let Some(secure_mode) = self.secure_mode && secure_mode { @@ -1084,7 +1083,7 @@ impl NetworkOptions { local_private_key: self.local_private_key.clone(), local_public_key: self.local_public_key.clone(), }; - cfg.set_secure_mode(Some(process_secure_mode_cfg(c)?)); + cfg.set_secure_mode(Some(normalize_secure_mode_config(c)?)); } let mut f = cfg.get_flags(); @@ -1353,7 +1352,7 @@ async fn run_main(cli: Cli) -> anyhow::Result<()> { defer!(dump_profile(0);); log::init(&cli.logging_options, true)?; - let manager = Arc::new(NetworkInstanceManager::new().with_config_path(cli.config_dir.clone())); + let manager = Arc::new(native_cli_instance_manager().with_config_path(cli.config_dir.clone())); let _rpc_server = ApiRpcServer::new( cli.rpc_portal_options.rpc_portal, @@ -1470,7 +1469,7 @@ async fn run_main(cli: Cli) -> anyhow::Result<()> { control.permission, cfg.dump() ); - manager.run_network_instance(cfg, true, control)?; + manager.run_network_instance(cfg, control)?; } if crate_cli_network { @@ -1487,7 +1486,7 @@ async fn run_main(cli: Cli) -> anyhow::Result<()> { ", cfg.dump() ); - manager.run_network_instance(cfg, true, ConfigFileControl::STATIC_CONFIG)?; + manager.run_network_instance(cfg, ConfigFileControl::STATIC_CONFIG)?; } #[cfg(unix)] @@ -1650,7 +1649,7 @@ async fn validate_config(cli: &Cli) -> anyhow::Result<()> { .context("failed to read config from stdin")?; TomlConfigLoader::new_from_str_with_source("stdin", stdin.as_str())?; } else { - TomlConfigLoader::new(config_file)?; + load_toml_config_from_path(config_file)?; }; } diff --git a/easytier/src/easytier-cli.rs b/easytier/src/easytier-cli.rs index 462e8cdb..fadc1175 100644 --- a/easytier/src/easytier-cli.rs +++ b/easytier/src/easytier-cli.rs @@ -18,6 +18,7 @@ use cidr::Ipv4Inet; use clap::{ArgAction, Args, CommandFactory, Parser, Subcommand, builder::BoolishValueParser}; use dashmap::DashMap; use easytier::ShellType; +use easytier_core::connectivity::stun::StunInfoProvider as _; use humansize::format_size; use rust_i18n::t; use service_manager::*; @@ -29,11 +30,7 @@ use easytier::service_manager::{Service, ServiceInstallOptions}; use tokio::time::timeout; use easytier::{ - common::{ - constants::EASYTIER_VERSION, - stun::{StunInfoCollector, StunInfoCollectorTrait}, - }, - peers, + common::{constants::EASYTIER_VERSION, stun::runtime_stun_info_collector}, proto::{ acl::AclStats, api::{ @@ -73,10 +70,10 @@ use easytier::{ }, common::{NatType, PortForwardConfigPb, SocketType}, peer_rpc::{GetGlobalPeerMapRequest, PeerCenterRpc, PeerCenterRpcClientFactory}, - rpc_impl::standalone::StandAloneClient, + rpc::standalone::{RuntimeRpcClient, runtime_rpc_client}, rpc_types::{controller::BaseController, error::Error as RpcError}, }, - tunnel::{TunnelScheme, tcp::TcpTunnelConnector}, + tunnel::TunnelScheme, utils::{PeerRoutePair, string::cost_to_str}, }; @@ -535,7 +532,7 @@ struct CommandHandler<'a> { resolved_target: Option, } -type RpcClient = StandAloneClient; +type RpcClient = RuntimeRpcClient; type LocalBoxFuture<'a, T> = Pin> + 'a>>; type ForeignNetworkMap = BTreeMap; type GlobalForeignNetworkMap = BTreeMap; @@ -1405,8 +1402,12 @@ impl<'a> CommandHandler<'a> { }; } - let a_is_public = a.hostname.starts_with(peers::PUBLIC_SERVER_HOSTNAME_PREFIX); - let b_is_public = b.hostname.starts_with(peers::PUBLIC_SERVER_HOSTNAME_PREFIX); + let a_is_public = a.hostname.starts_with( + easytier_core::peers::foreign_network::PUBLIC_SERVER_HOSTNAME_PREFIX, + ); + let b_is_public = b.hostname.starts_with( + easytier_core::peers::foreign_network::PUBLIC_SERVER_HOSTNAME_PREFIX, + ); if a_is_public != b_is_public { return if a_is_public { std::cmp::Ordering::Less @@ -2863,11 +2864,11 @@ async fn main() -> Result<(), Error> { rust_i18n::set_locale(&locale); let cli = Cli::parse(); - let client = RpcClient::new(TcpTunnelConnector::new( + let client = runtime_rpc_client( format!("tcp://{}:{}", cli.rpc_portal.ip(), cli.rpc_portal.port()) .parse() .unwrap(), - )); + ); let handler = CommandHandler { client: Arc::new(tokio::sync::Mutex::new(client)), verbose: cli.verbose, @@ -2941,7 +2942,7 @@ async fn main() -> Result<(), Error> { }, SubCommand::Stun => { timeout(Duration::from_secs(25), async move { - let collector = StunInfoCollector::new_with_default_servers(); + let collector = runtime_stun_info_collector(Default::default()); loop { let ret = collector.get_stun_info(); if ret.udp_nat_type != NatType::Unknown as i32 diff --git a/easytier/src/gateway/fast_socks5/util/mod.rs b/easytier/src/gateway/fast_socks5/util/mod.rs deleted file mode 100644 index e1c2f62c..00000000 --- a/easytier/src/gateway/fast_socks5/util/mod.rs +++ /dev/null @@ -1,2 +0,0 @@ -pub mod stream; -pub mod target_addr; diff --git a/easytier/src/gateway/fast_socks5/util/stream.rs b/easytier/src/gateway/fast_socks5/util/stream.rs deleted file mode 100644 index e76b34ab..00000000 --- a/easytier/src/gateway/fast_socks5/util/stream.rs +++ /dev/null @@ -1,65 +0,0 @@ -use std::time::Duration; -use tokio::io::ErrorKind as IOErrorKind; -use tokio::net::{TcpStream, ToSocketAddrs}; -use tokio::time::timeout; - -use crate::gateway::fast_socks5::{ReplyError, Result}; - -/// Easy to destructure bytes buffers by naming each fields: -/// -/// # Examples (before) -/// -/// ```ignore -/// let mut buf = [0u8; 2]; -/// stream.read_exact(&mut buf).await?; -/// let [version, method_len] = buf; -/// -/// assert_eq!(version, 0x05); -/// ``` -/// -/// # Examples (after) -/// -/// ```ignore -/// let [version, method_len] = read_exact!(stream, [0u8; 2]); -/// -/// assert_eq!(version, 0x05); -/// ``` -#[macro_export] -macro_rules! read_exact { - ($stream: expr, $array: expr) => {{ - let mut x = $array; - // $stream - // .read_exact(&mut x) - // .await - // .map_err(|_| io_err("lol"))?; - $stream.read_exact(&mut x).await.map(|_| x) - }}; -} - -pub async fn tcp_connect_with_timeout(addr: T, request_timeout_s: u64) -> Result -where - T: ToSocketAddrs, -{ - let fut = tcp_connect(addr); - match timeout(Duration::from_secs(request_timeout_s), fut).await { - Ok(result) => result, - Err(_) => Err(ReplyError::ConnectionTimeout.into()), - } -} - -pub async fn tcp_connect(addr: T) -> Result -where - T: ToSocketAddrs, -{ - match TcpStream::connect(addr).await { - Ok(o) => Ok(o), - Err(e) => match e.kind() { - // Match other TCP errors with ReplyError - IOErrorKind::ConnectionRefused => Err(ReplyError::ConnectionRefused.into()), - IOErrorKind::ConnectionAborted => Err(ReplyError::ConnectionNotAllowed.into()), - IOErrorKind::ConnectionReset => Err(ReplyError::ConnectionNotAllowed.into()), - IOErrorKind::NotConnected => Err(ReplyError::NetworkUnreachable.into()), - _ => Err(e.into()), // #[error("General failure")] ? - }, - } -} diff --git a/easytier/src/gateway/hedge.rs b/easytier/src/gateway/hedge.rs new file mode 100644 index 00000000..52e9622a --- /dev/null +++ b/easytier/src/gateway/hedge.rs @@ -0,0 +1,95 @@ +use std::{ + fmt::{self, Display}, + future::Future, + time::Duration, +}; + +use futures::{StreamExt, stream::FuturesUnordered}; +use tokio::time::sleep; + +#[derive(Debug)] +pub(super) struct ErrorCollection { + errors: Vec, +} + +impl ErrorCollection { + fn new() -> Self { + Self { errors: Vec::new() } + } + + fn push(&mut self, error: E) { + self.errors.push(error); + } +} + +impl Display for ErrorCollection { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + if self.errors.is_empty() { + return write!(f, "No errors"); + } + + write!(f, "{} error(s) occurred:", self.errors.len())?; + for (i, err) in self.errors.iter().enumerate() { + writeln!(f)?; + write!(f, " {}. {}", i + 1, err)?; + } + + Ok(()) + } +} + +impl std::error::Error for ErrorCollection {} + +pub(super) trait HedgeExt: Iterator + Sized { + async fn hedge(self, delay: Duration) -> Result> + where + Self::Item: Future>; +} + +impl HedgeExt for I +where + I: Iterator, +{ + async fn hedge(mut self, delay: Duration) -> Result> + where + Self::Item: Future>, + { + let mut tasks = FuturesUnordered::new(); + let mut errors = ErrorCollection::new(); + let mut exhausted = false; + + macro_rules! spawn { + () => { + if let Some(fut) = self.next() { + tasks.push(fut); + } else { + exhausted = true; + } + }; + } + + spawn!(); + + while !tasks.is_empty() { + tokio::select! { + res = tasks.next() => { + match res { + Some(Ok(v)) => return Ok(v), + Some(Err(e)) => errors.push(e), + None => unreachable!(), + } + + if !exhausted { + spawn!(); + } + } + + _ = sleep(delay), if !exhausted => { + spawn!(); + } + } + } + + Err(errors) + } +} diff --git a/easytier/src/gateway/icmp_proxy.rs b/easytier/src/gateway/icmp_proxy.rs index cff2f8db..4cc8f01e 100644 --- a/easytier/src/gateway/icmp_proxy.rs +++ b/easytier/src/gateway/icmp_proxy.rs @@ -1,492 +1,91 @@ use std::{ mem::MaybeUninit, net::{IpAddr, Ipv4Addr, SocketAddrV4}, - sync::{Arc, Weak}, - thread, - time::Duration, + sync::Arc, }; -use anyhow::Context; -use pnet::packet::{ - Packet, - icmp::{self, IcmpCode, IcmpTypes, MutableIcmpPacket, echo_reply::MutableEchoReplyPacket}, - ip::IpNextHeaderProtocols, - ipv4::Ipv4Packet, +use easytier_core::{ + gateway::proxy::icmp_host::{IcmpProxyHost, IcmpProxySocket, ProxyRuntimeError}, + socket::SocketContext, }; -use quanta::Instant; use socket2::Socket; -use tokio::{ - sync::{Mutex, mpsc::UnboundedSender}, - task::JoinSet, -}; -use tracing::Instrument; +use crate::common::netns::NetNS; -use crate::{ - common::{PeerId, error::Error, global_ctx::ArcGlobalCtx}, - gateway::ip_reassembler::ComposeIpv4PacketArgs, - peers::{PeerPacketFilter, peer_manager::PeerManager}, - tunnel::packet_def::{PacketType, ZCPacket}, -}; +#[derive(Debug, Default)] +pub(crate) struct RuntimeIcmpProxyHost; -use super::{ - CidrSet, - ip_reassembler::{IpReassembler, compose_ipv4_packet}, -}; - -#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] -struct IcmpNatKey { - real_dst_ip: std::net::IpAddr, - icmp_id: u16, - icmp_seq: u16, -} - -#[derive(Debug)] -struct IcmpNatEntry { - src_peer_id: PeerId, - my_peer_id: PeerId, - src_ip: IpAddr, - start_time: Instant, - mapped_dst_ip: std::net::Ipv4Addr, -} - -impl IcmpNatEntry { - fn new( - src_peer_id: PeerId, - my_peer_id: PeerId, - src_ip: IpAddr, - mapped_dst_ip: Ipv4Addr, - ) -> Result { - Ok(Self { - src_peer_id, - my_peer_id, - src_ip, - start_time: Instant::now(), - mapped_dst_ip, - }) - } -} - -type IcmpNatTable = Arc>; -type NewPacketSender = tokio::sync::mpsc::UnboundedSender; -type NewPacketReceiver = tokio::sync::mpsc::UnboundedReceiver; - -#[derive(Debug)] -pub struct IcmpProxy { - global_ctx: ArcGlobalCtx, - peer_manager: Weak, - - cidr_set: CidrSet, - socket: std::sync::Mutex>>, - - nat_table: IcmpNatTable, - - tasks: Mutex>, - - ip_resemmbler: Arc, - icmp_sender: Arc>>>, -} - -fn socket_recv( - socket: &Socket, - buf: &mut [MaybeUninit], -) -> Result<(usize, IpAddr), std::io::Error> { - let (size, addr) = socket.recv_from(buf)?; - let addr = match addr.as_socket() { - None => IpAddr::V4(Ipv4Addr::UNSPECIFIED), - Some(add) => add.ip(), - }; - Ok((size, addr)) -} - -fn socket_recv_loop( - socket: Arc, - nat_table: IcmpNatTable, - sender: UnboundedSender, -) { - let mut buf = [0u8; 8192]; - let data: &mut [MaybeUninit] = unsafe { std::mem::transmute(&mut buf[..]) }; - - loop { - let (len, peer_ip) = match socket_recv(&socket, data) { - Ok((len, peer_ip)) => (len, peer_ip), - Err(e) => { - tracing::error!("recv icmp packet failed: {:?}", e); - if sender.is_closed() { - break; - } else { - continue; - } - } - }; - - if len == 0 { - tracing::error!("recv empty packet, len: {}", len); - return; - } - - if !peer_ip.is_ipv4() { - continue; - } - - let Some(ipv4_packet) = Ipv4Packet::new(&buf[..len]) else { - continue; - }; - - let Some(icmp_packet) = icmp::echo_reply::EchoReplyPacket::new(ipv4_packet.payload()) - else { - continue; - }; - - if icmp_packet.get_icmp_type() != IcmpTypes::EchoReply { - continue; - } - - let key = IcmpNatKey { - real_dst_ip: peer_ip, - icmp_id: icmp_packet.get_identifier(), - icmp_seq: icmp_packet.get_sequence_number(), - }; - - let Some((_, v)) = nat_table.remove(&key) else { - continue; - }; - - // send packet back to the peer where this request origin. - let IpAddr::V4(dest_ip) = v.src_ip else { - continue; - }; - - let payload_len = len - ipv4_packet.get_header_length() as usize * 4; - let id = ipv4_packet.get_identification(); - let _ = compose_ipv4_packet( - ComposeIpv4PacketArgs { - buf: &mut buf[..], - src_v4: &v.mapped_dst_ip, - dst_v4: &dest_ip, - next_protocol: IpNextHeaderProtocols::Icmp, - payload_len, - payload_mtu: 1200, - ip_id: id, - }, - |buf| { - let mut p = ZCPacket::new_with_payload(buf); - p.fill_peer_manager_hdr(v.my_peer_id, v.src_peer_id, PacketType::Data as u8); - p.mut_peer_manager_header().unwrap().set_no_proxy(true); - - if let Err(e) = sender.send(p) { - tracing::error!("send icmp packet to peer failed: {:?}, may exiting..", e); - } - Ok(()) - }, - ); - } -} - -#[async_trait::async_trait] -impl PeerPacketFilter for IcmpProxy { - async fn try_process_packet_from_peer(&self, packet: ZCPacket) -> Option { - if self.try_handle_peer_packet(&packet).await.is_some() { - return None; - } else { - return Some(packet); - } - } -} - -impl IcmpProxy { - pub fn new( - global_ctx: ArcGlobalCtx, - peer_manager: Arc, - ) -> Result, Error> { - let cidr_set = CidrSet::new(global_ctx.clone()); - let ret = Self { - global_ctx, - peer_manager: Arc::downgrade(&peer_manager), - cidr_set, - socket: std::sync::Mutex::new(None), - - nat_table: Arc::new(dashmap::DashMap::new()), - tasks: Mutex::new(JoinSet::new()), - - ip_resemmbler: Arc::new(IpReassembler::new(Duration::from_secs(10))), - icmp_sender: Arc::new(std::sync::Mutex::new(None)), - }; - - Ok(Arc::new(ret)) - } - - fn create_raw_socket(self: &Arc) -> Result { - let _g = self.global_ctx.net_ns.guard(); - let socket = socket2::Socket::new( +impl RuntimeIcmpProxyHost { + fn create_raw_socket(context: &SocketContext) -> Result { + let _guard = NetNS::from_socket_context(context).guard(); + let socket = Socket::new( socket2::Domain::IPV4, socket2::Type::RAW, Some(socket2::Protocol::ICMPV4), )?; socket.bind(&socket2::SockAddr::from(SocketAddrV4::new( - std::net::Ipv4Addr::UNSPECIFIED, + Ipv4Addr::UNSPECIFIED, 0, )))?; Ok(socket) } +} - pub async fn start(self: &Arc) -> Result<(), Error> { - let socket = self.create_raw_socket(); - match socket { - Ok(socket) => { - self.socket.lock().unwrap().replace(Arc::new(socket)); - } - Err(e) => { - tracing::warn!("create icmp socket failed: {:?}", e); - if !self.global_ctx.no_tun() { - return Err(anyhow::anyhow!("create icmp socket failed: {:?}", e).into()); - } - } - } +#[derive(Debug)] +struct RuntimeIcmpSocket { + socket: Arc, +} - self.start_icmp_proxy().await?; - self.start_nat_table_cleaner().await?; - Ok(()) - } - - async fn start_nat_table_cleaner(self: &Arc) -> Result<(), Error> { - let nat_table = self.nat_table.clone(); - self.tasks.lock().await.spawn( - async move { - loop { - tokio::time::sleep(std::time::Duration::from_secs(1)).await; - nat_table.retain(|_, v| v.start_time.elapsed().as_secs() < 20); - nat_table.shrink_to_fit(); - } - } - .instrument(tracing::info_span!("icmp proxy nat table cleaner")), - ); - Ok(()) - } - - async fn start_icmp_proxy(self: &Arc) -> Result<(), Error> { - let (sender, mut receiver) = tokio::sync::mpsc::unbounded_channel(); - self.icmp_sender.lock().unwrap().replace(sender.clone()); - if let Some(socket) = self.socket.lock().unwrap().as_ref() { - let socket = socket.clone(); - let nat_table = self.nat_table.clone(); - thread::spawn(|| { - socket_recv_loop(socket, nat_table, sender); - }); - } - - let peer_manager = self.peer_manager.clone(); - let is_latency_first = self.global_ctx.latency_first(); - self.tasks.lock().await.spawn( - async move { - while let Some(mut msg) = receiver.recv().await { - let hdr = msg.mut_peer_manager_header().unwrap(); - hdr.set_latency_first(is_latency_first); - let to_peer_id = hdr.to_peer_id.into(); - let Some(pm) = peer_manager.upgrade() else { - tracing::warn!("peer manager is gone, icmp proxy send loop exit"); - return; - }; - let ret = pm.send_msg_for_proxy(msg, to_peer_id).await; - if ret.is_err() { - tracing::error!("send icmp packet to peer failed: {:?}", ret); - } - } - } - .instrument(tracing::info_span!("icmp proxy send loop")), - ); - - let ip_resembler = self.ip_resemmbler.clone(); - self.tasks.lock().await.spawn(async move { - loop { - tokio::time::sleep(Duration::from_secs(1)).await; - ip_resembler.remove_expired_packets(); - } - }); - - let Some(pm) = self.peer_manager.upgrade() else { - tracing::warn!("peer manager is gone, icmp proxy init failed"); - return Err(anyhow::anyhow!("peer manager is gone").into()); - }; - - pm.add_packet_process_pipeline(Box::new(self.clone())).await; - Ok(()) - } - - fn send_icmp_packet( - &self, - dst_ip: Ipv4Addr, - icmp_packet: &icmp::echo_request::EchoRequestPacket, - ) -> Result<(), Error> { +#[async_trait::async_trait] +impl IcmpProxySocket for RuntimeIcmpSocket { + async fn send(&self, destination: Ipv4Addr, packet: &[u8]) -> Result<(), ProxyRuntimeError> { self.socket - .lock() - .unwrap() - .as_ref() - .with_context(|| "icmp socket not created")? - .send_to(icmp_packet.packet(), &SocketAddrV4::new(dst_ip, 0).into())?; - + .send_to(packet, &SocketAddrV4::new(destination, 0).into())?; Ok(()) } - async fn send_icmp_reply_to_peer( + async fn recv(&self) -> Result<(IpAddr, Vec), ProxyRuntimeError> { + let socket = self.socket.clone(); + tokio::task::spawn_blocking(move || { + let mut buffer = vec![0_u8; 8192]; + let uninitialized: &mut [MaybeUninit] = + unsafe { std::mem::transmute(&mut buffer[..]) }; + let (length, peer_ip) = socket_recv(&socket, uninitialized)?; + buffer.truncate(length); + Ok((peer_ip, buffer)) + }) + .await + .map_err(|error| ProxyRuntimeError::Other(error.into()))? + } + + fn close(&self) { + let _ = self.socket.shutdown(std::net::Shutdown::Both); + } +} + +#[async_trait::async_trait] +impl IcmpProxyHost for RuntimeIcmpProxyHost { + async fn open_icmp_v4( &self, - src_ip: &Ipv4Addr, - dst_ip: &Ipv4Addr, - src_peer_id: PeerId, - dst_peer_id: PeerId, - icmp_packet: &icmp::echo_request::EchoRequestPacket<'_>, - ) { - let mut buf = vec![0u8; icmp_packet.packet().len() + 20]; - let mut reply_packet = MutableEchoReplyPacket::new(&mut buf[20..]).unwrap(); - reply_packet.set_icmp_type(IcmpTypes::EchoReply); - reply_packet.set_icmp_code(IcmpCode::new(0)); - reply_packet.set_identifier(icmp_packet.get_identifier()); - reply_packet.set_sequence_number(icmp_packet.get_sequence_number()); - reply_packet.set_payload(icmp_packet.payload()); - - let mut icmp_packet = MutableIcmpPacket::new(&mut buf[20..]).unwrap(); - icmp_packet.set_checksum(icmp::checksum(&icmp_packet.to_immutable())); - - let len = buf.len() - 20; - let _ = compose_ipv4_packet( - ComposeIpv4PacketArgs { - buf: &mut buf[..], - src_v4: src_ip, - dst_v4: dst_ip, - next_protocol: IpNextHeaderProtocols::Icmp, - payload_len: len, - payload_mtu: 1200, - ip_id: rand::random(), - }, - |buf| { - let mut packet = ZCPacket::new_with_payload(buf); - packet.fill_peer_manager_hdr(src_peer_id, dst_peer_id, PacketType::Data as u8); - let _ = self - .icmp_sender - .lock() - .unwrap() - .as_ref() - .unwrap() - .send(packet); - Ok(()) - }, - ); - } - - async fn try_handle_peer_packet(&self, packet: &ZCPacket) -> Option<()> { - if self.cidr_set.is_empty() - && !self.global_ctx.enable_exit_node() - && !self.global_ctx.no_tun() - { - return None; - } - - let _ = self.global_ctx.get_ipv4()?; - let hdr = packet.peer_manager_header().unwrap(); - let is_exit_node = hdr.is_exit_node(); - - if hdr.packet_type != PacketType::Data as u8 || hdr.is_no_proxy() { - return None; - }; - - let ipv4 = Ipv4Packet::new(packet.payload())?; - - if ipv4.get_version() != 4 || ipv4.get_next_level_protocol() != IpNextHeaderProtocols::Icmp - { - return None; - } - - let mut real_dst_ip = ipv4.get_destination(); - - if !(self - .cidr_set - .contains_v4(ipv4.get_destination(), &mut real_dst_ip) - || is_exit_node - || (self.global_ctx.no_tun() - && Some(ipv4.get_destination()) - == self - .global_ctx - .get_ipv4() - .as_ref() - .map(cidr::Ipv4Inet::address))) - { - return None; - } - - let resembled_buf: Option>; - let icmp_packet = if IpReassembler::is_packet_fragmented(&ipv4) { - resembled_buf = - self.ip_resemmbler - .add_fragment(ipv4.get_source(), ipv4.get_destination(), &ipv4); - resembled_buf.as_ref()?; - icmp::echo_request::EchoRequestPacket::new(resembled_buf.as_ref().unwrap())? - } else { - icmp::echo_request::EchoRequestPacket::new(ipv4.payload())? - }; - - if icmp_packet.get_icmp_type() != IcmpTypes::EchoRequest { - // if it's other icmp type, just ignore it. may forwarding network to network replay packet. - tracing::trace!("unsupported icmp type: {:?}", icmp_packet.get_icmp_type()); - return None; - } - - if self.global_ctx.no_tun() - && Some(ipv4.get_destination()) - == self - .global_ctx - .get_ipv4() - .as_ref() - .map(cidr::Ipv4Inet::address) - { - self.send_icmp_reply_to_peer( - &ipv4.get_destination(), - &ipv4.get_source(), - hdr.to_peer_id.get(), - hdr.from_peer_id.get(), - &icmp_packet, - ) - .await; - return Some(()); - } - - let icmp_id = icmp_packet.get_identifier(); - let icmp_seq = icmp_packet.get_sequence_number(); - - let key = IcmpNatKey { - real_dst_ip: real_dst_ip.into(), - icmp_id, - icmp_seq, - }; - - let value = IcmpNatEntry::new( - hdr.from_peer_id.into(), - hdr.to_peer_id.into(), - ipv4.get_source().into(), - ipv4.get_destination(), - ) - .ok()?; - - if let Some(old) = self.nat_table.insert(key, value) { - tracing::info!("icmp nat table entry replaced: {:?}", old); - } - - if let Err(e) = self.send_icmp_packet(real_dst_ip, &icmp_packet) { - tracing::error!("send icmp packet failed: {:?}", e); - } - - Some(()) + context: SocketContext, + ) -> Result, ProxyRuntimeError> { + let socket = Self::create_raw_socket(&context).inspect_err(|error| { + tracing::warn!(?error, "create ICMP socket failed"); + })?; + Ok(Arc::new(RuntimeIcmpSocket { + socket: Arc::new(socket), + })) } } -impl Drop for IcmpProxy { - fn drop(&mut self) { - tracing::info!( - "dropping icmp proxy, {:?}", - self.socket.lock().unwrap().as_ref() - ); - if let Some(s) = self.socket.lock().unwrap().as_ref() { - tracing::info!("shutting down icmp socket"); - let _ = s.shutdown(std::net::Shutdown::Both); - } - } +fn socket_recv( + socket: &Socket, + buffer: &mut [MaybeUninit], +) -> Result<(usize, IpAddr), std::io::Error> { + let (size, address) = socket.recv_from(buffer)?; + let peer_ip = address + .as_socket() + .map(|address| address.ip()) + .unwrap_or(IpAddr::V4(Ipv4Addr::UNSPECIFIED)); + Ok((size, peer_ip)) } diff --git a/easytier/src/gateway/ip_reassembler.rs b/easytier/src/gateway/ip_reassembler.rs deleted file mode 100644 index b45481a0..00000000 --- a/easytier/src/gateway/ip_reassembler.rs +++ /dev/null @@ -1,325 +0,0 @@ -use dashmap::DashMap; -use pnet::packet::Packet; -use pnet::packet::ip::IpNextHeaderProtocol; -use pnet::packet::ipv4::{self, Ipv4Flags, Ipv4Packet, MutableIpv4Packet}; -use quanta::Instant; -use std::net::Ipv4Addr; -use std::time::Duration; - -use crate::common::error::Error; - -#[derive(Debug, Clone)] -pub(crate) struct IpFragment { - id: u16, - offset: u16, - data: Vec, -} - -impl<'a> From<&Ipv4Packet<'a>> for IpFragment { - fn from(packet: &Ipv4Packet<'a>) -> Self { - let id = packet.get_identification(); - let offset = packet.get_fragment_offset() * 8; - let data = packet.payload().to_vec(); - IpFragment { id, offset, data } - } -} - -#[derive(Debug, Clone)] -struct IpPacket { - source: Ipv4Addr, - destination: Ipv4Addr, - total_length: Option, - fragments: Vec, -} - -impl IpPacket { - fn new(source: Ipv4Addr, destination: Ipv4Addr) -> Self { - IpPacket { - source, - destination, - total_length: None, - fragments: Vec::new(), - } - } - - fn add_fragment(&mut self, fragment: IpFragment) { - // make sure the fragment doesn't overlap with existing fragments - for f in &self.fragments { - if f.offset <= fragment.offset && fragment.offset < f.offset + f.data.len() as u16 { - tracing::trace!( - "fragment overlap 1, f.offset = {}, fragment.offset = {}, f.data.len() = {}, fragment.data.len() = {}", - f.offset, - fragment.offset, - f.data.len(), - fragment.data.len() - ); - return; - } - if fragment.offset <= f.offset - && f.offset < fragment.offset + fragment.data.len() as u16 - { - tracing::trace!( - "fragment overlap 2, f.offset = {}, fragment.offset = {}, f.data.len() = {}, fragment.data.len() = {}", - f.offset, - fragment.offset, - f.data.len(), - fragment.data.len() - ); - return; - } - } - self.fragments.push(fragment); - } - - fn is_complete(&self) -> bool { - if self.total_length.is_none() { - return false; - } - let mut total_length = 0; - for fragment in &self.fragments { - total_length += fragment.data.len() as u16; - } - tracing::trace!(?total_length, ?self.total_length, "ip resembler checking is_complete"); - Some(total_length) == self.total_length - } - - fn set_total_length(&mut self, total_length: u16) { - self.total_length = Some(total_length); - } - - fn assemble(&mut self) -> Option> { - if !self.is_complete() { - return None; - } - - // sort fragments by offset - self.fragments.sort_by_key(|f| f.offset); - - let mut packet = vec![0u8; self.total_length.unwrap() as usize]; - for fragment in &self.fragments { - let start = fragment.offset as usize; - let end = start + fragment.data.len(); - packet[start..end].copy_from_slice(&fragment.data); - } - - Some(packet) - } -} - -#[derive(Hash, Eq, PartialEq, Clone, Debug)] -struct IpResemblerKey { - source: Ipv4Addr, - destination: Ipv4Addr, - id: u16, -} - -#[derive(Debug)] -struct IpResemblerValue { - packet: IpPacket, - timestamp: Instant, -} - -#[derive(Debug)] -pub(crate) struct IpReassembler { - packets: DashMap, - timeout: Duration, -} - -impl IpReassembler { - pub fn new(timeout: Duration) -> Self { - IpReassembler { - packets: DashMap::new(), - timeout, - } - } - - pub fn is_packet_fragmented(packet: &Ipv4Packet) -> bool { - packet.get_fragment_offset() != 0 || packet.get_flags() & Ipv4Flags::MoreFragments != 0 - } - - pub fn is_last_fragment(packet: &Ipv4Packet) -> bool { - packet.get_flags() & Ipv4Flags::MoreFragments == 0 - } - - pub fn add_fragment( - &self, - source: Ipv4Addr, - destination: Ipv4Addr, - packet: &Ipv4Packet, - ) -> Option> { - let id = packet.get_identification(); - let total_length = packet.get_total_length() - packet.get_header_length() as u16 * 4; - if total_length != packet.payload().len() as u16 { - tracing::trace!( - ?packet, - ?total_length, - payload_len = ?packet.payload().len(), - "unexpected total length", - ); - return None; - } - - let fragment: IpFragment = packet.into(); - let key = IpResemblerKey { - source, - destination, - id, - }; - - tracing::trace!( - ?key, - "add fragment, offset = {}, total_length = {}", - fragment.offset, - total_length - ); - - let mut entry = self.packets.entry(key.clone()).or_insert_with(|| { - let packet = IpPacket::new(source, destination); - let timestamp = Instant::now(); - IpResemblerValue { packet, timestamp } - }); - let value_mut = entry.value_mut(); - - if Self::is_last_fragment(packet) { - value_mut - .packet - .set_total_length(total_length + fragment.offset); - } - - value_mut.packet.add_fragment(fragment); - if let Some(data) = value_mut.packet.assemble() { - drop(entry); - self.packets.remove(&key); - Some(data) - } else { - value_mut.timestamp = Instant::now(); - None - } - } - - pub fn remove_expired_packets(&self) { - let timeout = self.timeout; - self.packets.retain(|_, v| v.timestamp.elapsed() <= timeout); - self.packets.shrink_to_fit(); - } -} - -pub struct ComposeIpv4PacketArgs<'a> { - pub buf: &'a mut [u8], - pub src_v4: &'a Ipv4Addr, - pub dst_v4: &'a Ipv4Addr, - pub next_protocol: IpNextHeaderProtocol, - pub payload_len: usize, - pub payload_mtu: usize, - pub ip_id: u16, -} - -// ip payload should be in buf[20..] -pub fn compose_ipv4_packet(args: ComposeIpv4PacketArgs, cb: F) -> Result<(), Error> -where - F: Fn(&[u8]) -> Result<(), Error>, -{ - let total_pieces = args.payload_len.div_ceil(args.payload_mtu); - let mut buf_offset = 0; - let mut fragment_offset = 0; - let mut cur_piece = 0; - while fragment_offset < args.payload_len { - let next_fragment_offset = - std::cmp::min(fragment_offset + args.payload_mtu, args.payload_len); - let fragment_len = next_fragment_offset - fragment_offset; - let mut ipv4_packet = - MutableIpv4Packet::new(&mut args.buf[buf_offset..buf_offset + fragment_len + 20]) - .unwrap(); - ipv4_packet.set_version(4); - ipv4_packet.set_header_length(5); - ipv4_packet.set_total_length((fragment_len + 20) as u16); - ipv4_packet.set_identification(args.ip_id); - if total_pieces > 1 { - if cur_piece != total_pieces - 1 { - ipv4_packet.set_flags(Ipv4Flags::MoreFragments); - } else { - ipv4_packet.set_flags(0); - } - assert_eq!(0, fragment_offset % 8); - ipv4_packet.set_fragment_offset(fragment_offset as u16 / 8); - } else { - ipv4_packet.set_flags(Ipv4Flags::DontFragment); - ipv4_packet.set_fragment_offset(0); - } - ipv4_packet.set_ecn(0); - ipv4_packet.set_dscp(0); - ipv4_packet.set_ttl(32); - ipv4_packet.set_source(*args.src_v4); - ipv4_packet.set_destination(*args.dst_v4); - ipv4_packet.set_next_level_protocol(args.next_protocol); - ipv4_packet.set_checksum(ipv4::checksum(&ipv4_packet.to_immutable())); - - tracing::trace!(?ipv4_packet, "udp nat packet response send"); - - cb(ipv4_packet.packet())?; - - buf_offset += next_fragment_offset - fragment_offset; - fragment_offset = next_fragment_offset; - cur_piece += 1; - } - Ok(()) -} - -#[cfg(test)] -mod tests { - use super::*; - - #[test] - fn resembler() { - let raw_packets = [ - // last packet - vec![ - 0x45, 0x00, 0x00, 0x1c, 0x1c, 0x46, 0x20, 0x01, 0x40, 0x06, 0xb1, 0xe6, 0xc0, 0xa8, - 0x00, 0x01, 0xc0, 0xa8, 0x00, 0x02, 0x04, 0x05, 0x06, 0x07, 0x04, 0x05, 0x06, 0x07, - ], - // 1st packet - vec![ - 0x45, 0x00, 0x00, 0x1c, 0x1c, 0x46, 0x00, 0x02, 0x40, 0x06, 0xb1, 0xe6, 0xc0, 0xa8, - 0x00, 0x01, 0xc0, 0xa8, 0x00, 0x02, 0x08, 0x09, 0x0a, 0x0b, 0x04, 0x05, 0x06, 0x07, - ], - // 2nd packet - vec![ - 0x45, 0x00, 0x00, 0x1c, 0x1c, 0x46, 0x20, 0x00, 0x40, 0x06, 0xb1, 0xe6, 0xc0, 0xa8, - 0x00, 0x01, 0xc0, 0xa8, 0x00, 0x02, 0x00, 0x01, 0x02, 0x03, 0x04, 0x05, 0x06, 0x07, - ], - // expired packet - vec![ - 0x45, 0x00, 0x00, 0x1c, 0x1c, 0x47, 0x20, 0x00, 0x40, 0x06, 0xb1, 0xe6, 0xc0, 0xa8, - 0x00, 0x01, 0xc0, 0xa8, 0x00, 0x02, 0x00, 0x01, 0x02, 0x03, 0x04, 0x05, 0x06, 0x07, - ], - ]; - - let source = "192.168.0.1".parse().unwrap(); - let destination = "192.168.0.2".parse().unwrap(); - let resembler = IpReassembler::new(Duration::from_secs(1)); - - for (idx, raw_packet) in raw_packets.iter().enumerate() { - if let Some(packet) = Ipv4Packet::new(raw_packet) { - let ret = resembler.add_fragment(source, destination, &packet); - if idx != 2 { - assert!(ret.is_none()); - } else { - assert!(ret.is_some()); - } - println!( - "packet: {:?}, ret: {:?}, palyload_len: {}", - packet, - ret, - packet.payload().len() - ); - } - } - - resembler.remove_expired_packets(); - assert_eq!(1, resembler.packets.len()); - - std::thread::sleep(Duration::from_secs(2)); - resembler.remove_expired_packets(); - assert_eq!(0, resembler.packets.len()); - } -} diff --git a/easytier/src/gateway/kcp_proxy.rs b/easytier/src/gateway/kcp_proxy.rs index b1dceacc..e54e7777 100644 --- a/easytier/src/gateway/kcp_proxy.rs +++ b/easytier/src/gateway/kcp_proxy.rs @@ -1,47 +1,32 @@ -use std::{ - net::{IpAddr, Ipv4Addr, SocketAddr}, - sync::{Arc, Weak}, - time::Duration, -}; +use std::{net::SocketAddr, sync::Arc, time::Duration}; -use anyhow::{Context, anyhow, bail}; +use anyhow::{Context, anyhow}; use bytes::Bytes; -use dashmap::DashMap; -use guarden::defer; use kcp_sys::{ - endpoint::{ConnId, KcpEndpoint, KcpPacketReceiver}, + endpoint::{KcpEndpoint, KcpPacketReceiver}, ffi_safe::KcpConfig, packet_def::KcpPacket, stream::KcpStream, }; use prost::Message; -use tokio::task::JoinSet; +use tokio::{ + sync::Mutex, + task::{JoinHandle, JoinSet}, +}; +use tokio_util::sync::CancellationToken; -use super::{ - CidrSet, - tcp_proxy::{NatDstConnector, NatDstTcpConnector, TcpProxy}, -}; -use crate::utils::task::HedgeExt; -use crate::{ - common::{ - acl_processor::PacketInfo, - error::Result, - global_ctx::{ArcGlobalCtx, GlobalCtx}, +use easytier_core::{ + gateway::proxy::traits::TcpProxyStream, + gateway::proxy::wrapped_transport::{ + WrappedTransportAcceptedStream, WrappedTransportConnect, WrappedTransportDatagram, + WrappedTransportDatagramBuffer, WrappedTransportDestinationIngress, WrappedTransportEngine, + WrappedTransportEngineStart, WrappedTransportKind, WrappedTransportRole, }, - gateway::wrapped_proxy::{ProxyAclHandler, TcpProxyForWrappedSrcTrait}, - peers::{PeerPacketFilter, peer_manager::PeerManager}, - proto::{ - acl::{ChainType, Protocol}, - api::instance::{ - ListTcpProxyEntryRequest, ListTcpProxyEntryResponse, TcpProxyEntry, TcpProxyEntryState, - TcpProxyEntryTransportType, TcpProxyRpc, - }, - peer_rpc::KcpConnData, - rpc_types::{self, controller::BaseController}, - }, - tunnel::packet_def::{PacketType, PeerManagerHeader, ZCPacket}, }; +use super::hedge::HedgeExt; +use crate::proto::peer_rpc::KcpConnData; + fn create_kcp_endpoint() -> KcpEndpoint { let mut kcp_endpoint = KcpEndpoint::new(); kcp_endpoint.set_kcp_config_factory(Box::new(|conv| { @@ -52,177 +37,72 @@ fn create_kcp_endpoint() -> KcpEndpoint { kcp_endpoint } -struct KcpEndpointFilter { - kcp_endpoint: Arc, - is_src: bool, -} - -#[async_trait::async_trait] -impl PeerPacketFilter for KcpEndpointFilter { - async fn try_process_packet_from_peer(&self, packet: ZCPacket) -> Option { - let t = packet.peer_manager_header().unwrap().packet_type; - if t == PacketType::KcpSrc as u8 && !self.is_src { - // src packet, but we are dst - } else if t == PacketType::KcpDst as u8 && self.is_src { - // dst packet, but we are src - } else { - return Some(packet); - } - - let _ = self - .kcp_endpoint - .input_sender_ref() - .send(KcpPacket::from(packet.payload_bytes())) - .await; - - None - } -} - #[tracing::instrument] async fn handle_kcp_output( - peer_mgr: Arc, mut output_receiver: KcpPacketReceiver, - is_src: bool, + role: WrappedTransportRole, + datagrams: tokio::sync::mpsc::Sender, ) { while let Some(packet) = output_receiver.recv().await { - let dst_peer_id = if is_src { - packet.header().dst_session_id() - } else { - packet.header().src_session_id() + let peer_id = match role { + WrappedTransportRole::Source => packet.header().dst_session_id(), + WrappedTransportRole::Destination => packet.header().src_session_id(), }; - let packet_type = if is_src { - PacketType::KcpSrc as u8 - } else { - PacketType::KcpDst as u8 - }; - let mut packet = ZCPacket::new_with_payload(&packet.inner().freeze()); - packet.fill_peer_manager_hdr(peer_mgr.my_peer_id(), dst_peer_id, packet_type); - - if let Err(e) = peer_mgr.send_msg_for_proxy(packet, dst_peer_id).await { - tracing::error!("failed to send kcp packet to peer: {:?}", e); + if datagrams + .send(WrappedTransportDatagram { + transport: WrappedTransportKind::Kcp, + role, + peer_id, + buffer: WrappedTransportDatagramBuffer::copy_from_payload(&packet.inner().freeze()), + }) + .await + .is_err() + { + break; } } } -#[derive(Debug, Clone)] -pub struct NatDstKcpConnector { - pub(crate) kcp_endpoint: Arc, - pub(crate) peer_mgr: Weak, -} +async fn connect_kcp_source( + kcp_endpoint: Arc, + my_peer_id: u32, + dst_peer_id: u32, + src: SocketAddr, + dst: SocketAddr, +) -> anyhow::Result { + let conn_data = KcpConnData { + src: Some(src.into()), + dst: Some(dst.into()), + }; -#[async_trait::async_trait] -impl NatDstConnector for NatDstKcpConnector { - type DstStream = KcpStream; + (0..5) + .map(|_| { + let kcp_endpoint = kcp_endpoint.clone(); + async move { + let conn_id = kcp_endpoint + .connect( + Duration::from_secs(10), + my_peer_id, + dst_peer_id, + Bytes::from(conn_data.encode_to_vec()), + ) + .await?; - async fn connect( - &self, - src: SocketAddr, - nat_dst: SocketAddr, - ) -> anyhow::Result { - let peer_mgr = self - .peer_mgr - .upgrade() - .ok_or_else(|| anyhow!("peer manager is not available"))?; - - let dst_peer = { - let SocketAddr::V4(addr) = nat_dst else { - bail!("ipv6 is not supported"); - }; - peer_mgr - .get_peer_map() - .get_peer_id_by_ipv4(addr.ip()) - .await - .ok_or_else(|| anyhow!("no peer found for nat dst: {}", nat_dst))? - }; - - tracing::trace!(?nat_dst, ?dst_peer, "kcp nat"); - - let conn_data = KcpConnData { - src: Some(src.into()), - dst: Some(nat_dst.into()), - }; - - let stream = (0..5) - .map(|_| { - let kcp_endpoint = self.kcp_endpoint.clone(); - let my_peer_id = peer_mgr.my_peer_id(); - - async move { - let conn_id = kcp_endpoint - .connect( - Duration::from_secs(10), - my_peer_id, - dst_peer, - Bytes::from(conn_data.encode_to_vec()), - ) - .await?; - - KcpStream::new(&kcp_endpoint, conn_id).context("failed to create kcp stream") - } - }) - .hedge(Duration::from_millis(200)) - .await - .context("failed to connect to peer")?; - - Ok(stream) - } - - fn check_packet_from_peer_fast(&self, _cidr_set: &CidrSet, _global_ctx: &GlobalCtx) -> bool { - true - } - - fn check_packet_from_peer( - &self, - _cidr_set: &CidrSet, - _global_ctx: &GlobalCtx, - hdr: &PeerManagerHeader, - _ipv4: &Ipv4Addr, - _real_dst_ip: &mut Ipv4Addr, - ) -> bool { - hdr.from_peer_id == hdr.to_peer_id && hdr.is_kcp_src_modified() - } - - fn transport_type(&self) -> TcpProxyEntryTransportType { - TcpProxyEntryTransportType::Kcp - } -} - -#[derive(Clone)] -struct TcpProxyForKcpSrc(Arc>); - -#[async_trait::async_trait] -impl TcpProxyForWrappedSrcTrait for TcpProxyForKcpSrc { - type Connector = NatDstKcpConnector; - - fn get_tcp_proxy(&self) -> &Arc> { - &self.0 - } - - fn mark_src_modified(hdr: &mut PeerManagerHeader) -> &mut PeerManagerHeader { - hdr.mark_kcp_src_modified() - } - - async fn check_dst_allow_wrapped_input(&self, dst_ip: &Ipv4Addr) -> bool { - let Some(peer_manager) = self.0.get_peer_manager() else { - return false; - }; - peer_manager - .check_allow_kcp_to_dst(&IpAddr::V4(*dst_ip)) - .await - } + KcpStream::new(&kcp_endpoint, conn_id).context("failed to create kcp stream") + } + }) + .hedge(Duration::from_millis(200)) + .await + .context("failed to connect to peer") } pub struct KcpProxySrc { kcp_endpoint: Arc, - peer_manager: Arc, - - tcp_proxy: TcpProxyForKcpSrc, tasks: JoinSet<()>, } impl KcpProxySrc { - pub async fn new(peer_manager: Arc) -> Self { + pub async fn new(datagrams: tokio::sync::mpsc::Sender) -> Self { let mut kcp_endpoint = create_kcp_endpoint(); kcp_endpoint.run().await; @@ -230,96 +110,69 @@ impl KcpProxySrc { let mut tasks = JoinSet::new(); tasks.spawn(handle_kcp_output( - peer_manager.clone(), output_receiver, - true, + WrappedTransportRole::Source, + datagrams, )); let kcp_endpoint = Arc::new(kcp_endpoint); - let tcp_proxy = TcpProxy::new( - peer_manager.clone(), - NatDstKcpConnector { - kcp_endpoint: kcp_endpoint.clone(), - peer_mgr: Arc::downgrade(&peer_manager), - }, - ); - Self { kcp_endpoint, - peer_manager, - tcp_proxy: TcpProxyForKcpSrc(tcp_proxy), tasks, } } - pub async fn start(&self) { - self.peer_manager - .add_nic_packet_process_pipeline(Box::new(self.tcp_proxy.clone())) - .await; - self.peer_manager - .add_packet_process_pipeline(Box::new(self.tcp_proxy.0.clone())) - .await; - self.peer_manager - .add_packet_process_pipeline(Box::new(KcpEndpointFilter { - kcp_endpoint: self.kcp_endpoint.clone(), - is_src: true, - })) - .await; - self.tcp_proxy.0.start(false).await.unwrap(); - } - - pub fn get_tcp_proxy(&self) -> Arc> { - self.tcp_proxy.0.clone() - } - pub fn get_kcp_endpoint(&self) -> Arc { self.kcp_endpoint.clone() } + + async fn stop(&mut self) { + self.tasks.shutdown().await; + } } pub struct KcpProxyDst { kcp_endpoint: Arc, - peer_manager: Arc, - proxy_entries: Arc>, - cidr_set: Arc, + destination_ingress: WrappedTransportDestinationIngress, tasks: JoinSet<()>, + accept_cancel: CancellationToken, + accept_task: Option>, } impl KcpProxyDst { - pub async fn new(peer_manager: Arc) -> Self { + pub async fn new( + destination_ingress: WrappedTransportDestinationIngress, + datagrams: tokio::sync::mpsc::Sender, + ) -> Self { let mut kcp_endpoint = create_kcp_endpoint(); kcp_endpoint.run().await; let mut tasks = JoinSet::new(); let output_receiver = kcp_endpoint.output_receiver().unwrap(); tasks.spawn(handle_kcp_output( - peer_manager.clone(), output_receiver, - false, + WrappedTransportRole::Destination, + datagrams, )); - let cidr_set = CidrSet::new(peer_manager.get_global_ctx()); Self { kcp_endpoint: Arc::new(kcp_endpoint), - peer_manager, - proxy_entries: Arc::new(DashMap::new()), - cidr_set: Arc::new(cidr_set), + destination_ingress, tasks, + accept_cancel: CancellationToken::new(), + accept_task: None, } } - #[tracing::instrument(ret, skip(route))] + #[tracing::instrument(ret, skip(destination_ingress))] async fn handle_one_in_stream( kcp_stream: KcpStream, - global_ctx: ArcGlobalCtx, - proxy_entries: Arc>, - cidr_set: Arc, - route: Arc, - ) -> Result<()> { + destination_ingress: WrappedTransportDestinationIngress, + ) -> anyhow::Result<()> { let mut conn_data = kcp_stream.conn_data().clone(); let parsed_conn_data = KcpConnData::decode(&mut conn_data) .with_context(|| format!("failed to decode kcp conn data: {:?}", conn_data))?; - let mut dst_socket: SocketAddr = parsed_conn_data + let dst_socket: SocketAddr = parsed_conn_data .dst .ok_or(anyhow::anyhow!( "failed to get dst socket from kcp conn data: {:?}", @@ -328,154 +181,188 @@ impl KcpProxyDst { .into(); let src_socket: SocketAddr = parsed_conn_data.src.unwrap_or_default().into(); - if let IpAddr::V4(dst_v4_ip) = dst_socket.ip() { - let mut real_ip = dst_v4_ip; - if cidr_set.contains_v4(dst_v4_ip, &mut real_ip) { - dst_socket.set_ip(real_ip.into()); - } - }; - - let conn_id = kcp_stream.conn_id(); - proxy_entries.insert( - conn_id, - TcpProxyEntry { - src: parsed_conn_data.src, - dst: parsed_conn_data.dst, - start_time: chrono::Local::now().timestamp() as u64, - state: TcpProxyEntryState::ConnectingDst.into(), - transport_type: TcpProxyEntryTransportType::Kcp.into(), - }, - ); - defer! { - proxy_entries.remove(&conn_id); - if proxy_entries.capacity() - proxy_entries.len() > 16 { - proxy_entries.shrink_to_fit(); - } - } - - let src_ip = src_socket.ip(); - let dst_ip = dst_socket.ip(); - let (src_groups, dst_groups) = tokio::join!( - route.get_peer_groups_by_ip(&src_ip), - route.get_peer_groups_by_ip(&dst_ip) - ); - - if global_ctx.should_deny_proxy(&dst_socket, false) { - return Err(anyhow::anyhow!( - "dst socket {:?} is in running listeners, ignore it", - dst_socket - ) - .into()); - } - - let send_to_self = global_ctx.is_ip_local_virtual_ip(&dst_ip); - if send_to_self && global_ctx.no_tun() { - dst_socket = format!("127.0.0.1:{}", dst_socket.port()).parse().unwrap(); - } - - let acl_handler = ProxyAclHandler { - acl_filter: global_ctx.get_acl_filter().clone(), - packet_info: PacketInfo { - src_ip, - dst_ip, - src_port: Some(src_socket.port()), - dst_port: Some(dst_socket.port()), - protocol: Protocol::Tcp, - packet_size: conn_data.len(), - src_groups, - dst_groups, - }, - chain_type: if send_to_self { - ChainType::Inbound - } else { - ChainType::Forward - }, - }; - acl_handler.handle_packet(&conn_data)?; - - tracing::debug!("kcp connect to dst socket: {:?}", dst_socket); - - let _g = global_ctx.net_ns.guard(); - let connector = NatDstTcpConnector {}; - let ret = connector - .connect("0.0.0.0:0".parse().unwrap(), dst_socket) - .await?; - - if let Some(mut e) = proxy_entries.get_mut(&kcp_stream.conn_id()) { - e.state = TcpProxyEntryState::Connected.into(); - } - - acl_handler - .copy_bidirection_with_acl(kcp_stream, ret) - .await?; - - Ok(()) + destination_ingress + .submit(WrappedTransportAcceptedStream { + src: src_socket, + dst: dst_socket, + initial_acl_packet_size: conn_data.len(), + stream: Box::new(kcp_stream), + }) + .await } async fn run_accept_task(&mut self) { let kcp_endpoint = self.kcp_endpoint.clone(); - let global_ctx = self.peer_manager.get_global_ctx(); - let proxy_entries = self.proxy_entries.clone(); - let cidr_set = self.cidr_set.clone(); - let route = Arc::new(self.peer_manager.get_route()); - self.tasks.spawn(async move { - while let Ok(conn) = kcp_endpoint.accept().await { - let stream = KcpStream::new(&kcp_endpoint, conn) - .ok_or(anyhow::anyhow!("failed to create kcp stream")) - .unwrap(); + let destination_ingress = self.destination_ingress.clone(); + let cancel = self.accept_cancel.clone(); + self.accept_task = Some(tokio::spawn(async move { + let mut streams = JoinSet::new(); + loop { + tokio::select! { + biased; + _ = cancel.cancelled() => { + streams.shutdown().await; + break; + } + accepted = kcp_endpoint.accept() => { + let Ok(conn) = accepted else { + streams.shutdown().await; + break; + }; + let Some(stream) = KcpStream::new(&kcp_endpoint, conn) else { + tracing::warn!("failed to create accepted kcp stream"); + continue; + }; - let global_ctx = global_ctx.clone(); - let proxy_entries = proxy_entries.clone(); - let cidr_set = cidr_set.clone(); - let route = route.clone(); - tokio::spawn(async move { - let _ = Self::handle_one_in_stream( - stream, - global_ctx, - proxy_entries, - cidr_set, - route, - ) - .await; - }); + let destination_ingress = destination_ingress.clone(); + streams.spawn(async move { + let _ = Self::handle_one_in_stream(stream, destination_ingress).await; + }); + } + _ = streams.join_next(), if !streams.is_empty() => {} + } } - }); + })); } pub async fn start(&mut self) { self.run_accept_task().await; - self.peer_manager - .add_packet_process_pipeline(Box::new(KcpEndpointFilter { - kcp_endpoint: self.kcp_endpoint.clone(), - is_src: false, - })) - .await; + } + + async fn stop(&mut self) { + self.accept_cancel.cancel(); + if let Some(task) = self.accept_task.as_mut() { + let _ = task.await; + } + self.accept_task.take(); + self.tasks.shutdown().await; } } -#[derive(Clone)] -pub struct KcpProxyDstRpcService(Weak>); +#[derive(Default)] +struct KcpProxyServiceState { + src: Option, + dst: Option, +} -impl KcpProxyDstRpcService { - pub fn new(kcp_proxy_dst: &KcpProxyDst) -> Self { - Self(Arc::downgrade(&kcp_proxy_dst.proxy_entries)) +impl KcpProxyServiceState { + async fn stop(&mut self) { + if let Some(dst) = &mut self.dst { + dst.stop().await; + } + if let Some(src) = &mut self.src { + src.stop().await; + } + } +} + +pub struct KcpProxyService { + state: Mutex>, +} + +impl KcpProxyService { + pub fn new() -> Self { + Self { + state: Mutex::new(None), + } } } #[async_trait::async_trait] -impl TcpProxyRpc for KcpProxyDstRpcService { - type Controller = BaseController; - async fn list_tcp_proxy_entry( +impl WrappedTransportEngine for KcpProxyService { + async fn prepare(&self, options: WrappedTransportEngineStart) -> anyhow::Result<()> { + let mut state = self.state.lock().await; + if state.is_some() { + return Ok(()); + } + let directions = options.directions; + + let src = if directions.source { + let src = KcpProxySrc::new(options.datagrams.clone()).await; + Some(src) + } else { + None + }; + let dst = if directions.destination { + let destination_ingress = options + .destination_ingress + .ok_or_else(|| anyhow!("KCP destination ingress is required"))?; + let dst = KcpProxyDst::new(destination_ingress, options.datagrams).await; + Some(dst) + } else { + None + }; + + *state = Some(KcpProxyServiceState { src, dst }); + Ok(()) + } + + async fn activate(&self) -> anyhow::Result<()> { + let mut state = self.state.lock().await; + let state = state + .as_mut() + .ok_or_else(|| anyhow!("KCP engine is not prepared"))?; + if let Some(dst) = state.dst.as_mut() { + dst.start().await; + } + Ok(()) + } + + async fn inject_peer_datagram( &self, - _: BaseController, - _request: ListTcpProxyEntryRequest, // Accept request of type HelloRequest - ) -> std::result::Result { - let mut reply = ListTcpProxyEntryResponse::default(); - if let Some(tcp_proxy) = self.0.upgrade() { - for item in tcp_proxy.iter() { - reply.entries.push(*item.value()); + role: WrappedTransportRole, + _from_peer_id: u32, + payload: Bytes, + ) -> anyhow::Result<()> { + let endpoint = { + let state = self.state.lock().await; + match (state.as_ref(), role) { + (Some(state), WrappedTransportRole::Source) => { + state.src.as_ref().map(KcpProxySrc::get_kcp_endpoint) + } + (Some(state), WrappedTransportRole::Destination) => { + state.dst.as_ref().map(|dst| dst.kcp_endpoint.clone()) + } + (None, _) => None, } } - Ok(reply) + .ok_or_else(|| anyhow!("KCP {role:?} endpoint is not active"))?; + + endpoint + .input_sender_ref() + .send(KcpPacket::from(bytes::BytesMut::from(payload))) + .await + .map_err(|error| anyhow!("failed to inject KCP datagram: {error}")) + } + + async fn connect_source( + &self, + request: WrappedTransportConnect, + ) -> anyhow::Result> { + let endpoint = { + let state = self.state.lock().await; + state + .as_ref() + .and_then(|state| state.src.as_ref()) + .map(KcpProxySrc::get_kcp_endpoint) + } + .ok_or_else(|| anyhow!("KCP source endpoint is not prepared"))?; + let stream = connect_kcp_source( + endpoint, + request.my_peer_id, + request.dst_peer_id, + request.src, + request.dst, + ) + .await?; + Ok(Box::new(stream)) + } + + async fn stop(&self) { + let mut state = self.state.lock().await; + if let Some(active) = state.as_mut() { + active.stop().await; + } + state.take(); } } diff --git a/easytier/src/gateway/mod.rs b/easytier/src/gateway/mod.rs index ab057635..c0e3e367 100644 --- a/easytier/src/gateway/mod.rs +++ b/easytier/src/gateway/mod.rs @@ -1,101 +1,11 @@ -use dashmap::DashMap; -use std::sync::{Arc, Mutex}; -use tokio::task::JoinSet; - -use crate::common::global_ctx::ArcGlobalCtx; +#[cfg(any(feature = "kcp", feature = "quic"))] +mod hedge; +#[cfg(feature = "icmp-proxy")] pub mod icmp_proxy; -pub mod ip_reassembler; -pub mod tcp_proxy; -#[cfg(feature = "smoltcp")] -pub mod tokio_smoltcp; -pub mod udp_proxy; - -#[cfg(feature = "socks5")] -pub mod fast_socks5; -#[cfg(feature = "socks5")] -pub mod socks5; #[cfg(feature = "kcp")] pub mod kcp_proxy; -mod wrapped_proxy; #[cfg(feature = "quic")] pub mod quic_proxy; - -#[derive(Debug)] -pub(crate) struct CidrSet { - global_ctx: ArcGlobalCtx, - cidr_set: Arc>>, - tasks: JoinSet<()>, - - mapped_to_real: Arc>, -} - -impl CidrSet { - pub fn new(global_ctx: ArcGlobalCtx) -> Self { - let mut ret = Self { - global_ctx, - cidr_set: Arc::new(Mutex::new(vec![])), - tasks: JoinSet::new(), - - mapped_to_real: Arc::new(DashMap::new()), - }; - ret.run_cidr_updater(); - ret - } - - fn run_cidr_updater(&mut self) { - let global_ctx = self.global_ctx.clone(); - let cidr_set = self.cidr_set.clone(); - let mapped_to_real = self.mapped_to_real.clone(); - self.tasks.spawn(async move { - let mut last_cidrs = vec![]; - loop { - let cidrs = global_ctx.config.get_proxy_cidrs(); - if cidrs != last_cidrs { - last_cidrs = cidrs.clone(); - mapped_to_real.clear(); - cidr_set.lock().unwrap().clear(); - for cidr in cidrs.iter() { - let real_cidr = cidr.cidr; - let mapped = cidr.mapped_cidr.unwrap_or(real_cidr); - cidr_set.lock().unwrap().push(mapped); - - if mapped != real_cidr { - mapped_to_real.insert(mapped, real_cidr); - } - } - } - tokio::time::sleep(std::time::Duration::from_secs(1)).await; - } - }); - } - - pub fn contains_v4(&self, ipv4: std::net::Ipv4Addr, real_ip: &mut std::net::Ipv4Addr) -> bool { - let ip = ipv4; - let s = self.cidr_set.lock().unwrap(); - for cidr in s.iter() { - if cidr.contains(&ip) { - if let Some(real_cidr) = self.mapped_to_real.get(cidr).map(|v| *v.value()) { - let origin_network_bits = real_cidr.first().address().to_bits(); - let network_mask = cidr.mask().to_bits(); - - let mut converted_ip = ipv4.to_bits(); - converted_ip &= !network_mask; - converted_ip |= origin_network_bits; - - *real_ip = std::net::Ipv4Addr::from(converted_ip); - } else { - *real_ip = ipv4; - } - return true; - } - } - false - } - - pub fn is_empty(&self) -> bool { - self.cidr_set.lock().unwrap().is_empty() - } -} diff --git a/easytier/src/gateway/quic_proxy.rs b/easytier/src/gateway/quic_proxy.rs index a943422c..ab06d1ce 100644 --- a/easytier/src/gateway/quic_proxy.rs +++ b/easytier/src/gateway/quic_proxy.rs @@ -1,56 +1,46 @@ -use crate::common::PeerId; -use crate::common::acl_processor::PacketInfo; -use crate::common::global_ctx::{ArcGlobalCtx, GlobalCtx}; -use crate::gateway::CidrSet; -use crate::gateway::tcp_proxy::{NatDstConnector, TcpProxy}; -use crate::gateway::wrapped_proxy::{ProxyAclHandler, TcpProxyForWrappedSrcTrait}; -use crate::peers::PeerPacketFilter; -use crate::peers::peer_manager::PeerManager; -use crate::proto::acl::{ChainType, Protocol}; -use crate::proto::api::instance::{ - ListTcpProxyEntryRequest, ListTcpProxyEntryResponse, TcpProxyEntry, TcpProxyEntryState, - TcpProxyEntryTransportType, TcpProxyRpc, -}; +use super::hedge::HedgeExt; use crate::proto::peer_rpc::KcpConnData as QuicConnData; -use crate::proto::rpc_types; -use crate::proto::rpc_types::controller::BaseController; -use crate::tunnel::packet_def::{ - PacketType, PeerManagerHeader, TAIL_RESERVED_SIZE, ZCPacket, ZCPacketType, -}; use crate::tunnel::quic::{client_config, endpoint_config, server_config}; -use crate::utils::task::HedgeExt; -use anyhow::{Context, Error, anyhow, bail, ensure}; +use anyhow::{Context, Error, anyhow, ensure}; use atomic_refcell::AtomicRefCell; use bytes::{BufMut, Bytes, BytesMut}; -use dashmap::DashMap; -use derivative::Derivative; use derive_more::{Constructor, Deref, DerefMut, From, Into}; -use guarden::defer; +use easytier_core::config::PeerId; +use easytier_core::packet::{PacketType, TAIL_RESERVED_SIZE, ZCPacket, ZCPacketType}; use moka::future::Cache; use prost::Message; use quinn::udp::{EcnCodepoint, RecvMeta, Transmit}; use quinn::{ - AsyncUdpSocket, Connection, ConnectionError, Endpoint, RecvStream, SendStream, StreamId, - UdpPoller, WriteError, default_runtime, + AsyncUdpSocket, Connection, ConnectionError, Endpoint, RecvStream, SendStream, UdpPoller, + WriteError, default_runtime, }; use std::cmp::min; -use std::future::Future; use std::io::IoSliceMut; use std::net::{IpAddr, Ipv4Addr, SocketAddr}; use std::pin::Pin; use std::ptr::copy_nonoverlapping; -use std::sync::{Arc, Weak}; +use std::sync::Arc; use std::task::Poll; use std::time::Duration; use tokio::io::{AsyncReadExt, Join, join}; +use tokio::select; +use tokio::sync::Mutex; use tokio::sync::mpsc::error::TrySendError; use tokio::sync::mpsc::{Receiver, Sender, channel}; -use tokio::task::JoinSet; +use tokio::task::{JoinHandle, JoinSet}; use tokio::time::timeout; -use tokio::{join, select}; -use tokio_util::sync::PollSender; +use tokio_util::sync::{CancellationToken, PollSender}; use tracing::{debug, error, info, instrument, trace, warn}; +use easytier_core::{ + gateway::proxy::traits::TcpProxyStream, + gateway::proxy::wrapped_transport::{ + WrappedTransportAcceptedStream, WrappedTransportConnect, WrappedTransportDatagram, + WrappedTransportDatagramBuffer, WrappedTransportDestinationIngress, WrappedTransportEngine, + WrappedTransportEngineStart, WrappedTransportKind, WrappedTransportRole, + }, +}; + //region packet #[derive(Debug, Constructor)] struct QuicPacket { @@ -263,13 +253,6 @@ struct QuicStream { inner: QuicStreamInner, } -impl QuicStream { - #[inline] - fn id(&self) -> (StreamId, StreamId) { - (self.reader().id(), self.writer().id()) - } -} - impl From<(SendStream, RecvStream)> for QuicStream { #[inline] fn from(value: (SendStream, RecvStream)) -> Self { @@ -281,35 +264,16 @@ impl From<(SendStream, RecvStream)> for QuicStream { #[derive(Debug, Clone)] pub struct NatDstQuicConnector { pub(crate) endpoint: Endpoint, - pub(crate) peer_mgr: Weak, pub(crate) conn_map: Cache, } -#[async_trait::async_trait] -impl NatDstConnector for NatDstQuicConnector { - type DstStream = QuicStreamInner; - - async fn connect( +impl NatDstQuicConnector { + async fn connect_to_peer( &self, + dst_peer: PeerId, src: SocketAddr, nat_dst: SocketAddr, - ) -> anyhow::Result { - let peer_mgr = self - .peer_mgr - .upgrade() - .ok_or_else(|| anyhow!("peer manager is not available"))?; - - let dst_peer = { - let SocketAddr::V4(addr) = nat_dst else { - bail!("ipv6 is not supported"); - }; - peer_mgr - .get_peer_map() - .get_peer_id_by_ipv4(addr.ip()) - .await - .ok_or_else(|| anyhow!("no peer found for nat dst: {}", nat_dst))? - }; - + ) -> anyhow::Result { tracing::trace!(?nat_dst, ?dst_peer, "quic nat"); let header = { @@ -405,56 +369,6 @@ impl NatDstConnector for NatDstQuicConnector { break result; } } - - #[inline] - fn check_packet_from_peer_fast(&self, _cidr_set: &CidrSet, _global_ctx: &GlobalCtx) -> bool { - true - } - - #[inline] - fn check_packet_from_peer( - &self, - _cidr_set: &CidrSet, - _global_ctx: &GlobalCtx, - hdr: &PeerManagerHeader, - _ipv4: &Ipv4Addr, - _real_dst_ip: &mut Ipv4Addr, - ) -> bool { - hdr.from_peer_id == hdr.to_peer_id && hdr.is_quic_src_modified() - } - - #[inline] - fn transport_type(&self) -> TcpProxyEntryTransportType { - TcpProxyEntryTransportType::Quic - } -} - -#[derive(Clone)] -struct TcpProxyForQuicSrc(Arc>); - -#[async_trait::async_trait] -impl TcpProxyForWrappedSrcTrait for TcpProxyForQuicSrc { - type Connector = NatDstQuicConnector; - - #[inline] - fn get_tcp_proxy(&self) -> &Arc> { - &self.0 - } - - #[inline] - fn mark_src_modified(hdr: &mut PeerManagerHeader) -> &mut PeerManagerHeader { - hdr.mark_quic_src_modified() - } - - #[inline] - async fn check_dst_allow_wrapped_input(&self, dst_ip: &Ipv4Addr) -> bool { - let Some(peer_manager) = self.0.get_peer_manager() else { - return false; - }; - peer_manager - .check_allow_quic_to_dst(&IpAddr::V4(*dst_ip)) - .await - } } #[derive(Debug)] @@ -464,14 +378,6 @@ enum QuicProxyRole { } impl QuicProxyRole { - #[inline] - const fn incoming(&self) -> PacketType { - match self { - QuicProxyRole::Src => PacketType::QuicDst, - QuicProxyRole::Dst => PacketType::QuicSrc, - } - } - #[inline] const fn outgoing(&self) -> PacketType { match self { @@ -481,41 +387,10 @@ impl QuicProxyRole { } } -// Receive packets from peers and forward them to the QUIC endpoint -#[derive(Debug)] -struct QuicPacketReceiver { - tx: Sender, - role: QuicProxyRole, -} - -#[async_trait::async_trait] -impl PeerPacketFilter for QuicPacketReceiver { - async fn try_process_packet_from_peer(&self, packet: ZCPacket) -> Option { - let header = packet.peer_manager_header().unwrap(); - - if header.packet_type != self.role.incoming() as u8 { - return Some(packet); - } - - let addr = QuicAddr::new(header.from_peer_id.get(), self.role.outgoing()); - - if let Err(e) = self.tx.try_send(QuicPacket::new( - addr.into(), - packet.payload_bytes(), - None, - None, - )) { - debug!("failed to send quic packet to endpoint: {:?}", e); - } - - None - } -} - // Send to peers packets received from the QUIC endpoint #[derive(Debug)] struct QuicPacketSender { - peer_mgr: Arc, + datagrams: Sender, rx: Receiver, header: Bytes, @@ -542,48 +417,40 @@ impl QuicPacketSender { let mut payload = payload.split_to(len); payload[..self.margins.header].copy_from_slice(&self.header); payload.truncate(len - self.margins.trailer); - let mut packet = ZCPacket::new_from_buf(payload, self.zc_packet_type); - - packet.fill_peer_manager_hdr( - self.peer_mgr.my_peer_id(), - addr.peer_id, - addr.packet_type as u8, - ); - - if let Err(e) = self.peer_mgr.send_msg_for_proxy(packet, addr.peer_id).await { - error!("failed to send QUIC packet to peer: {:?}", e); + let role = match addr.packet_type { + PacketType::QuicSrc => WrappedTransportRole::Source, + PacketType::QuicDst => WrappedTransportRole::Destination, + packet_type => { + error!(?packet_type, "invalid QUIC proxy output packet type"); + continue; + } + }; + if self + .datagrams + .send(WrappedTransportDatagram { + transport: WrappedTransportKind::Quic, + role, + peer_id: addr.peer_id, + buffer: WrappedTransportDatagramBuffer::from_packet_buffer( + payload, + self.zc_packet_type, + ), + }) + .await + .is_err() + { + break; } } } } } -#[derive(Derivative, Clone)] -#[derivative(Debug)] -struct QuicStreamContext { - global_ctx: ArcGlobalCtx, - proxy_entries: Arc>, - cidr_set: Arc, - #[derivative(Debug = "ignore")] - route: Arc, -} - -impl QuicStreamContext { - fn new(peer_mgr: Arc) -> Self { - let global_ctx = peer_mgr.get_global_ctx(); - Self { - global_ctx: global_ctx.clone(), - proxy_entries: Arc::new(DashMap::new()), - cidr_set: Arc::new(CidrSet::new(global_ctx.clone())), - route: Arc::new(peer_mgr.get_route()), - } - } -} - struct QuicStreamReceiver { endpoint: Endpoint, tasks: JoinSet<()>, - ctx: Arc, + destination_ingress: WrappedTransportDestinationIngress, + cancel: CancellationToken, } impl QuicStreamReceiver { @@ -592,6 +459,8 @@ impl QuicStreamReceiver { select! { biased; + _ = self.cancel.cancelled() => break, + Some(incoming) = self.endpoint.accept() => { let addr = incoming.remote_address(); let connection = match incoming.accept() { @@ -603,21 +472,30 @@ impl QuicStreamReceiver { }; let addr = connection.remote_address(); - let connection = match connection.await { - Ok(connection) => connection, - Err(e) => { - error!("failed to accept quic connection from {:?}: {:?}", addr, e); - continue; + let connection = select! { + biased; + _ = self.cancel.cancelled() => break, + result = connection => { + match result { + Ok(connection) => connection, + Err(e) => { + error!("failed to accept quic connection from {:?}: {:?}", addr, e); + continue; + } + } } }; - let ctx = self.ctx.clone(); + let destination_ingress = self.destination_ingress.clone(); + let cancel = self.cancel.clone(); self.tasks.spawn(async move { let mut tasks = JoinSet::new(); loop { select! { biased; + _ = cancel.cancelled() => break, + e = connection.closed() => { info!("connection to {:?} closed: {:?}", addr, e); break; @@ -632,15 +510,10 @@ impl QuicStreamReceiver { } }; - let ctx = ctx.clone(); + let destination_ingress = destination_ingress.clone(); tasks.spawn(async move { - match Self::establish_stream(stream, ctx).await { - Ok(transfer_fut) => { - if let Err(e) = transfer_fut.await { - warn!("quic stream transfer error: {:?}", e); - } - } - Err(e) => warn!("failed to establish quic stream: {:?}", e), + if let Err(e) = Self::submit_stream(stream, destination_ingress).await { + warn!("failed to submit quic stream: {:?}", e); } }); } @@ -651,6 +524,7 @@ impl QuicStreamReceiver { } } + tasks.shutdown().await; connection.close(1u32.into(), b"error"); }); } @@ -663,6 +537,11 @@ impl QuicStreamReceiver { } } } + if self.cancel.is_cancelled() { + while self.tasks.join_next().await.is_some() {} + } else { + self.tasks.shutdown().await; + } } async fn read_stream_header(stream: &mut QuicStream) -> Result { @@ -687,144 +566,66 @@ impl QuicStreamReceiver { Ok(header.into()) } - async fn establish_stream( + async fn submit_stream( mut stream: QuicStream, - ctx: Arc, - ) -> Result>, Error> { + destination_ingress: WrappedTransportDestinationIngress, + ) -> Result<(), Error> { let conn_data = Self::read_stream_header(&mut stream).await?; let conn_data_parsed = QuicConnData::decode(conn_data.as_ref()) .context("failed to decode quic stream header")?; - let handle = stream.id(); - let proxy_entries = &ctx.proxy_entries; - proxy_entries.insert( - handle, - TcpProxyEntry { - src: conn_data_parsed.src, - dst: conn_data_parsed.dst, - start_time: chrono::Local::now().timestamp() as u64, - state: TcpProxyEntryState::ConnectingDst.into(), - transport_type: TcpProxyEntryTransportType::Quic.into(), - }, - ); - defer! { - proxy_entries.remove(&handle); - if proxy_entries.capacity() - proxy_entries.len() > 16 { - proxy_entries.shrink_to_fit(); - } - } - let src_socket: SocketAddr = conn_data_parsed .src .ok_or_else(|| anyhow!("missing src addr in quic stream header"))? .into(); - let mut dst_socket: SocketAddr = conn_data_parsed + let dst_socket: SocketAddr = conn_data_parsed .dst .ok_or_else(|| anyhow!("missing dst addr in quic stream header"))? .into(); - if let IpAddr::V4(dst_v4_ip) = dst_socket.ip() { - let mut real_ip = dst_v4_ip; - if ctx.cidr_set.contains_v4(dst_v4_ip, &mut real_ip) { - dst_socket.set_ip(real_ip.into()); - } - }; - - let src_ip = src_socket.ip(); - let dst_ip = dst_socket.ip(); - - let route = ctx.route.clone(); - let (src_groups, dst_groups) = join!( - route.get_peer_groups_by_ip(&src_ip), - route.get_peer_groups_by_ip(&dst_ip) - ); - - let global_ctx = ctx.global_ctx.clone(); - if global_ctx.should_deny_proxy(&dst_socket, false) { - return Err(anyhow::anyhow!( - "dst socket {:?} is in running listeners, ignore it", - dst_socket - )); - } - - let send_to_self = global_ctx.is_ip_local_virtual_ip(&dst_ip); - if send_to_self && global_ctx.no_tun() { - dst_socket = format!("127.0.0.1:{}", dst_socket.port()).parse()?; - } - - let acl_handler = ProxyAclHandler { - acl_filter: global_ctx.get_acl_filter().clone(), - packet_info: PacketInfo { - src_ip, - dst_ip, - src_port: Some(src_socket.port()), - dst_port: Some(dst_socket.port()), - protocol: Protocol::Tcp, - packet_size: conn_data.len(), - src_groups, - dst_groups, - }, - chain_type: if send_to_self { - ChainType::Inbound - } else { - ChainType::Forward - }, - }; - acl_handler.handle_packet(&conn_data)?; - - debug!("quic connect to dst socket: {:?}", dst_socket); - - let _g = global_ctx.net_ns.guard(); - let connector = crate::gateway::tcp_proxy::NatDstTcpConnector {}; - let ret = connector.connect("0.0.0.0:0".parse()?, dst_socket).await?; - - if let Some(mut e) = proxy_entries.get_mut(&handle) { - e.state = TcpProxyEntryState::Connected.into(); - } - - Ok(async move { - acl_handler - .copy_bidirection_with_acl(stream.inner, ret) - .await - }) + destination_ingress + .submit(WrappedTransportAcceptedStream { + src: src_socket, + dst: dst_socket, + initial_acl_packet_size: conn_data.len(), + stream: Box::new(stream.inner), + }) + .await } } pub struct QuicProxy { - peer_mgr: Arc, - endpoint: Option, + input_tx: Option>>, - src: Option, - dst: Option, + source_connector: Option, + destination_ingress: Option, tasks: JoinSet<()>, + stream_cancel: CancellationToken, + stream_task: Option>, } impl QuicProxy { - #[inline] - pub fn src(&self) -> Option<&QuicProxySrc> { - self.src.as_ref() - } - - #[inline] - pub fn dst(&self) -> Option<&QuicProxyDst> { - self.dst.as_ref() - } -} - -impl QuicProxy { - pub fn new(peer_mgr: Arc) -> Self { + pub fn new() -> Self { Self { - peer_mgr, endpoint: None, - src: None, - dst: None, + input_tx: None, + source_connector: None, + destination_ingress: None, tasks: JoinSet::new(), + stream_cancel: CancellationToken::new(), + stream_task: None, } } - pub async fn run(&mut self, src: bool, dst: bool) { + pub async fn prepare( + &mut self, + my_peer_id: u32, + src: bool, + destination_ingress: Option, + datagrams: Sender, + ) { trace!("quic proxy starting"); if self.endpoint.is_some() { @@ -845,10 +646,12 @@ impl QuicProxy { let margins = (header.len(), TAIL_RESERVED_SIZE).into(); let (in_tx, in_rx) = channel(1024); + let in_tx = Arc::new(in_tx); + self.input_tx = Some(in_tx.clone()); let (out_tx, out_rx) = channel(1024); let socket = QuicSocket { - addr: SocketAddr::new(Ipv4Addr::from(self.peer_mgr.my_peer_id()).into(), 0), + addr: SocketAddr::new(Ipv4Addr::from(my_peer_id).into(), 0), rx: AtomicRefCell::new(in_rx), tx: out_tx, margins, @@ -864,10 +667,9 @@ impl QuicProxy { endpoint.set_default_client_config(client_config()); self.endpoint = Some(endpoint.clone()); - let peer_mgr = self.peer_mgr.clone(); self.tasks.spawn( QuicPacketSender { - peer_mgr, + datagrams, rx: out_rx, header, zc_packet_type, @@ -876,141 +678,165 @@ impl QuicProxy { .run(), ); - let peer_mgr = self.peer_mgr.clone(); - if src { - if self.src.is_some() { + if self.source_connector.is_some() { error!("quic proxy src already running"); return; } - let tcp_proxy = TcpProxyForQuicSrc(TcpProxy::new( - peer_mgr.clone(), - NatDstQuicConnector { - endpoint: endpoint.clone(), - peer_mgr: Arc::downgrade(&peer_mgr), - conn_map: Cache::builder() - .max_capacity(u8::MAX.into()) // cf. quinn transport config (max_concurrent_bidi_streams) - .time_to_idle(Duration::from_secs(600)) // cf. quinn transport config (max_idle_timeout) - .build(), - }, - )); - - let src = QuicProxySrc { - peer_mgr: peer_mgr.clone(), - tcp_proxy, - tx: in_tx.clone(), - }; - src.run().await; - - self.src = Some(src); + self.source_connector = Some(NatDstQuicConnector { + endpoint: endpoint.clone(), + conn_map: Cache::builder() + .max_capacity(u8::MAX.into()) // cf. quinn transport config (max_concurrent_bidi_streams) + .time_to_idle(Duration::from_secs(600)) // cf. quinn transport config (max_idle_timeout) + .build(), + }); } - if dst { - if self.dst.is_some() { + if let Some(destination_ingress) = destination_ingress { + if self.destination_ingress.is_some() { error!("quic proxy dst already running"); return; } + self.destination_ingress = Some(destination_ingress); + } + } - let stream_ctx = Arc::new(QuicStreamContext::new(peer_mgr.clone())); - - let dst = QuicProxyDst { - peer_mgr: peer_mgr.clone(), - tx: in_tx.clone(), - stream_ctx: stream_ctx.clone(), - }; - dst.run().await; - - self.tasks.spawn( + async fn activate(&mut self) -> anyhow::Result<()> { + if let Some(destination_ingress) = self.destination_ingress.as_ref() { + let endpoint = self + .endpoint + .as_ref() + .cloned() + .ok_or_else(|| anyhow!("QUIC endpoint is not prepared"))?; + self.stream_task = Some(tokio::spawn( QuicStreamReceiver { - endpoint: endpoint.clone(), + endpoint, tasks: JoinSet::new(), - ctx: stream_ctx, + destination_ingress: destination_ingress.clone(), + cancel: self.stream_cancel.clone(), } .run(), - ); + )); + } + Ok(()) + } - self.dst = Some(dst); + async fn stop(&mut self) { + self.stream_cancel.cancel(); + if let Some(task) = self.stream_task.as_mut() { + let _ = task.await; + } + self.stream_task.take(); + self.tasks.shutdown().await; + if let Some(endpoint) = self.endpoint.take() { + endpoint.close(1u32.into(), b"stopped"); } } } -pub struct QuicProxySrc { - peer_mgr: Arc, - tcp_proxy: TcpProxyForQuicSrc, - - tx: Sender, +pub struct QuicProxyService { + state: Mutex>, } -impl QuicProxySrc { - #[inline] - pub fn get_tcp_proxy(&self) -> Arc> { - self.tcp_proxy.get_tcp_proxy().clone() - } -} - -impl QuicProxySrc { - async fn run(&self) { - trace!("quic proxy src starting"); - self.peer_mgr - .add_nic_packet_process_pipeline(Box::new(self.tcp_proxy.clone())) - .await; - self.peer_mgr - .add_packet_process_pipeline(Box::new(self.tcp_proxy.0.clone())) - .await; - self.peer_mgr - .add_packet_process_pipeline(Box::new(QuicPacketReceiver { - tx: self.tx.clone(), - role: QuicProxyRole::Src, - })) - .await; - self.tcp_proxy.0.start(false).await.unwrap(); - } -} - -pub struct QuicProxyDst { - peer_mgr: Arc, - - tx: Sender, - stream_ctx: Arc, -} - -impl QuicProxyDst { - async fn run(&self) { - trace!("quic proxy dst starting"); - self.peer_mgr - .add_packet_process_pipeline(Box::new(QuicPacketReceiver { - tx: self.tx.clone(), - role: QuicProxyRole::Dst, - })) - .await; - } -} - -#[derive(Clone, Deref, DerefMut, From, Into)] -pub struct QuicProxyDstRpcService(Weak>); - -impl QuicProxyDstRpcService { - pub fn new(quic_proxy_dst: &QuicProxyDst) -> Self { - Self(Arc::downgrade(&quic_proxy_dst.stream_ctx.proxy_entries)) +impl QuicProxyService { + pub fn new() -> Self { + Self { + state: Mutex::new(None), + } } } #[async_trait::async_trait] -impl TcpProxyRpc for QuicProxyDstRpcService { - type Controller = BaseController; - async fn list_tcp_proxy_entry( - &self, - _: BaseController, - _request: ListTcpProxyEntryRequest, // Accept request of type HelloRequest - ) -> Result { - let mut reply = ListTcpProxyEntryResponse::default(); - if let Some(tcp_proxy) = self.0.upgrade() { - for item in tcp_proxy.iter() { - reply.entries.push(*item.value()); - } +impl WrappedTransportEngine for QuicProxyService { + async fn prepare(&self, options: WrappedTransportEngineStart) -> anyhow::Result<()> { + let mut state = self.state.lock().await; + if state.is_some() { + return Ok(()); } - Ok(reply) + let directions = options.directions; + let destination_ingress = if directions.destination { + Some( + options + .destination_ingress + .ok_or_else(|| anyhow!("QUIC destination ingress is required"))?, + ) + } else { + None + }; + + let mut proxy = QuicProxy::new(); + if directions.source || directions.destination { + proxy + .prepare( + options.my_peer_id, + directions.source, + destination_ingress, + options.datagrams, + ) + .await; + } + + *state = Some(proxy); + Ok(()) + } + + async fn activate(&self) -> anyhow::Result<()> { + let mut state = self.state.lock().await; + state + .as_mut() + .ok_or_else(|| anyhow!("QUIC engine is not prepared"))? + .activate() + .await + } + + async fn inject_peer_datagram( + &self, + role: WrappedTransportRole, + from_peer_id: u32, + payload: Bytes, + ) -> anyhow::Result<()> { + let tx = { + let state = self.state.lock().await; + state.as_ref().and_then(|proxy| proxy.input_tx.clone()) + } + .ok_or_else(|| anyhow!("QUIC endpoint is not active"))?; + let role = match role { + WrappedTransportRole::Source => QuicProxyRole::Src, + WrappedTransportRole::Destination => QuicProxyRole::Dst, + }; + tx.try_send(QuicPacket::new( + QuicAddr::new(from_peer_id, role.outgoing()).into(), + payload.into(), + None, + None, + )) + .map_err(|error| anyhow!("failed to inject QUIC datagram: {error}")) + } + + async fn connect_source( + &self, + request: WrappedTransportConnect, + ) -> anyhow::Result> { + let connector = { + let state = self.state.lock().await; + state + .as_ref() + .and_then(|proxy| proxy.source_connector.clone()) + } + .ok_or_else(|| anyhow!("QUIC source endpoint is not prepared"))?; + let stream = connector + .connect_to_peer(request.dst_peer_id, request.src, request.dst) + .await?; + Ok(Box::new(stream)) + } + + async fn stop(&self) { + let mut state = self.state.lock().await; + if let Some(active) = state.as_mut() { + active.stop().await; + } + state.take(); } } diff --git a/easytier/src/gateway/socks5.rs b/easytier/src/gateway/socks5.rs deleted file mode 100644 index ebcda26b..00000000 --- a/easytier/src/gateway/socks5.rs +++ /dev/null @@ -1,1773 +0,0 @@ -use std::{ - any::Any, - net::{IpAddr, Ipv4Addr, SocketAddr}, - sync::{ - Arc, Weak, - atomic::{AtomicBool, AtomicUsize, Ordering}, - }, - time::Duration, -}; - -use crossbeam::atomic::AtomicCell; -#[cfg(feature = "kcp")] -use kcp_sys::{endpoint::KcpEndpoint, stream::KcpStream}; -use quanta::Instant; -use tokio_util::sync::{CancellationToken, DropGuard}; -use tokio_util::task::AbortOnDropHandle; - -#[cfg(feature = "kcp")] -use crate::gateway::kcp_proxy::NatDstKcpConnector; -use crate::{ - common::{config::PortForwardConfig, global_ctx::GlobalCtxEvent, join_joinset_background}, - gateway::{ - fast_socks5::{ - server::{ - AcceptAuthentication, AsyncTcpConnector, Config, SimpleUserPassword, Socks5Socket, - }, - util::stream::tcp_connect_with_timeout, - }, - ip_reassembler::IpReassembler, - tokio_smoltcp::{BufferSize, Net, NetConfig, channel_device}, - }, - tunnel::packet_def::{PacketType, ZCPacket}, -}; -use anyhow::Context; -use dashmap::{DashMap, mapref::entry::Entry}; -use pnet::packet::{ - Packet, ip::IpNextHeaderProtocols, ipv4::Ipv4Packet, tcp::TcpPacket, udp::UdpPacket, -}; -use tokio::{ - io::{AsyncRead, AsyncWrite}, - net::{TcpListener, UdpSocket}, - select, - sync::{Mutex, Notify, mpsc}, - task::JoinSet, - time::timeout, -}; - -#[cfg(feature = "kcp")] -use super::tcp_proxy::NatDstConnector as _; -use crate::tunnel::common::bind; -use crate::{ - common::{error::Error, global_ctx::GlobalCtx}, - peers::{PeerPacketFilter, peer_manager::PeerManager}, -}; - -#[cfg(feature = "ffi-dataplane")] -mod dataplane; - -#[cfg(feature = "ffi-dataplane")] -pub use dataplane::{DataPlaneTcpListener, DataPlaneTcpStream, DataPlaneUdpSocket}; - -enum SocksUdpSocket { - UdpSocket(Arc), - SmolUdpSocket(super::tokio_smoltcp::UdpSocket), -} - -impl SocksUdpSocket { - pub async fn send_to(&self, buf: &[u8], addr: SocketAddr) -> Result { - match self { - SocksUdpSocket::UdpSocket(socket) => socket.send_to(buf, addr).await, - SocksUdpSocket::SmolUdpSocket(socket) => socket.send_to(buf, addr).await, - } - } - - pub async fn recv_from(&self, buf: &mut [u8]) -> Result<(usize, SocketAddr), std::io::Error> { - match self { - SocksUdpSocket::UdpSocket(socket) => socket.recv_from(buf).await, - SocksUdpSocket::SmolUdpSocket(socket) => socket.recv_from(buf).await, - } - } -} - -enum SocksTcpStream { - Tcp(tokio::net::TcpStream), - SmolTcp(super::tokio_smoltcp::TcpStream), - #[cfg(feature = "kcp")] - Kcp(KcpStream), -} - -impl AsyncRead for SocksTcpStream { - fn poll_read( - self: std::pin::Pin<&mut Self>, - cx: &mut std::task::Context<'_>, - buf: &mut tokio::io::ReadBuf<'_>, - ) -> std::task::Poll> { - match self.get_mut() { - SocksTcpStream::Tcp(stream) => std::pin::Pin::new(stream).poll_read(cx, buf), - SocksTcpStream::SmolTcp(stream) => std::pin::Pin::new(stream).poll_read(cx, buf), - #[cfg(feature = "kcp")] - SocksTcpStream::Kcp(stream) => std::pin::Pin::new(stream).poll_read(cx, buf), - } - } -} - -impl AsyncWrite for SocksTcpStream { - fn poll_write( - self: std::pin::Pin<&mut Self>, - cx: &mut std::task::Context<'_>, - buf: &[u8], - ) -> std::task::Poll> { - match self.get_mut() { - SocksTcpStream::Tcp(stream) => std::pin::Pin::new(stream).poll_write(cx, buf), - SocksTcpStream::SmolTcp(stream) => std::pin::Pin::new(stream).poll_write(cx, buf), - #[cfg(feature = "kcp")] - SocksTcpStream::Kcp(stream) => std::pin::Pin::new(stream).poll_write(cx, buf), - } - } - - fn poll_flush( - self: std::pin::Pin<&mut Self>, - cx: &mut std::task::Context<'_>, - ) -> std::task::Poll> { - match self.get_mut() { - SocksTcpStream::Tcp(stream) => std::pin::Pin::new(stream).poll_flush(cx), - SocksTcpStream::SmolTcp(stream) => std::pin::Pin::new(stream).poll_flush(cx), - #[cfg(feature = "kcp")] - SocksTcpStream::Kcp(stream) => std::pin::Pin::new(stream).poll_flush(cx), - } - } - - fn poll_shutdown( - self: std::pin::Pin<&mut Self>, - cx: &mut std::task::Context<'_>, - ) -> std::task::Poll> { - match self.get_mut() { - SocksTcpStream::Tcp(stream) => std::pin::Pin::new(stream).poll_shutdown(cx), - SocksTcpStream::SmolTcp(stream) => std::pin::Pin::new(stream).poll_shutdown(cx), - #[cfg(feature = "kcp")] - SocksTcpStream::Kcp(stream) => std::pin::Pin::new(stream).poll_shutdown(cx), - } - } -} - -enum Socks5EntryData { - Tcp(TcpListener), // hold a binded socket to hold the tcp port - #[cfg(feature = "ffi-dataplane")] - // a data-plane routing entry that owns no resource. the entry_type in the - // key distinguishes a listen route from an actively outbound route. - DataPlaneRoute, - Udp((Arc, UdpClientKey)), // hold the socket to send data to dst -} - -const UDP_ENTRY: u8 = 1; -const TCP_ENTRY: u8 = 2; -#[cfg(feature = "ffi-dataplane")] -const TCP_LISTEN_ENTRY: u8 = 3; - -#[derive(Debug, Eq, PartialEq, Hash, Clone)] -struct Socks5Entry { - src: SocketAddr, - dst: SocketAddr, - entry_type: u8, -} - -type Socks5EntrySet = Arc>; - -fn increment_entry_count(entry_count: &AtomicUsize) -> (usize, usize) { - let old_entry_count = entry_count - .fetch_update(Ordering::Relaxed, Ordering::Relaxed, |count| { - count.checked_add(1) - }) - .unwrap_or_else(|count| count); - (old_entry_count, old_entry_count.saturating_add(1)) -} - -fn decrement_entry_count(entry_count: &AtomicUsize) -> (usize, usize) { - decrement_entry_count_by(entry_count, 1) -} - -fn decrement_entry_count_by(entry_count: &AtomicUsize, delta: usize) -> (usize, usize) { - if delta == 0 { - let current = entry_count.load(Ordering::Relaxed); - return (current, current); - } - - let old_entry_count = entry_count - .fetch_update(Ordering::Relaxed, Ordering::Relaxed, |count| { - Some(count.saturating_sub(delta)) - }) - .unwrap_or_else(|count| count); - (old_entry_count, old_entry_count.saturating_sub(delta)) -} - -fn insert_entry_and_increment_count( - entries: &Socks5EntrySet, - entry_count: &AtomicUsize, - entry: Socks5Entry, - data: Socks5EntryData, -) -> (bool, usize, usize) { - match entries.entry(entry) { - Entry::Occupied(mut occupied) => { - occupied.insert(data); - let current = entry_count.load(Ordering::Relaxed); - (true, current, current) - } - Entry::Vacant(vacant) => { - // Keep the count update inside the VacantEntry shard lock so bulk clear - // cannot observe the inserted entry before its count is reserved. - let (old_entry_count, new_entry_count) = increment_entry_count(entry_count); - vacant.insert(data); - (false, old_entry_count, new_entry_count) - } - } -} - -fn try_insert_entry_and_increment_count( - entries: &Socks5EntrySet, - entry_count: &AtomicUsize, - entry: Socks5Entry, - data: Socks5EntryData, -) -> bool { - match entries.entry(entry) { - Entry::Occupied(_) => false, - Entry::Vacant(vacant) => { - // See insert_entry_and_increment_count for why the count is reserved first. - increment_entry_count(entry_count); - vacant.insert(data); - true - } - } -} - -fn remove_entry_and_decrement_count( - entries: &Socks5EntrySet, - entry_count: &AtomicUsize, - entry: &Socks5Entry, -) -> (bool, usize, usize) { - let removed = entries.remove(entry).is_some(); - let (old_entry_count, new_entry_count) = if removed { - decrement_entry_count(entry_count) - } else { - let current = entry_count.load(Ordering::Relaxed); - (current, current) - }; - (removed, old_entry_count, new_entry_count) -} - -struct SmolTcpConnector { - net: Arc, - entries: Socks5EntrySet, - entry_count: Arc, - current_entry: std::sync::Mutex>, -} - -#[async_trait::async_trait] -impl AsyncTcpConnector for SmolTcpConnector { - type S = SocksTcpStream; - - async fn tcp_connect( - &self, - addr: SocketAddr, - timeout_s: u64, - ) -> crate::gateway::fast_socks5::Result { - let tmp_listener = TcpListener::bind("0.0.0.0:0").await?; - let local_addr = self.net.get_address(); - let port = tmp_listener.local_addr()?.port(); - - let entry = Socks5Entry { - src: SocketAddr::new(local_addr, port), - dst: addr, - entry_type: TCP_ENTRY, - }; - *self.current_entry.lock().unwrap() = Some(entry.clone()); - let (replaced, old_entry_count, new_entry_count) = insert_entry_and_increment_count( - &self.entries, - &self.entry_count, - entry.clone(), - Socks5EntryData::Tcp(tmp_listener), - ); - tracing::trace!( - ?entry, - replaced, - old_entry_count, - new_entry_count, - entries_len = self.entries.len(), - "socks5 inserted smoltcp tcp connector entry" - ); - - if addr.ip() == local_addr { - let modified_addr = - SocketAddr::new(IpAddr::V4(Ipv4Addr::new(127, 0, 0, 1)), addr.port()); - - Ok(SocksTcpStream::Tcp( - tcp_connect_with_timeout(modified_addr, timeout_s).await?, - )) - } else { - let remote_socket = timeout( - Duration::from_secs(timeout_s), - self.net.tcp_connect(addr, port), - ) - .await - .with_context(|| "connect to remote timeout")?; - - Ok(SocksTcpStream::SmolTcp(remote_socket.map_err(|e| { - super::fast_socks5::SocksError::Other(e.into()) - })?)) - } - } -} - -impl Drop for SmolTcpConnector { - fn drop(&mut self) { - if let Some(entry) = self.current_entry.lock().unwrap().take() { - tracing::debug!("drop smoltcp connector entry {:?}", entry); - let (removed, old_entry_count, new_entry_count) = - remove_entry_and_decrement_count(&self.entries, &self.entry_count, &entry); - tracing::trace!( - ?entry, - removed, - old_entry_count, - new_entry_count, - entries_len = self.entries.len(), - "socks5 removed smoltcp tcp connector entry" - ); - } - } -} - -#[cfg(feature = "kcp")] -struct Socks5KcpConnector { - kcp_endpoint: Weak, - peer_mgr: Weak, - src_addr: SocketAddr, -} - -#[cfg(feature = "kcp")] -#[async_trait::async_trait] -impl AsyncTcpConnector for Socks5KcpConnector { - type S = SocksTcpStream; - - async fn tcp_connect( - &self, - addr: SocketAddr, - _timeout_s: u64, - ) -> crate::gateway::fast_socks5::Result { - let Some(kcp_endpoint) = self.kcp_endpoint.upgrade() else { - return Err(anyhow::anyhow!("kcp endpoint is not ready").into()); - }; - let c = NatDstKcpConnector { - kcp_endpoint, - peer_mgr: self.peer_mgr.clone(), - }; - let ret = c - .connect(self.src_addr, addr) - .await - .map_err(super::fast_socks5::SocksError::Other)?; - Ok(SocksTcpStream::Kcp(ret)) - } -} - -struct Socks5AutoConnector { - #[cfg(feature = "kcp")] - kcp_endpoint: Option>, - peer_mgr: Weak, - entries: Socks5EntrySet, - entry_count: Arc, - smoltcp_net: Option>, - src_addr: SocketAddr, - - inner_connector: parking_lot::Mutex>>, -} - -#[async_trait::async_trait] -impl AsyncTcpConnector for Socks5AutoConnector { - type S = SocksTcpStream; - - async fn tcp_connect( - &self, - mut addr: SocketAddr, - timeout_s: u64, - ) -> crate::gateway::fast_socks5::Result { - if self.inner_connector.lock().is_some() { - return Err(anyhow::anyhow!("inner connector is already set").into()); - } - - let Some(peer_mgr_arc) = self.peer_mgr.upgrade() else { - tracing::error!("peer manager is dropped"); - return Err(anyhow::anyhow!("peer manager is dropped").into()); - }; - - if let Some(local_addr) = self.smoltcp_net.as_ref().map(|n| n.get_address()) - && local_addr == addr.ip() - { - addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::new(127, 0, 0, 1)), addr.port()); - } - - let has_smoltcp_net = self.smoltcp_net.is_some(); - let dst_peers = if has_smoltcp_net && !addr.ip().is_loopback() { - Some(peer_mgr_arc.get_msg_dst_peer(&addr.ip()).await.0) - } else { - None - }; - - if !has_smoltcp_net - || dst_peers.as_ref().is_some_and(Vec::is_empty) - || addr.ip().is_loopback() - { - // cannot find dst in virtual network, so try connect to dst directly - tracing::trace!( - ?addr, - src_addr = ?self.src_addr, - has_smoltcp_net, - dst_peer_count = dst_peers.as_ref().map(Vec::len), - is_loopback = addr.ip().is_loopback(), - "socks5 auto connector falling back to kernel tcp connect" - ); - return Ok(SocksTcpStream::Tcp( - tcp_connect_with_timeout(addr, timeout_s).await?, - )); - } - - let dst_allow_kcp = peer_mgr_arc.check_allow_kcp_to_dst(&addr.ip()).await; - tracing::debug!("dst_allow_kcp: {:?}", dst_allow_kcp); - - #[cfg(feature = "kcp")] - let connector: Box + Send> = - match (&self.kcp_endpoint, dst_allow_kcp) { - (Some(kcp_endpoint), true) => { - tracing::trace!( - ?addr, - src_addr = ?self.src_addr, - dst_peer_count = dst_peers.as_ref().map(Vec::len), - "socks5 auto connector selected kcp" - ); - Box::new(Socks5KcpConnector { - kcp_endpoint: kcp_endpoint.clone(), - peer_mgr: self.peer_mgr.clone(), - src_addr: self.src_addr, - }) - } - (_, _) => { - tracing::trace!( - ?addr, - src_addr = ?self.src_addr, - dst_peer_count = dst_peers.as_ref().map(Vec::len), - dst_allow_kcp, - has_kcp_endpoint = self.kcp_endpoint.is_some(), - "socks5 auto connector selected smoltcp" - ); - Box::new(SmolTcpConnector { - net: self.smoltcp_net.clone().unwrap(), - entries: self.entries.clone(), - entry_count: self.entry_count.clone(), - current_entry: std::sync::Mutex::new(None), - }) - } - }; - #[cfg(not(feature = "kcp"))] - let connector = { - tracing::trace!( - ?addr, - src_addr = ?self.src_addr, - dst_peer_count = dst_peers.as_ref().map(Vec::len), - "socks5 auto connector selected smoltcp" - ); - Box::new(SmolTcpConnector { - net: self.smoltcp_net.clone().unwrap(), - entries: self.entries.clone(), - entry_count: self.entry_count.clone(), - current_entry: std::sync::Mutex::new(None), - }) - }; - - let ret = connector.tcp_connect(addr, timeout_s).await; - self.inner_connector.lock().replace(Box::new(connector)); - ret - } -} - -struct Socks5ServerNet { - ipv4_addr: cidr::Ipv4Inet, - auth: Option, - - smoltcp_net: Arc, - forward_tasks: Arc>>, - - entries: Socks5EntrySet, -} - -impl Socks5ServerNet { - pub fn new( - ipv4_addr: cidr::Ipv4Inet, - auth: Option, - peer_manager: Weak, - packet_recv: Arc>>, - entries: Socks5EntrySet, - ) -> Self { - let mut forward_tasks = JoinSet::new(); - let mut cap = smoltcp::phy::DeviceCapabilities::default(); - cap.max_transmission_unit = 1284; // 1284 - 20 can be divided by 8 (fragment offset unit) - cap.medium = smoltcp::phy::Medium::Ip; - let (dev, stack_sink, mut stack_stream) = channel_device::ChannelDevice::new(cap); - - forward_tasks.spawn(async move { - let mut smoltcp_stack_receiver = packet_recv.lock().await; - while let Some(packet) = smoltcp_stack_receiver.recv().await { - tracing::trace!(?packet, "receive from peer send to smoltcp packet"); - if let Err(e) = stack_sink.send(Ok(packet.payload().to_vec())).await { - tracing::error!("send to smoltcp stack failed: {:?}", e); - } - } - tracing::warn!("smoltcp stack sink exited"); - }); - - forward_tasks.spawn(async move { - while let Some(data) = stack_stream.recv().await { - tracing::trace!( - ?data, - "receive from smoltcp stack and send to peer mgr packet, len = {}", - data.len() - ); - let Some(ipv4) = Ipv4Packet::new(&data) else { - tracing::error!(?data, "smoltcp stack stream get non ipv4 packet"); - continue; - }; - - let dst = ipv4.get_destination(); - let packet = ZCPacket::new_with_payload(&data); - let Some(peer_manager) = peer_manager.upgrade() else { - tracing::warn!("peer manager is gone, smoltcp sender exited"); - return; - }; - if let Err(e) = peer_manager - .send_msg_by_ip(packet, IpAddr::V4(dst), false) - .await - { - tracing::error!("send to peer failed in smoltcp sender: {:?}", e); - } - } - tracing::warn!("smoltcp stack stream exited"); - }); - - let interface_config = smoltcp::iface::Config::new(smoltcp::wire::HardwareAddress::Ip); - let net = Net::new( - dev, - NetConfig::new( - interface_config, - format!("{}/{}", ipv4_addr.address(), ipv4_addr.network_length()) - .parse() - .unwrap(), - vec![format!("{}", ipv4_addr.address()).parse().unwrap()], - Some(BufferSize { - tcp_rx_size: 1024 * 128, - tcp_tx_size: 1024 * 128, - ..Default::default() - }), - ), - ); - - let forward_tasks = Arc::new(std::sync::Mutex::new(forward_tasks)); - join_joinset_background(forward_tasks.clone(), "Socks5ServerNet".to_string()); - - Self { - ipv4_addr, - auth, - - smoltcp_net: Arc::new(net), - forward_tasks, - - entries, - } - } - - async fn handle_tcp_stream_task(stream: tokio::net::TcpStream, connector: Socks5AutoConnector) { - let mut config = Config::::default(); - config.set_request_timeout(10); - config.set_skip_auth(false); - config.set_allow_no_auth(true); - - let socket = Socks5Socket::new(stream, Arc::new(config), connector); - - match socket.upgrade_to_socks5().await { - Ok(_) => { - tracing::info!("socks5 handle success"); - } - Err(e) => { - tracing::error!("socks5 handshake failed: {:?}", e); - } - }; - } - - fn handle_tcp_stream(&self, stream: tokio::net::TcpStream, connector: Socks5AutoConnector) { - self.forward_tasks - .lock() - .unwrap() - .spawn(Self::handle_tcp_stream_task(stream, connector)); - } -} - -struct UdpClientInfo { - client_addr: SocketAddr, - port_holder_socket: Arc, - local_addr: SocketAddr, - last_active: AtomicCell, - entries: Socks5EntrySet, - entry_key: Socks5Entry, -} - -#[derive(Debug, Eq, PartialEq, Hash, Clone)] -struct UdpClientKey { - client_addr: SocketAddr, - dst_addr: SocketAddr, -} - -pub struct Socks5Server { - global_ctx: Arc, - peer_manager: Weak, - auth: Option, - - tasks: Arc>>, - packet_sender: mpsc::Sender, - packet_recv: Arc>>, - - net: Arc>>, - entries: Socks5EntrySet, - - udp_client_map: Arc>>, - udp_forward_task: Arc>>, - - #[cfg(feature = "kcp")] - kcp_endpoint: Mutex>>, - - socks5_enabled: Arc, - #[cfg(feature = "ffi-dataplane")] - data_plane_refs: Arc, - // Tracks whether the smoltcp `net` is ready for data-plane callers. - #[cfg(feature = "ffi-dataplane")] - data_plane_net_ready: tokio::sync::watch::Sender, - cancel_tokens: Arc>, - port_forward_list_change_notifier: Arc, - entry_count: Arc, -} - -#[async_trait::async_trait] -impl PeerPacketFilter for Socks5Server { - async fn try_process_packet_from_peer(&self, packet: ZCPacket) -> Option { - let entry_count = self.entry_count.load(Ordering::Relaxed); - let socks5_enabled = self.socks5_enabled.load(Ordering::Relaxed); - if entry_count == 0 && !socks5_enabled && self.entries.is_empty() { - if tracing::enabled!(tracing::Level::TRACE) - && let Some(hdr) = packet.peer_manager_header() - && matches!( - hdr.packet_type, - x if x == PacketType::Data as u8 - || x == PacketType::DataWithKcpSrcModified as u8 - || x == PacketType::DataWithQuicSrcModified as u8 - ) - { - if let Some(ipv4) = Ipv4Packet::new(packet.payload()) { - let (tcp_src_port, tcp_dst_port, tcp_flags) = - if ipv4.get_next_level_protocol() == IpNextHeaderProtocols::Tcp { - TcpPacket::new(ipv4.payload()) - .map(|tcp| { - ( - Some(tcp.get_source()), - Some(tcp.get_destination()), - Some(tcp.get_flags()), - ) - }) - .unwrap_or((None, None, None)) - } else { - (None, None, None) - }; - tracing::trace!( - packet_type = hdr.packet_type, - from_peer_id = hdr.from_peer_id.get(), - to_peer_id = hdr.to_peer_id.get(), - ipv4_src = %ipv4.get_source(), - ipv4_dst = %ipv4.get_destination(), - next_protocol = ?ipv4.get_next_level_protocol(), - ?tcp_src_port, - ?tcp_dst_port, - ?tcp_flags, - entry_count, - socks5_enabled, - "socks5 fast gate passed packet from peer" - ); - } else { - tracing::trace!( - packet_type = hdr.packet_type, - from_peer_id = hdr.from_peer_id.get(), - to_peer_id = hdr.to_peer_id.get(), - entry_count, - socks5_enabled, - "socks5 fast gate passed non-ipv4 packet from peer" - ); - } - } - return Some(packet); - } - let hdr = packet.peer_manager_header().unwrap(); - let is_modified_src_packet = matches!( - hdr.packet_type, - x if x == PacketType::DataWithKcpSrcModified as u8 - || x == PacketType::DataWithQuicSrcModified as u8 - ); - if hdr.packet_type != PacketType::Data as u8 && !is_modified_src_packet { - return Some(packet); - } - if is_modified_src_packet && hdr.from_peer_id != hdr.to_peer_id { - tracing::trace!( - packet_type = hdr.packet_type, - from_peer_id = hdr.from_peer_id.get(), - to_peer_id = hdr.to_peer_id.get(), - "socks5 passed non-loopback modified-source packet from peer" - ); - return Some(packet); - } - - let payload_bytes = packet.payload(); - - let Some(ipv4) = Ipv4Packet::new(payload_bytes) else { - return Some(packet); - }; - if ipv4.get_version() != 4 { - return Some(packet); - } - - let (entry_key, tcp_flags) = match ipv4.get_next_level_protocol() { - IpNextHeaderProtocols::Tcp => { - let Some(tcp_packet) = TcpPacket::new(ipv4.payload()) else { - return Some(packet); - }; - let entry = Socks5Entry { - dst: SocketAddr::new(ipv4.get_source().into(), tcp_packet.get_source()), - src: SocketAddr::new( - ipv4.get_destination().into(), - tcp_packet.get_destination(), - ), - entry_type: TCP_ENTRY, - }; - #[cfg(feature = "ffi-dataplane")] - let entry = if self.entries.contains_key(&entry) { - // Case 1: it is an established connection that has an exactly matched inbound. - entry - } else { - // Case 2: it could be a new TCP SYN packet that has not been accepted. - Socks5Entry { - src: entry.src, - dst: SocketAddr::new(IpAddr::V4(Ipv4Addr::UNSPECIFIED), 0), - entry_type: TCP_LISTEN_ENTRY, - } - }; - (entry, Some(tcp_packet.get_flags())) - } - - IpNextHeaderProtocols::Udp => { - if IpReassembler::is_packet_fragmented(&ipv4) { - let ipv4_src: IpAddr = ipv4.get_source().into(); - // only send to smoltcp if the ipv4 src is in the entries - let is_in_entries = self.entries.iter().any(|x| x.key().dst.ip() == ipv4_src); - tracing::trace!( - ?is_in_entries, - "ipv4 src = {:?}, check need send both smoltcp and kernel tun", - ipv4_src - ); - if is_in_entries { - // if the packet is fragmented, no matther what the payload is, need send it to both smoltcp and kernel tun. because - // we cannot determine the udp port of the packet. - match self.packet_sender.try_send(packet.clone()) { - Ok(()) => tracing::trace!( - ?ipv4_src, - entry_count = self.entry_count.load(Ordering::Relaxed), - "socks5 delivered fragmented packet from peer to smoltcp" - ), - Err(err) => tracing::trace!( - ?ipv4_src, - ?err, - entry_count = self.entry_count.load(Ordering::Relaxed), - "socks5 failed to deliver fragmented packet from peer to smoltcp" - ), - } - } - return Some(packet); - } - - let Some(udp_packet) = UdpPacket::new(ipv4.payload()) else { - return Some(packet); - }; - ( - Socks5Entry { - dst: SocketAddr::new(ipv4.get_source().into(), udp_packet.get_source()), - src: SocketAddr::new( - ipv4.get_destination().into(), - udp_packet.get_destination(), - ), - entry_type: UDP_ENTRY, - }, - None, - ) - } - _ => { - return Some(packet); - } - }; - - if !self.entries.contains_key(&entry_key) { - tracing::trace!( - ?entry_key, - ?tcp_flags, - ipv4_src = %ipv4.get_source(), - ipv4_dst = %ipv4.get_destination(), - entry_count = self.entry_count.load(Ordering::Relaxed), - socks5_enabled = self.socks5_enabled.load(Ordering::Relaxed), - "socks5 no entry for packet from peer" - ); - return Some(packet); - } - - tracing::trace!( - ?entry_key, - ?tcp_flags, - ?ipv4, - entry_count = self.entry_count.load(Ordering::Relaxed), - "socks5 found entry for packet from peer" - ); - - match self.packet_sender.try_send(packet) { - Ok(()) => tracing::trace!( - ?entry_key, - ?tcp_flags, - entry_count = self.entry_count.load(Ordering::Relaxed), - "socks5 delivered packet from peer to smoltcp" - ), - Err(err) => tracing::trace!( - ?entry_key, - ?tcp_flags, - ?err, - entry_count = self.entry_count.load(Ordering::Relaxed), - "socks5 failed to deliver packet from peer to smoltcp" - ), - } - - None - } -} - -impl Socks5Server { - pub fn new( - global_ctx: Arc, - peer_manager: Arc, - auth: Option, - ) -> Arc { - let (packet_sender, packet_recv) = mpsc::channel(1024); - Arc::new(Self { - global_ctx, - peer_manager: Arc::downgrade(&peer_manager), - auth, - - tasks: Arc::new(std::sync::Mutex::new(JoinSet::new())), - packet_recv: Arc::new(Mutex::new(packet_recv)), - packet_sender, - - net: Arc::new(Mutex::new(None)), - entries: Arc::new(DashMap::new()), - - udp_client_map: Arc::new(DashMap::new()), - udp_forward_task: Arc::new(DashMap::new()), - - #[cfg(feature = "kcp")] - kcp_endpoint: Mutex::new(None), - - socks5_enabled: Arc::new(AtomicBool::new(false)), - #[cfg(feature = "ffi-dataplane")] - data_plane_refs: Arc::new(AtomicUsize::new(0)), - #[cfg(feature = "ffi-dataplane")] - data_plane_net_ready: tokio::sync::watch::channel(false).0, - cancel_tokens: Arc::new(DashMap::new()), - port_forward_list_change_notifier: Arc::new(Notify::new()), - entry_count: Arc::new(AtomicUsize::new(0)), - }) - } - - async fn run_net_update_task(self: &Arc) { - let net = self.net.clone(); - let global_ctx = self.global_ctx.clone(); - let peer_manager = self.peer_manager.clone(); - let packet_recv = self.packet_recv.clone(); - let entries = self.entries.clone(); - let entry_count = self.entry_count.clone(); - let udp_client_map = self.udp_client_map.clone(); - let cancel_tokens = self.cancel_tokens.clone(); - let port_forward_list_change_notifier = self.port_forward_list_change_notifier.clone(); - let socks5_enabled = self.socks5_enabled.clone(); - #[cfg(feature = "ffi-dataplane")] - let data_plane_refs = self.data_plane_refs.clone(); - #[cfg(feature = "ffi-dataplane")] - let data_plane_net_ready = self.data_plane_net_ready.clone(); - self.tasks.lock().unwrap().spawn(async move { - let mut prev_ipv4 = None; - loop { - #[cfg(feature = "ffi-dataplane")] - let data_plane_active = data_plane_refs.load(Ordering::Relaxed) > 0; - #[cfg(not(feature = "ffi-dataplane"))] - let data_plane_active = false; - - let active_port_forwards = cancel_tokens.len(); - let is_socks5_enabled = socks5_enabled.load(Ordering::Relaxed); - if active_port_forwards == 0 && !is_socks5_enabled && !data_plane_active { - let had_net = { - let mut net_guard = net.lock().await; - net_guard.take().is_some() - }; - tracing::trace!( - had_net, - active_port_forwards, - is_socks5_enabled, - data_plane_active, - entry_count = entry_count.load(Ordering::Relaxed), - entries_len = entries.len(), - "socks5 net update waiting for consumers" - ); - #[cfg(feature = "ffi-dataplane")] - let _ = data_plane_net_ready.send_replace(false); - port_forward_list_change_notifier.notified().await; - continue; - } - - let mut event_recv = global_ctx.subscribe(); - - let cur_ipv4 = global_ctx.get_ipv4(); - if prev_ipv4 != cur_ipv4 { - let old_ipv4 = prev_ipv4; - prev_ipv4 = cur_ipv4; - - tracing::trace!( - ?old_ipv4, - ?cur_ipv4, - old_entry_count = entry_count.load(Ordering::Relaxed), - old_entries_len = entries.len(), - udp_client_count = udp_client_map.len(), - "socks5 net update resetting entries for ipv4 change" - ); - let mut removed_entries = 0; - entries.retain(|_, _| { - removed_entries += 1; - false - }); - let (_, new_entry_count) = - decrement_entry_count_by(&entry_count, removed_entries); - udp_client_map.clear(); - tracing::trace!( - ?old_ipv4, - ?cur_ipv4, - removed_entries, - new_entry_count, - new_entries_len = entries.len(), - udp_client_count = udp_client_map.len(), - "socks5 net update reset entries complete" - ); - - if let Some(cur_ipv4) = cur_ipv4 { - net.lock().await.replace(Socks5ServerNet::new( - cur_ipv4, - None, - peer_manager.clone(), - packet_recv.clone(), - entries.clone(), - )); - tracing::trace!( - ?cur_ipv4, - entry_count = entry_count.load(Ordering::Relaxed), - entries_len = entries.len(), - "socks5 net update installed smoltcp net" - ); - // Wake any data-plane callers waiting in - // `wait_data_plane_net` for the smoltcp net to appear. - #[cfg(feature = "ffi-dataplane")] - let _ = data_plane_net_ready.send_replace(true); - } else { - let _ = net.lock().await.take(); - tracing::trace!( - entry_count = entry_count.load(Ordering::Relaxed), - entries_len = entries.len(), - "socks5 net update removed smoltcp net" - ); - #[cfg(feature = "ffi-dataplane")] - let _ = data_plane_net_ready.send_replace(false); - } - } - - select! { - _ = event_recv.recv() => {} - _ = tokio::time::sleep(Duration::from_secs(120)) => {} - } - } - }); - } - - pub async fn run( - self: &Arc, - #[cfg(feature = "kcp")] kcp_endpoint: Option>, - ) -> Result<(), Error> { - #[cfg(feature = "kcp")] - { - *self.kcp_endpoint.lock().await = kcp_endpoint.clone(); - } - if let Some(proxy_url) = self.global_ctx.config.get_socks5_portal() { - let bind_addr = format!( - "{}:{}", - proxy_url.host_str().unwrap(), - proxy_url.port().unwrap() - ); - - let listener = bind::() - .addr(bind_addr.parse::().unwrap()) - .net_ns(self.global_ctx.net_ns.clone()) - .call()?; - - let entries = self.entries.clone(); - let entry_count = self.entry_count.clone(); - let peer_manager = self.peer_manager.clone(); - let net = self.net.clone(); - self.tasks.lock().unwrap().spawn(async move { - loop { - match listener.accept().await { - Ok((socket, addr)) => { - tracing::info!("accept a new connection, {:?}", socket); - let connector = Socks5AutoConnector { - smoltcp_net: net - .lock() - .await - .as_ref() - .map(|net| net.smoltcp_net.clone()), - entries: entries.clone(), - #[cfg(feature = "kcp")] - kcp_endpoint: kcp_endpoint.clone(), - peer_mgr: peer_manager.clone(), - src_addr: addr, - inner_connector: parking_lot::Mutex::new(None), - entry_count: entry_count.clone(), - }; - if let Some(net) = net.lock().await.as_ref() { - net.handle_tcp_stream(socket, connector); - } else { - tokio::spawn(Socks5ServerNet::handle_tcp_stream_task( - socket, connector, - )); - } - } - Err(err) => tracing::error!("accept error = {:?}", err), - } - } - }); - - self.socks5_enabled.store(true, Ordering::Relaxed); - join_joinset_background(self.tasks.clone(), "socks5 server".to_string()); - }; - - let cfgs = self.global_ctx.config.get_port_forwards(); - self.reload_port_forwards(&cfgs).await?; - - let Some(peer_manager) = self.peer_manager.upgrade() else { - return Err(anyhow::anyhow!("peer manager is gone").into()); - }; - peer_manager - .add_packet_process_pipeline(Box::new(self.clone())) - .await; - tracing::trace!( - cfg_count = cfgs.len(), - cancel_token_count = self.cancel_tokens.len(), - entry_count = self.entry_count.load(Ordering::Relaxed), - entries_len = self.entries.len(), - "socks5 peer packet pipeline registered" - ); - - self.run_net_update_task().await; - - Ok(()) - } - - pub async fn reload_port_forwards(&self, cfgs: &Vec) -> Result<(), Error> { - // remove entries not in new cfg - self.cancel_tokens.retain(|k, _| { - cfgs.iter().any(|cfg| { - if cfg.dst_addr.ip().is_unspecified() { - k.bind_addr == cfg.bind_addr && k.proto == cfg.proto - } else { - k == cfg - } - }) - }); - // add new ones - for cfg in cfgs { - if !self.cancel_tokens.contains_key(cfg) { - self.add_port_forward(cfg.clone()).await?; - } - } - self.port_forward_list_change_notifier.notify_one(); - Ok(()) - } - - async fn handle_port_forward_connection( - mut incoming_socket: tokio::net::TcpStream, - connector: Box + Send>, - dst_addr: SocketAddr, - ) { - tracing::trace!(?dst_addr, "port forward: connecting to destination"); - let outgoing_socket = match connector.tcp_connect(dst_addr, 10).await { - Ok(socket) => socket, - Err(e) => { - tracing::error!("port forward: failed to connect to destination: {:?}", e); - return; - } - }; - tracing::trace!(?dst_addr, "port forward: connected to destination"); - - let mut outgoing_socket = outgoing_socket; - match tokio::io::copy_bidirectional(&mut incoming_socket, &mut outgoing_socket).await { - Ok((from_client, from_server)) => { - tracing::info!( - "port forward connection finished: client->server: {} bytes, server->client: {} bytes", - from_client, - from_server - ); - } - Err(e) => { - tracing::error!("port forward connection error: {:?}", e); - } - } - } - - pub async fn add_port_forward(&self, cfg: PortForwardConfig) -> Result<(), Error> { - match cfg.proto.to_lowercase().as_str() { - "tcp" => { - self.add_tcp_port_forward(&cfg).await?; - } - "udp" => { - self.add_udp_port_forward(&cfg).await?; - } - _ => { - return Err(anyhow::anyhow!( - "unsupported protocol: {}, only support udp / tcp", - cfg.proto - ) - .into()); - } - } - self.global_ctx - .issue_event(GlobalCtxEvent::PortForwardAdded(cfg.clone().into())); - Ok(()) - } - - pub fn remove_port_forward(&self, cfg: PortForwardConfig) { - let _ = self.cancel_tokens.remove(&cfg); - } - - pub async fn add_tcp_port_forward(&self, cfg: &PortForwardConfig) -> Result<(), Error> { - let (bind_addr, dst_addr) = (cfg.bind_addr, cfg.dst_addr); - let listener = bind::() - .addr(bind_addr) - .net_ns(self.global_ctx.net_ns.clone()) - .call()?; - - let net = self.net.clone(); - let entries = self.entries.clone(); - let entry_count = self.entry_count.clone(); - let tasks = Arc::new(std::sync::Mutex::new(JoinSet::new())); - join_joinset_background(tasks.clone(), "tcp port forward".to_string()); - let forward_tasks = tasks; - #[cfg(feature = "kcp")] - let kcp_endpoint = self.kcp_endpoint.lock().await.clone(); - let peer_mgr = self.peer_manager.clone(); - let cancel_token = CancellationToken::new(); - self.cancel_tokens - .insert(cfg.clone(), cancel_token.clone().drop_guard()); - - self.tasks.lock().unwrap().spawn(async move { - loop { - let (incoming_socket, addr) = select! { - biased; - _ = cancel_token.cancelled() => { - tracing::info!("port forward for {:?} cancelled", bind_addr); - break; - } - res = listener.accept() => { - match res { - Ok(result) => result, - Err(err) => { - tracing::error!("port forward accept error = {:?}", err); - continue; - } - } - } - }; - - tracing::info!( - "port forward: accept new connection from {:?} to {:?}", - bind_addr, - dst_addr - ); - - let (smoltcp_net, net_ipv4) = { - let net_guard = net.lock().await; - ( - net_guard.as_ref().map(|net| net.smoltcp_net.clone()), - net_guard.as_ref().map(|net| net.ipv4_addr), - ) - }; - tracing::trace!( - ?bind_addr, - ?dst_addr, - client_addr = ?addr, - has_smoltcp_net = smoltcp_net.is_some(), - ?net_ipv4, - entry_count = entry_count.load(Ordering::Relaxed), - entries_len = entries.len(), - "port forward: preparing connector" - ); - - let connector = Socks5AutoConnector { - #[cfg(feature = "kcp")] - kcp_endpoint: kcp_endpoint.clone(), - peer_mgr: peer_mgr.clone(), - entries: entries.clone(), - smoltcp_net, - src_addr: addr, - entry_count: entry_count.clone(), - inner_connector: parking_lot::Mutex::new(None), - }; - - forward_tasks - .lock() - .unwrap() - .spawn(Self::handle_port_forward_connection( - incoming_socket, - Box::new(connector), - dst_addr, - )); - } - }); - - Ok(()) - } - - #[tracing::instrument(name = "add_udp_port_forward", skip(self))] - pub async fn add_udp_port_forward(&self, cfg: &PortForwardConfig) -> Result<(), Error> { - let (bind_addr, dst_addr) = (cfg.bind_addr, cfg.dst_addr); - let socket = Arc::new( - bind::() - .addr(bind_addr) - .net_ns(self.global_ctx.net_ns.clone()) - .call()?, - ); - - let entries = self.entries.clone(); - let entry_count = self.entry_count.clone(); - let net_ns = self.global_ctx.net_ns.clone(); - let net = self.net.clone(); - let udp_client_map = self.udp_client_map.clone(); - let udp_forward_task = self.udp_forward_task.clone(); - let cancel_token = CancellationToken::new(); - self.cancel_tokens - .insert(cfg.clone(), cancel_token.clone().drop_guard()); - - self.tasks.lock().unwrap().spawn(async move { - loop { - // we set the max buffer size of smoltcp to 8192, so we need to use a buffer size that is less than 8192 here. - let mut buf = vec![0u8; 8192]; - let (len, addr) = select! { - biased; - _ = cancel_token.cancelled() => { - tracing::info!("udp port forward for {:?} cancelled", bind_addr); - break; - } - res = socket.recv_from(&mut buf) => { - match res { - Ok(result) => result, - Err(err) => { - tracing::error!("udp port forward recv error = {:?}", err); - continue; - } - } - } - }; - - tracing::trace!( - "udp port forward recv packet from {:?}, len = {}", - addr, - len - ); - - let udp_client_key = UdpClientKey { - client_addr: addr, - dst_addr, - }; - - let binded_socket = udp_client_map.get(&udp_client_key); - let client_info = match binded_socket { - Some(s) => s.clone(), - None => { - let _g = net_ns.guard(); - // reserve a port so os will not use it to connect to the virtual network - let binded_socket = tokio::net::UdpSocket::bind("0.0.0.0:0").await; - if binded_socket.is_err() { - tracing::error!("udp port forward bind error = {:?}", binded_socket); - continue; - } - let binded_socket = binded_socket.unwrap(); - let mut local_addr = binded_socket.local_addr().unwrap(); - let Some(cur_ipv4) = net.lock().await.as_ref().map(|net| net.ipv4_addr) else { - continue; - }; - local_addr.set_ip(cur_ipv4.address().into()); - - let entry_key = Socks5Entry { - src: local_addr, - dst: dst_addr, - entry_type: UDP_ENTRY, - }; - - tracing::debug!("udp port forward binded socket = {:?}, entry_key = {:?}", local_addr, entry_key); - - let client_info = Arc::new(UdpClientInfo { - client_addr: addr, - port_holder_socket: Arc::new(binded_socket), - local_addr, - last_active: AtomicCell::new(Instant::now()), - entries: entries.clone(), - entry_key, - }); - udp_client_map.insert(udp_client_key.clone(), client_info.clone()); - client_info - } - }; - - client_info.last_active.store(Instant::now()); - - let entry_data = match entries.get(&client_info.entry_key) { - Some(data) => data, - None => { - let guard = net.lock().await; - let Some(net) = guard.as_ref() else { - continue; - }; - let local_addr = net.ipv4_addr; - let sokcs_udp = if dst_addr.ip() == local_addr.address() { - SocksUdpSocket::UdpSocket(client_info.port_holder_socket.clone()) - } else { - tracing::debug!("udp port forward bind new smol udp socket, {:?}", local_addr); - SocksUdpSocket::SmolUdpSocket( - net.smoltcp_net - .udp_bind(SocketAddr::new( - IpAddr::V4(local_addr.address()), - client_info.local_addr.port(), - )) - .await - .unwrap(), - ) - }; - let socks_udp = Arc::new(sokcs_udp); - insert_entry_and_increment_count( - &entries, - &entry_count, - client_info.entry_key.clone(), - Socks5EntryData::Udp((socks_udp.clone(), udp_client_key.clone())), - ); - - let socks = socket.clone(); - let client_addr = addr; - udp_forward_task.insert( - udp_client_key.clone(), - AbortOnDropHandle::new(tokio::spawn(async move { - loop { - let mut buf = vec![0u8; 8192]; - match socks_udp.recv_from(&mut buf).await { - Ok((len, dst_addr)) => { - tracing::trace!( - "udp port forward recv response packet from {:?}, len = {}, client_addr = {:?}", - dst_addr, - len, - client_addr - ); - if let Err(e) = socks.send_to(&buf[..len], client_addr).await { - tracing::error!("udp forward send error = {:?}", e); - } - } - Err(e) => { - tracing::error!("udp forward recv error = {:?}", e); - } - } - } - })), - ); - - entries.get(&client_info.entry_key).unwrap() - } - }; - - let s = match entry_data.value() { - Socks5EntryData::Udp((s, _)) => s.clone(), - _ => { - panic!("udp entry data is not udp entry data"); - } - }; - drop(entry_data); - - if let Err(e) = s.send_to(&buf[..len], dst_addr).await { - tracing::error!(?dst_addr, ?len, "udp port forward send error = {:?}", e); - } else { - tracing::trace!(?dst_addr, ?len, "udp port forward send packet success"); - } - } - }); - - // clean up task - let udp_client_map = self.udp_client_map.clone(); - let udp_forward_task = self.udp_forward_task.clone(); - let entries = self.entries.clone(); - let entry_count = self.entry_count.clone(); - let cancel_tokens = self.cancel_tokens.clone(); - self.tasks.lock().unwrap().spawn(async move { - loop { - tokio::time::sleep(Duration::from_secs(30)).await; - let now = Instant::now(); - udp_client_map.retain(|_, client_info| { - now.duration_since(client_info.last_active.load()).as_secs() < 600 - }); - udp_forward_task.retain(|k, _| udp_client_map.contains_key(k)); - let mut removed_entries = 0; - entries.retain(|_, data| match data { - Socks5EntryData::Udp((_, udp_client_key)) => { - let keep = udp_client_map.contains_key(udp_client_key); - if !keep { - removed_entries += 1; - } - keep - } - _ => true, - }); - decrement_entry_count_by(&entry_count, removed_entries); - - udp_client_map.shrink_to_fit(); - udp_forward_task.shrink_to_fit(); - entries.shrink_to_fit(); - cancel_tokens.shrink_to_fit(); - } - }); - - Ok(()) - } -} - -#[cfg(test)] -mod tests { - use std::net::{IpAddr, Ipv4Addr, SocketAddr}; - - use pnet::packet::{ - MutablePacket, - ip::IpNextHeaderProtocols, - ipv4::{self, MutableIpv4Packet}, - tcp::{self, MutableTcpPacket, TcpFlags}, - }; - - use super::*; - use crate::peers::tests::create_mock_peer_manager; - - fn build_tcp_packet(src: SocketAddr, dst: SocketAddr) -> Vec { - let mut buf = vec![0u8; 40]; - let src_ip = match src.ip() { - IpAddr::V4(ip) => ip, - IpAddr::V6(_) => panic!("test only supports ipv4"), - }; - let dst_ip = match dst.ip() { - IpAddr::V4(ip) => ip, - IpAddr::V6(_) => panic!("test only supports ipv4"), - }; - - { - let mut ip_packet = MutableIpv4Packet::new(&mut buf).unwrap(); - ip_packet.set_version(4); - ip_packet.set_header_length(5); - ip_packet.set_total_length(40); - ip_packet.set_ttl(64); - ip_packet.set_next_level_protocol(IpNextHeaderProtocols::Tcp); - ip_packet.set_source(src_ip); - ip_packet.set_destination(dst_ip); - - let mut tcp_packet = MutableTcpPacket::new(ip_packet.payload_mut()).unwrap(); - tcp_packet.set_source(src.port()); - tcp_packet.set_destination(dst.port()); - tcp_packet.set_data_offset(5); - tcp_packet.set_flags(TcpFlags::SYN | TcpFlags::ACK); - tcp_packet.set_window(65535); - tcp_packet.set_checksum(tcp::ipv4_checksum( - &tcp_packet.to_immutable(), - &src_ip, - &dst_ip, - )); - - ip_packet.set_checksum(ipv4::checksum(&ip_packet.to_immutable())); - } - - buf - } - - fn build_udp_followup_fragment(src: Ipv4Addr, dst: Ipv4Addr) -> Vec { - let mut buf = vec![0u8; 28]; - { - let mut ip_packet = MutableIpv4Packet::new(&mut buf).unwrap(); - ip_packet.set_version(4); - ip_packet.set_header_length(5); - ip_packet.set_total_length(28); - ip_packet.set_ttl(64); - ip_packet.set_next_level_protocol(IpNextHeaderProtocols::Udp); - ip_packet.set_fragment_offset(1); - ip_packet.set_source(src); - ip_packet.set_destination(dst); - ip_packet - .payload_mut() - .copy_from_slice(&[0xde, 0xad, 0xbe, 0xef, 0xca, 0xfe, 0xba, 0xbe]); - - ip_packet.set_checksum(ipv4::checksum(&ip_packet.to_immutable())); - } - - buf - } - - #[tokio::test] - async fn socks5_consumes_modified_data_when_entry_matches() { - let peer_manager = create_mock_peer_manager().await; - let server = Socks5Server::new(peer_manager.get_global_ctx(), peer_manager, None); - - let local = SocketAddr::new(IpAddr::V4(Ipv4Addr::new(10, 144, 144, 1)), 40000); - let remote = SocketAddr::new(IpAddr::V4(Ipv4Addr::new(10, 144, 144, 3)), 22); - let entry = Socks5Entry { - src: local, - dst: remote, - entry_type: TCP_ENTRY, - }; - let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); - insert_entry_and_increment_count( - &server.entries, - &server.entry_count, - entry, - Socks5EntryData::Tcp(listener), - ); - - for packet_type in [ - PacketType::DataWithKcpSrcModified, - PacketType::DataWithQuicSrcModified, - ] { - let mut packet = ZCPacket::new_with_payload(&build_tcp_packet(remote, local)); - packet.fill_peer_manager_hdr(1, 1, packet_type as u8); - - let result = server.try_process_packet_from_peer(packet).await; - assert!(result.is_none()); - - let mut receiver = server.packet_recv.lock().await; - let received = receiver.try_recv().unwrap(); - assert_eq!( - received.peer_manager_header().unwrap().packet_type, - packet_type as u8 - ); - } - } - - #[tokio::test] - async fn socks5_passes_through_unmatched_or_malformed_modified_data() { - let peer_manager = create_mock_peer_manager().await; - let server = Socks5Server::new(peer_manager.get_global_ctx(), peer_manager, None); - let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); - insert_entry_and_increment_count( - &server.entries, - &server.entry_count, - Socks5Entry { - src: SocketAddr::new(IpAddr::V4(Ipv4Addr::new(10, 144, 144, 1)), 40000), - dst: SocketAddr::new(IpAddr::V4(Ipv4Addr::new(10, 144, 144, 3)), 22), - entry_type: TCP_ENTRY, - }, - Socks5EntryData::Tcp(listener), - ); - - let unmatched_local = SocketAddr::new(IpAddr::V4(Ipv4Addr::new(10, 144, 144, 1)), 40001); - let remote = SocketAddr::new(IpAddr::V4(Ipv4Addr::new(10, 144, 144, 3)), 22); - let mut unmatched_packet = - ZCPacket::new_with_payload(&build_tcp_packet(remote, unmatched_local)); - unmatched_packet.fill_peer_manager_hdr(1, 2, PacketType::DataWithKcpSrcModified as u8); - let result = server.try_process_packet_from_peer(unmatched_packet).await; - assert!(result.is_some()); - - let mut malformed_packet = ZCPacket::new_with_payload(&[0u8; 8]); - malformed_packet.fill_peer_manager_hdr(1, 2, PacketType::DataWithQuicSrcModified as u8); - let result = server.try_process_packet_from_peer(malformed_packet).await; - assert!(result.is_some()); - - let mut receiver = server.packet_recv.lock().await; - assert!(receiver.try_recv().is_err()); - } - - #[tokio::test] - async fn socks5_passes_through_non_loopback_modified_data_even_when_entry_matches() { - let peer_manager = create_mock_peer_manager().await; - let server = Socks5Server::new(peer_manager.get_global_ctx(), peer_manager, None); - - let local = SocketAddr::new(IpAddr::V4(Ipv4Addr::new(10, 144, 144, 1)), 40000); - let remote = SocketAddr::new(IpAddr::V4(Ipv4Addr::new(10, 144, 144, 3)), 22); - let entry = Socks5Entry { - src: local, - dst: remote, - entry_type: TCP_ENTRY, - }; - let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); - insert_entry_and_increment_count( - &server.entries, - &server.entry_count, - entry, - Socks5EntryData::Tcp(listener), - ); - - let mut packet = ZCPacket::new_with_payload(&build_tcp_packet(remote, local)); - packet.fill_peer_manager_hdr(1, 2, PacketType::DataWithKcpSrcModified as u8); - - let result = server.try_process_packet_from_peer(packet).await; - assert!(result.is_some()); - - let mut receiver = server.packet_recv.lock().await; - assert!(receiver.try_recv().is_err()); - } - - #[tokio::test] - async fn socks5_mirrors_fragmented_udp_even_when_entry_count_is_stale_zero() { - let peer_manager = create_mock_peer_manager().await; - let server = Socks5Server::new(peer_manager.get_global_ctx(), peer_manager, None); - - let local = SocketAddr::new(IpAddr::V4(Ipv4Addr::new(10, 144, 144, 1)), 40000); - let remote = SocketAddr::new(IpAddr::V4(Ipv4Addr::new(10, 144, 144, 3)), 53); - let udp_socket = Arc::new(tokio::net::UdpSocket::bind("127.0.0.1:0").await.unwrap()); - server.entries.insert( - Socks5Entry { - src: local, - dst: remote, - entry_type: UDP_ENTRY, - }, - Socks5EntryData::Udp(( - Arc::new(SocksUdpSocket::UdpSocket(udp_socket)), - UdpClientKey { - client_addr: local, - dst_addr: remote, - }, - )), - ); - assert_eq!(server.entry_count.load(Ordering::Relaxed), 0); - - let mut packet = ZCPacket::new_with_payload(&build_udp_followup_fragment( - match remote.ip() { - IpAddr::V4(ip) => ip, - IpAddr::V6(_) => unreachable!(), - }, - match local.ip() { - IpAddr::V4(ip) => ip, - IpAddr::V6(_) => unreachable!(), - }, - )); - packet.fill_peer_manager_hdr(1, 2, PacketType::Data as u8); - - let result = server.try_process_packet_from_peer(packet).await; - assert!(result.is_some()); - - let mut receiver = server.packet_recv.lock().await; - let received = receiver.try_recv().unwrap(); - assert_eq!( - received.peer_manager_header().unwrap().packet_type, - PacketType::Data as u8 - ); - } - - #[test] - fn decrement_entry_count_does_not_underflow() { - let entry_count = AtomicUsize::new(0); - - let (old_entry_count, new_entry_count) = decrement_entry_count(&entry_count); - - assert_eq!(old_entry_count, 0); - assert_eq!(new_entry_count, 0); - assert_eq!(entry_count.load(Ordering::Relaxed), 0); - } - - #[tokio::test] - async fn removing_missing_entry_does_not_decrement_entry_count() { - let entries = Arc::new(DashMap::new()); - let entry_count = AtomicUsize::new(1); - let entry = Socks5Entry { - src: SocketAddr::new(IpAddr::V4(Ipv4Addr::new(10, 42, 0, 2)), 40000), - dst: SocketAddr::new(IpAddr::V4(Ipv4Addr::new(10, 42, 0, 1)), 22), - entry_type: TCP_ENTRY, - }; - - let (removed, old_entry_count, new_entry_count) = - remove_entry_and_decrement_count(&entries, &entry_count, &entry); - - assert!(!removed); - assert_eq!(old_entry_count, 1); - assert_eq!(new_entry_count, 1); - assert_eq!(entry_count.load(Ordering::Relaxed), 1); - } - - #[tokio::test] - async fn removing_present_entry_decrements_entry_count_once() { - let entries = Arc::new(DashMap::new()); - let entry_count = AtomicUsize::new(0); - let entry = Socks5Entry { - src: SocketAddr::new(IpAddr::V4(Ipv4Addr::new(10, 42, 0, 2)), 40000), - dst: SocketAddr::new(IpAddr::V4(Ipv4Addr::new(10, 42, 0, 1)), 22), - entry_type: TCP_ENTRY, - }; - let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); - insert_entry_and_increment_count( - &entries, - &entry_count, - entry.clone(), - Socks5EntryData::Tcp(listener), - ); - - let (removed, old_entry_count, new_entry_count) = - remove_entry_and_decrement_count(&entries, &entry_count, &entry); - let (removed_again, old_entry_count_again, new_entry_count_again) = - remove_entry_and_decrement_count(&entries, &entry_count, &entry); - - assert!(removed); - assert_eq!(old_entry_count, 1); - assert_eq!(new_entry_count, 0); - assert!(!removed_again); - assert_eq!(old_entry_count_again, 0); - assert_eq!(new_entry_count_again, 0); - assert_eq!(entry_count.load(Ordering::Relaxed), 0); - } - - #[tokio::test] - async fn replacing_present_entry_does_not_increment_entry_count() { - let entries = Arc::new(DashMap::new()); - let entry_count = AtomicUsize::new(0); - let entry = Socks5Entry { - src: SocketAddr::new(IpAddr::V4(Ipv4Addr::new(10, 42, 0, 2)), 40000), - dst: SocketAddr::new(IpAddr::V4(Ipv4Addr::new(10, 42, 0, 1)), 22), - entry_type: TCP_ENTRY, - }; - let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); - let replacement = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); - - let (replaced, old_entry_count, new_entry_count) = insert_entry_and_increment_count( - &entries, - &entry_count, - entry.clone(), - Socks5EntryData::Tcp(listener), - ); - let (replaced_again, old_entry_count_again, new_entry_count_again) = - insert_entry_and_increment_count( - &entries, - &entry_count, - entry, - Socks5EntryData::Tcp(replacement), - ); - - assert!(!replaced); - assert_eq!(old_entry_count, 0); - assert_eq!(new_entry_count, 1); - assert!(replaced_again); - assert_eq!(old_entry_count_again, 1); - assert_eq!(new_entry_count_again, 1); - assert_eq!(entry_count.load(Ordering::Relaxed), 1); - } -} diff --git a/easytier/src/gateway/socks5/dataplane.rs b/easytier/src/gateway/socks5/dataplane.rs deleted file mode 100644 index ae78ef3f..00000000 --- a/easytier/src/gateway/socks5/dataplane.rs +++ /dev/null @@ -1,541 +0,0 @@ -//! Data-plane access built on top of the `Socks5Server` smoltcp stack. -//! -//! This module exposes TCP streams and UDP sockets (mainly for FFI callers that -//! send traffic through EasyTier without creating OS-level proxy listeners). -//! -//! Typical usage: -//! -//! ```ignore -//! let instance = Instance::new(cfg); -//! instance.run().await?; -//! let socks5_server = instance.get_socks5_server(); -//! -//! let socket = socks5_server.data_plane_udp_bind(local_port, timeout).await?; -//! socket.send_to(buf, peer_addr).await?; -//! ``` - -use std::{ - net::{IpAddr, Ipv4Addr, SocketAddr}, - pin::Pin, - sync::{ - Arc, - atomic::{AtomicUsize, Ordering}, - }, - task::{Context, Poll}, - time::Duration, -}; - -use anyhow::Context as _; -use quanta::Instant; -use tokio::io::{AsyncRead, AsyncWrite}; - -use crate::{common::error::Error, gateway::fast_socks5::server::AsyncTcpConnector}; - -use super::{ - Socks5AutoConnector, Socks5Entry, Socks5EntryData, Socks5EntrySet, Socks5Server, - SocksTcpStream, SocksUdpSocket, TCP_ENTRY, TCP_LISTEN_ENTRY, UDP_ENTRY, UdpClientKey, - decrement_entry_count, insert_entry_and_increment_count, try_insert_entry_and_increment_count, -}; -use crate::gateway::tokio_smoltcp::{Net, TcpListener}; - -struct DataPlaneRef { - refs: Arc, - notifier: Arc, -} - -/// A route-table entry whose lifetime is tied to this value: constructing it -/// reserves the route and bumps the active-entry count, dropping it removes the -/// route and drops the count back. -struct OwnedRouteEntry { - entries: Socks5EntrySet, - entry_count: Arc, - entry: Socks5Entry, -} - -impl OwnedRouteEntry { - /// Inserts the route, replacing any existing entry for the same key. - fn register( - entries: Socks5EntrySet, - entry_count: Arc, - entry: Socks5Entry, - ) -> Self { - insert_entry_and_increment_count( - &entries, - &entry_count, - entry.clone(), - Socks5EntryData::DataPlaneRoute, - ); - Self { - entries, - entry_count, - entry, - } - } - - /// Inserts the route only if the key is free, returning `None` on conflict. - fn try_register( - entries: Socks5EntrySet, - entry_count: Arc, - entry: Socks5Entry, - ) -> Option { - if !try_insert_entry_and_increment_count( - &entries, - &entry_count, - entry.clone(), - Socks5EntryData::DataPlaneRoute, - ) { - return None; - } - Some(Self { - entries, - entry_count, - entry, - }) - } -} - -impl Drop for OwnedRouteEntry { - fn drop(&mut self) { - if self.entries.remove(&self.entry).is_some() { - decrement_entry_count(&self.entry_count); - } - } -} - -/// Tracks how an established data-plane TCP stream keeps its inbound route alive. -/// -/// The two variants capture the intrinsic asymmetry between the connect and -/// accept paths. An outbound stream reserved a source port through the -/// [`Socks5AutoConnector`], which owns the matching route entry and clears it on -/// drop. An accepted stream instead inherits its port and peer from the -/// listener, so it carries merely an [`OwnedRouteEntry`]. -enum DataPlaneTcpStreamRoute { - Outbound(Socks5AutoConnector), - Accepted(OwnedRouteEntry), -} - -/// A TCP stream created by the data plane API. -/// Can be either an actively requested outbound connection or an outbound request accepted from a TCP listener. -pub struct DataPlaneTcpStream { - stream: SocksTcpStream, - local_addr: SocketAddr, - _data_plane_ref: DataPlaneRef, - _route: DataPlaneTcpStreamRoute, -} - -/// A TCP listener created by the data plane API. -/// It accepts inbound connections and produces [`DataPlaneTcpStream`]s. -pub struct DataPlaneTcpListener { - listener: TcpListener, - local_addr: SocketAddr, - entries: Socks5EntrySet, - entry_count: Arc, - _listen_route: OwnedRouteEntry, - _data_plane_ref: DataPlaneRef, -} - -pub struct DataPlaneUdpSocket { - socket: Arc, - entries: Socks5EntrySet, - entry_count: Arc, - local_addr: SocketAddr, - _data_plane_ref: DataPlaneRef, -} - -impl Drop for DataPlaneRef { - fn drop(&mut self) { - if self.refs.fetch_sub(1, Ordering::Relaxed) == 1 { - self.notifier.notify_one(); - } - } -} - -impl DataPlaneTcpStream { - pub fn local_addr(&self) -> SocketAddr { - self.local_addr - } -} - -impl DataPlaneTcpListener { - pub fn local_addr(&self) -> SocketAddr { - self.local_addr - } - - pub async fn accept(&mut self) -> Result<(DataPlaneTcpStream, SocketAddr), std::io::Error> { - let (stream, peer_addr) = self.listener.accept().await?; - let local_addr = stream.local_addr()?; - let route = OwnedRouteEntry::register( - self.entries.clone(), - self.entry_count.clone(), - Socks5Entry { - src: local_addr, - dst: peer_addr, - entry_type: TCP_ENTRY, - }, - ); - let accepted = DataPlaneTcpStream { - stream: SocksTcpStream::SmolTcp(stream), - local_addr, - _data_plane_ref: self._data_plane_ref.clone(), - _route: DataPlaneTcpStreamRoute::Accepted(route), - }; - Ok((accepted, peer_addr)) - } -} - -impl AsyncRead for DataPlaneTcpStream { - fn poll_read( - self: Pin<&mut Self>, - cx: &mut Context<'_>, - buf: &mut tokio::io::ReadBuf<'_>, - ) -> Poll> { - Pin::new(&mut self.get_mut().stream).poll_read(cx, buf) - } -} - -impl AsyncWrite for DataPlaneTcpStream { - fn poll_write( - self: Pin<&mut Self>, - cx: &mut Context<'_>, - buf: &[u8], - ) -> Poll> { - Pin::new(&mut self.get_mut().stream).poll_write(cx, buf) - } - - fn poll_flush(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { - Pin::new(&mut self.get_mut().stream).poll_flush(cx) - } - - fn poll_shutdown( - self: Pin<&mut Self>, - cx: &mut Context<'_>, - ) -> Poll> { - Pin::new(&mut self.get_mut().stream).poll_shutdown(cx) - } -} - -impl DataPlaneUdpSocket { - pub fn local_addr(&self) -> SocketAddr { - self.local_addr - } - - pub async fn send_to(&self, buf: &[u8], addr: SocketAddr) -> Result { - let key = Socks5Entry { - src: self.local_addr, - dst: addr, - entry_type: UDP_ENTRY, - }; - try_insert_entry_and_increment_count( - &self.entries, - &self.entry_count, - key, - Socks5EntryData::Udp(( - self.socket.clone(), - UdpClientKey { - client_addr: self.local_addr, - dst_addr: addr, - }, - )), - ); - self.socket.send_to(buf, addr).await - } - - pub async fn recv_from(&self, buf: &mut [u8]) -> Result<(usize, SocketAddr), std::io::Error> { - self.socket.recv_from(buf).await - } -} - -impl Drop for DataPlaneUdpSocket { - fn drop(&mut self) { - let mut removed_entries = 0; - self.entries.retain(|_, data| match data { - Socks5EntryData::Udp((socket, _)) if Arc::ptr_eq(socket, &self.socket) => { - removed_entries += 1; - false - } - _ => true, - }); - super::decrement_entry_count_by(&self.entry_count, removed_entries); - } -} - -impl Clone for DataPlaneRef { - fn clone(&self) -> Self { - self.refs.fetch_add(1, Ordering::Relaxed); - Self { - refs: self.refs.clone(), - notifier: self.notifier.clone(), - } - } -} - -impl Socks5Server { - fn acquire_data_plane_ref(&self) -> DataPlaneRef { - self.data_plane_refs.fetch_add(1, Ordering::Relaxed); - self.port_forward_list_change_notifier.notify_one(); - DataPlaneRef { - refs: self.data_plane_refs.clone(), - notifier: self.port_forward_list_change_notifier.clone(), - } - } - - async fn wait_data_plane_net( - &self, - deadline: Instant, - ) -> Result<(cidr::Ipv4Inet, Arc), Error> { - let mut ready = self.data_plane_net_ready.subscribe(); - loop { - if let Some(net) = self - .net - .lock() - .await - .as_ref() - .map(|net| (net.ipv4_addr, net.smoltcp_net.clone())) - { - return Ok(net); - } - - let now = Instant::now(); - if now >= deadline { - return Err(anyhow::anyhow!("data plane net is not ready").into()); - } - let _ = tokio::time::timeout(deadline - now, ready.wait_for(|ready| *ready)).await; - } - } - - pub async fn data_plane_tcp_connect( - &self, - dst_addr: SocketAddr, - timeout: Duration, - ) -> Result { - let data_plane_ref = self.acquire_data_plane_ref(); - let deadline = Instant::now() + timeout; - let (ipv4_addr, smoltcp_net) = self.wait_data_plane_net(deadline).await?; - // FIXME: This is the data-plane source address reserved for route - // matching. `Socks5AutoConnector` may fall back to direct TCP for - // non-virtual destinations, so this is not always the OS socket's - // local address. - let local_port = smoltcp_net.get_port(); - let local_addr = SocketAddr::new(IpAddr::V4(ipv4_addr.address()), local_port); - let connector = Socks5AutoConnector { - #[cfg(feature = "kcp")] - kcp_endpoint: self.kcp_endpoint.lock().await.clone(), - peer_mgr: self.peer_manager.clone(), - entries: self.entries.clone(), - smoltcp_net: Some(smoltcp_net), - src_addr: local_addr, - entry_count: self.entry_count.clone(), - inner_connector: parking_lot::Mutex::new(None), - }; - - let remaining = deadline.saturating_duration_since(Instant::now()); - let inner_timeout_s = remaining.as_secs().saturating_add(1); - let stream = - tokio::time::timeout(remaining, connector.tcp_connect(dst_addr, inner_timeout_s)) - .await - .with_context(|| "data plane tcp connect timeout")? - .map_err(anyhow::Error::from)?; - Ok(DataPlaneTcpStream { - stream, - local_addr, - _data_plane_ref: data_plane_ref, - _route: DataPlaneTcpStreamRoute::Outbound(connector), - }) - } - - pub async fn data_plane_tcp_bind( - &self, - local_port: u16, - timeout: Duration, - ) -> Result { - let data_plane_ref = self.acquire_data_plane_ref(); - let deadline = Instant::now() + timeout; - let (ipv4_addr, smoltcp_net) = self.wait_data_plane_net(deadline).await?; - let bind_addr = SocketAddr::new(IpAddr::V4(ipv4_addr.address()), local_port); - let listener = smoltcp_net.tcp_bind(bind_addr).await?; - let local_addr = listener.local_addr()?; - let listen_route = OwnedRouteEntry::try_register( - self.entries.clone(), - self.entry_count.clone(), - Socks5Entry { - src: local_addr, - dst: SocketAddr::new(IpAddr::V4(Ipv4Addr::UNSPECIFIED), 0), - entry_type: TCP_LISTEN_ENTRY, - }, - ) - .ok_or_else(|| anyhow::anyhow!("data plane tcp listener already exists"))?; - - Ok(DataPlaneTcpListener { - listener, - local_addr, - entries: self.entries.clone(), - entry_count: self.entry_count.clone(), - _listen_route: listen_route, - _data_plane_ref: data_plane_ref, - }) - } - - pub async fn data_plane_udp_bind( - &self, - local_port: u16, - timeout: Duration, - ) -> Result { - let data_plane_ref = self.acquire_data_plane_ref(); - let deadline = Instant::now() + timeout; - let (ipv4_addr, smoltcp_net) = self.wait_data_plane_net(deadline).await?; - let bind_addr = SocketAddr::new(IpAddr::V4(ipv4_addr.address()), local_port); - let smol = smoltcp_net.udp_bind(bind_addr).await?; - let local_addr = smol.local_addr()?; - let socket = Arc::new(SocksUdpSocket::SmolUdpSocket(smol)); - - Ok(DataPlaneUdpSocket { - socket, - entries: self.entries.clone(), - entry_count: self.entry_count.clone(), - local_addr, - _data_plane_ref: data_plane_ref, - }) - } -} - -#[cfg(test)] -mod tests { - use std::time::Duration; - - use tokio::io::{AsyncReadExt, AsyncWriteExt}; - - use super::Socks5Server; - use crate::peers::peer_manager::PeerManager; - use crate::peers::tests::{connect_peer_manager, create_mock_peer_manager}; - use crate::tunnel::common::tests::wait_for_condition; - - /// A peer and its data-plane server. `Socks5Server` only holds a `Weak` - /// reference to the `PeerManager`, so the manager must be kept alive by the - /// test for the server's smoltcp <-> peer routing to work. - struct Endpoint { - _peer: std::sync::Arc, - server: std::sync::Arc, - ip: cidr::Ipv4Inet, - } - - /// Brings up two peers connected by a ring tunnel, each with a virtual IPv4 - /// and a running `Socks5Server`, and waits until the route to `b`'s IPv4 is - /// visible from `a`. `run(None)` leaves the kcp endpoint unset, so the - /// connect path goes through smoltcp, matching the listener side under test. - async fn setup_pair() -> (Endpoint, Endpoint) { - let a = create_mock_peer_manager().await; - let b = create_mock_peer_manager().await; - connect_peer_manager(a.clone(), b.clone()).await; - - let a_ip: cidr::Ipv4Inet = "10.126.126.1/24".parse().unwrap(); - let b_ip: cidr::Ipv4Inet = "10.126.126.2/24".parse().unwrap(); - a.get_global_ctx().set_ipv4(Some(a_ip)); - b.get_global_ctx().set_ipv4(Some(b_ip)); - - let server_a = Socks5Server::new(a.get_global_ctx(), a.clone(), None); - let server_b = Socks5Server::new(b.get_global_ctx(), b.clone(), None); - server_a.run(None).await.unwrap(); - server_b.run(None).await.unwrap(); - - wait_for_condition( - || async { - a.get_route() - .get_peer_id_by_ipv4(&b_ip.address()) - .await - .is_some() - }, - Duration::from_secs(10), - ) - .await; - - ( - Endpoint { - _peer: a, - server: server_a, - ip: a_ip, - }, - Endpoint { - _peer: b, - server: server_b, - ip: b_ip, - }, - ) - } - - #[tokio::test] - async fn data_plane_tcp_pingpong() { - let (ep_a, ep_b) = setup_pair().await; - let (server_a, server_b, b_ip) = (ep_a.server, ep_b.server, ep_b.ip); - let timeout = Duration::from_secs(10); - - let mut listener = server_b.data_plane_tcp_bind(0, timeout).await.unwrap(); - let listen_addr = - std::net::SocketAddr::new(b_ip.address().into(), listener.local_addr().port()); - - let accept = tokio::spawn(async move { - let (mut stream, _peer) = listener.accept().await.unwrap(); - let mut buf = [0u8; 4]; - stream.read_exact(&mut buf).await.unwrap(); - assert_eq!(&buf, b"ping"); - stream.write_all(b"pong").await.unwrap(); - stream.flush().await.unwrap(); - // Hold the listener and stream until the client has read the reply. - tokio::time::sleep(Duration::from_secs(1)).await; - }); - - let mut client = server_a - .data_plane_tcp_connect(listen_addr, timeout) - .await - .unwrap(); - client.write_all(b"ping").await.unwrap(); - client.flush().await.unwrap(); - let mut buf = [0u8; 4]; - client.read_exact(&mut buf).await.unwrap(); - assert_eq!(&buf, b"pong"); - - accept.await.unwrap(); - } - - #[tokio::test] - async fn data_plane_udp_pingpong() { - let (ep_a, ep_b) = setup_pair().await; - let (server_a, a_ip, server_b, b_ip) = (ep_a.server, ep_a.ip, ep_b.server, ep_b.ip); - let timeout = Duration::from_secs(10); - - let sock_a = server_a.data_plane_udp_bind(0, timeout).await.unwrap(); - let sock_b = server_b.data_plane_udp_bind(0, timeout).await.unwrap(); - let addr_a = std::net::SocketAddr::new(a_ip.address().into(), sock_a.local_addr().port()); - let addr_b = std::net::SocketAddr::new(b_ip.address().into(), sock_b.local_addr().port()); - - // UDP data-plane routes are connected-style: a socket only accepts - // inbound datagrams from a peer it has already sent to, because the - // route entry is registered by `send_to`. Prime b's route toward a so - // the upcoming ping is routed instead of dropped at b's packet filter. - // This datagram is dropped at a (a has no route yet) and is not awaited. - sock_b.send_to(b"warmup", addr_a).await.unwrap(); - - sock_a.send_to(b"ping", addr_b).await.unwrap(); - let mut buf = [0u8; 16]; - let (n, from) = tokio::time::timeout(timeout, sock_b.recv_from(&mut buf)) - .await - .expect("recv ping timed out") - .unwrap(); - assert_eq!(&buf[..n], b"ping"); - assert_eq!(from, addr_a); - - sock_b.send_to(b"pong", addr_a).await.unwrap(); - // a may also receive the stray warmup datagram (it arrives once a has - // registered its route by sending the ping above), so skip anything - // that is not the reply. - loop { - let (n, from) = tokio::time::timeout(timeout, sock_a.recv_from(&mut buf)) - .await - .expect("recv pong timed out") - .unwrap(); - if &buf[..n] == b"pong" { - assert_eq!(from, addr_b); - break; - } - } - } -} diff --git a/easytier/src/gateway/tcp_proxy.rs b/easytier/src/gateway/tcp_proxy.rs deleted file mode 100644 index d5af4340..00000000 --- a/easytier/src/gateway/tcp_proxy.rs +++ /dev/null @@ -1,1036 +0,0 @@ -use anyhow::Context; -use cidr::Ipv4Inet; -use core::panic; -use crossbeam::atomic::AtomicCell; -use dashmap::DashMap; -use pnet::packet::MutablePacket; -use pnet::packet::Packet; -use pnet::packet::ip::IpNextHeaderProtocols; -use pnet::packet::ipv4::{Ipv4Packet, MutableIpv4Packet}; -use pnet::packet::tcp::{MutableTcpPacket, TcpPacket, ipv4_checksum}; -use quanta::Instant; -use socket2::{SockRef, TcpKeepalive}; -use std::net::{IpAddr, Ipv4Addr, SocketAddr, SocketAddrV4}; -use std::sync::atomic::{AtomicBool, AtomicU16}; -use std::sync::{Arc, Weak}; -use std::time::Duration; -use tokio::io::{AsyncRead, AsyncWrite, AsyncWriteExt, copy_bidirectional}; -use tokio::net::{TcpListener, TcpSocket, TcpStream}; -use tokio::sync::{Mutex, mpsc}; -use tokio::task::JoinSet; -use tokio::time::timeout; -use tracing::Instrument; - -use crate::common::error::Result; -use crate::common::global_ctx::{ArcGlobalCtx, GlobalCtx}; -use crate::common::join_joinset_background; -use crate::common::log; -use crate::common::stats_manager::{LabelSet, LabelType, MetricName}; -use crate::peers::peer_manager::PeerManager; -use crate::peers::{NicPacketFilter, PeerPacketFilter}; -use crate::proto::api::instance::{ - ListTcpProxyEntryRequest, ListTcpProxyEntryResponse, TcpProxyEntry, TcpProxyEntryState, - TcpProxyEntryTransportType, TcpProxyRpc, -}; -use crate::proto::rpc_types; -use crate::proto::rpc_types::controller::BaseController; -use crate::tunnel::packet_def::{PacketType, PeerManagerHeader, ZCPacket}; - -use super::CidrSet; - -#[cfg(feature = "smoltcp")] -use super::tokio_smoltcp::{self, Net, NetConfig, channel_device}; - -#[async_trait::async_trait] -pub(crate) trait NatDstConnector: Send + Sync + Clone + 'static { - type DstStream: AsyncRead + AsyncWrite + Unpin + Send; - - async fn connect(&self, src: SocketAddr, dst: SocketAddr) -> anyhow::Result; - fn check_packet_from_peer_fast(&self, cidr_set: &CidrSet, global_ctx: &GlobalCtx) -> bool; - fn check_packet_from_peer( - &self, - cidr_set: &CidrSet, - global_ctx: &GlobalCtx, - hdr: &PeerManagerHeader, - ipv4: &Ipv4Addr, - real_dst_ip: &mut Ipv4Addr, - ) -> bool; - fn transport_type(&self) -> TcpProxyEntryTransportType; -} - -#[derive(Debug, Clone)] -pub struct NatDstTcpConnector; - -#[async_trait::async_trait] -impl NatDstConnector for NatDstTcpConnector { - type DstStream = TcpStream; - async fn connect( - &self, - _src: SocketAddr, - nat_dst: SocketAddr, - ) -> anyhow::Result { - let socket = TcpSocket::new_v4() - .inspect_err(|error| log::error!(?error, "create v4 socket failed"))?; - - let stream = timeout(Duration::from_secs(10), socket.connect(nat_dst)) - .await? - .with_context(|| format!("connect to nat dst failed: {:?}", nat_dst))?; - - prepare_kernel_tcp_socket(&stream)?; - - Ok(stream) - } - - fn check_packet_from_peer_fast(&self, cidr_set: &CidrSet, global_ctx: &GlobalCtx) -> bool { - !cidr_set.is_empty() || global_ctx.enable_exit_node() || global_ctx.no_tun() - } - - fn check_packet_from_peer( - &self, - cidr_set: &CidrSet, - global_ctx: &GlobalCtx, - hdr: &PeerManagerHeader, - ipv4: &Ipv4Addr, - real_dst_ip: &mut Ipv4Addr, - ) -> bool { - let is_exit_node = hdr.is_exit_node(); - - if !(cidr_set.contains_v4(*ipv4, real_dst_ip) - || is_exit_node - || global_ctx.no_tun() - && Some(*ipv4) == global_ctx.get_ipv4().as_ref().map(Ipv4Inet::address)) - { - return false; - } - - true - } - - fn transport_type(&self) -> TcpProxyEntryTransportType { - TcpProxyEntryTransportType::Tcp - } -} - -type NatDstEntryState = TcpProxyEntryState; - -#[derive(Debug)] -pub struct NatDstEntry { - id: uuid::Uuid, - src: SocketAddr, - real_dst: SocketAddr, - mapped_dst: SocketAddr, - start_time: Instant, - start_time_local: chrono::DateTime, - tasks: Mutex>, - state: AtomicCell, -} - -impl NatDstEntry { - pub fn new(src: SocketAddr, real_dst: SocketAddr, mapped_dst: SocketAddr) -> Self { - Self { - id: uuid::Uuid::new_v4(), - src, - real_dst, - mapped_dst, - start_time: Instant::now(), - start_time_local: chrono::Local::now(), - tasks: Mutex::new(JoinSet::new()), - state: AtomicCell::new(NatDstEntryState::SynReceived), - } - } - - fn parse_as_pb(&self, transport_type: TcpProxyEntryTransportType) -> TcpProxyEntry { - TcpProxyEntry { - src: Some(self.src.into()), - dst: Some(self.real_dst.into()), - start_time: self.start_time_local.timestamp() as u64, - state: self.state.load().into(), - transport_type: transport_type.into(), - } - } -} - -enum ProxyTcpStream { - KernelTcpStream(TcpStream), - #[cfg(feature = "smoltcp")] - SmolTcpStream(tokio_smoltcp::TcpStream), -} - -impl ProxyTcpStream { - pub fn set_nodelay(&self, nodelay: bool) -> Result<()> { - match self { - Self::KernelTcpStream(stream) => stream.set_nodelay(nodelay).map_err(Into::into), - #[cfg(feature = "smoltcp")] - Self::SmolTcpStream(_stream) => { - tracing::warn!("smol tcp stream set_nodelay not implemented"); - Ok(()) - } - } - } - - pub async fn shutdown(&mut self) -> Result<()> { - match self { - Self::KernelTcpStream(stream) => { - stream.shutdown().await?; - Ok(()) - } - #[cfg(feature = "smoltcp")] - Self::SmolTcpStream(stream) => { - stream.shutdown().await?; - Ok(()) - } - } - } - - pub async fn copy_bidirectional( - &mut self, - dst: &mut D, - ) -> Result<()> { - match self { - Self::KernelTcpStream(stream) => { - copy_bidirectional(stream, dst).await?; - Ok(()) - } - #[cfg(feature = "smoltcp")] - Self::SmolTcpStream(stream) => { - copy_bidirectional(stream, dst).await?; - Ok(()) - } - } - } -} - -#[cfg(feature = "smoltcp")] -type SmolTcpAcceptResult = Result<(tokio_smoltcp::TcpStream, SocketAddr)>; -#[cfg(feature = "smoltcp")] -struct SmolTcpListener { - stream_tx: mpsc::UnboundedSender, - stream_rx: mpsc::UnboundedReceiver, - - tasks: Arc>>, -} - -#[cfg(feature = "smoltcp")] -impl SmolTcpListener { - pub async fn new() -> Self { - let tasks = Arc::new(std::sync::Mutex::new(JoinSet::new())); - join_joinset_background(tasks.clone(), "smoltcp listener".to_owned()); - - let (tx, rx) = mpsc::unbounded_channel(); - - Self { - stream_tx: tx, - stream_rx: rx, - tasks, - } - } - - pub async fn accept(&mut self) -> SmolTcpAcceptResult { - self.stream_rx.recv().await.unwrap() - } - - pub fn stream_tx(&self) -> mpsc::UnboundedSender { - self.stream_tx.clone() - } - - pub async fn add_listener( - tx: mpsc::UnboundedSender, - net: Arc>>, - tasks: Arc>>, - ) { - let locked_net = net.lock().await; - let mut tcp = locked_net - .as_ref() - .unwrap() - .tcp_bind("0.0.0.0:8899".parse().unwrap()) - .await - .unwrap(); - tasks.lock().unwrap().spawn(async move { - let ret = timeout(Duration::from_secs(10), tcp.accept()).await; - if let Ok(accept_ret) = ret { - tx.send(accept_ret.map_err(|e| { - anyhow::anyhow!("smol tcp listener accept failed: {:?}", e).into() - })) - .unwrap(); - } else { - tracing::error!("smol tcp listener accept timeout"); - } - }); - } -} - -enum ProxyTcpListener { - KernelTcpListener(TcpListener), - #[cfg(feature = "smoltcp")] - SmolTcpListener(SmolTcpListener), -} - -fn prepare_kernel_tcp_socket(stream: &TcpStream) -> Result<()> { - const TCP_KEEPALIVE_TIME: Duration = Duration::from_secs(5); - const TCP_KEEPALIVE_INTERVAL: Duration = Duration::from_secs(2); - const TCP_KEEPALIVE_RETRIES: u32 = 2; - - let ka = TcpKeepalive::new() - .with_time(TCP_KEEPALIVE_TIME) - .with_interval(TCP_KEEPALIVE_INTERVAL); - - #[cfg(not(target_os = "windows"))] - let ka = ka.with_retries(TCP_KEEPALIVE_RETRIES); - - let sf = SockRef::from(&stream); - sf.set_tcp_keepalive(&ka)?; - if let Err(e) = sf.set_nodelay(true) { - tracing::warn!("set_nodelay failed, ignore it: {:?}", e); - } - - Ok(()) -} - -impl ProxyTcpListener { - pub async fn accept(&mut self) -> Result<(ProxyTcpStream, SocketAddr)> { - match self { - Self::KernelTcpListener(listener) => { - let (stream, addr) = listener.accept().await?; - prepare_kernel_tcp_socket(&stream)?; - Ok((ProxyTcpStream::KernelTcpStream(stream), addr)) - } - #[cfg(feature = "smoltcp")] - Self::SmolTcpListener(listener) => { - let Ok((stream, src)) = listener.accept().await else { - return Err(anyhow::anyhow!("smol tcp listener closed").into()); - }; - tracing::info!(?src, "smol tcp listener accepted"); - Ok((ProxyTcpStream::SmolTcpStream(stream), src)) - } - } - } -} - -type ArcNatDstEntry = Arc; - -type SynSockMap = Arc>; -type ConnSockMap = Arc>; -// peer src addr to nat entry, when respond tcp packet, should modify the tcp src addr to the nat entry's dst addr -type AddrConnSockMap = Arc>; - -#[derive(Debug)] -pub struct TcpProxy { - global_ctx: Arc, - peer_manager: Weak, - local_port: AtomicU16, - - tasks: Arc>>, - - syn_map: SynSockMap, - conn_map: ConnSockMap, - addr_conn_map: AddrConnSockMap, - - cidr_set: CidrSet, - - smoltcp_stack_sender: Option>, - smoltcp_stack_receiver: Arc>>>, - #[cfg(feature = "smoltcp")] - smoltcp_net: Arc>>, - #[cfg(feature = "smoltcp")] - smoltcp_listener_tx: std::sync::Mutex>>, - enable_smoltcp: Arc, - - connector: C, -} - -#[async_trait::async_trait] -impl PeerPacketFilter for TcpProxy { - async fn try_process_packet_from_peer(&self, mut packet: ZCPacket) -> Option { - if self.try_handle_peer_packet(&mut packet).await.is_some() { - if self.is_smoltcp_enabled() { - let smoltcp_stack_sender = self.smoltcp_stack_sender.as_ref().unwrap(); - if let Err(e) = smoltcp_stack_sender.try_send(packet) { - tracing::error!("send to smoltcp stack failed: {:?}", e); - } - } else if let Some(peer_manager) = self.get_peer_manager() - && let Err(e) = peer_manager.get_nic_channel().send(packet).await - { - tracing::error!("send to nic failed: {:?}", e); - } - return None; - } else { - Some(packet) - } - } -} - -#[async_trait::async_trait] -impl NicPacketFilter for TcpProxy { - async fn try_process_packet_from_nic(&self, zc_packet: &mut ZCPacket) -> bool { - let Some(my_ipv4_inet) = self.get_local_inet() else { - return false; - }; - let my_ipv4 = my_ipv4_inet.address(); - - let data = zc_packet.payload(); - let ip_packet = Ipv4Packet::new(data).unwrap(); - if ip_packet.get_version() != 4 - || ip_packet.get_source() != my_ipv4 - || ip_packet.get_next_level_protocol() != IpNextHeaderProtocols::Tcp - { - return false; - } - - let tcp_packet = TcpPacket::new(ip_packet.payload()).unwrap(); - if tcp_packet.get_source() != self.get_local_port() { - return false; - } - - let mut dst_addr = SocketAddr::V4(SocketAddrV4::new( - ip_packet.get_destination(), - tcp_packet.get_destination(), - )); - let mut need_transform_dst = false; - - // for kcp proxy, the src ip of nat entry will be converted from my ip to fake ip - // here we need to convert it back - if !self.is_smoltcp_enabled() && dst_addr.ip() == Self::get_fake_local_ipv4(&my_ipv4_inet) { - dst_addr.set_ip(IpAddr::V4(my_ipv4)); - need_transform_dst = true; - } - - tracing::trace!(dst_addr = ?dst_addr, "tcp packet try find entry"); - let entry = if let Some(entry) = self.addr_conn_map.get(&dst_addr) { - entry - } else { - let Some(syn_entry) = self.syn_map.get(&dst_addr) else { - return false; - }; - syn_entry - }; - let nat_entry = entry.clone(); - drop(entry); - assert_eq!(nat_entry.src, dst_addr); - - let IpAddr::V4(ip) = nat_entry.mapped_dst.ip() else { - panic!("v4 nat entry src ip is not v4"); - }; - - zc_packet - .mut_peer_manager_header() - .unwrap() - .set_no_proxy(true); - if need_transform_dst { - zc_packet.mut_peer_manager_header().unwrap().to_peer_id = self.get_my_peer_id().into(); - } - - let mut ip_packet = MutableIpv4Packet::new(zc_packet.mut_payload()).unwrap(); - ip_packet.set_source(ip); - if need_transform_dst { - ip_packet.set_destination(my_ipv4); - } - let dst = ip_packet.get_destination(); - - let mut tcp_packet = MutableTcpPacket::new(ip_packet.payload_mut()).unwrap(); - tcp_packet.set_source(nat_entry.real_dst.port()); - - Self::update_tcp_packet_checksum(&mut tcp_packet, &ip, &dst); - drop(tcp_packet); - Self::update_ip_packet_checksum(&mut ip_packet); - - tracing::trace!(dst_addr = ?dst_addr, nat_entry = ?nat_entry, packet = ?ip_packet, "tcp packet after modified"); - - true - } -} - -impl TcpProxy { - pub fn new(peer_manager: Arc, connector: C) -> Arc { - let (smoltcp_stack_sender, smoltcp_stack_receiver) = mpsc::channel::(1000); - let global_ctx = peer_manager.get_global_ctx(); - - Arc::new(Self { - global_ctx: global_ctx.clone(), - peer_manager: Arc::downgrade(&peer_manager), - - local_port: AtomicU16::new(0), - tasks: Arc::new(std::sync::Mutex::new(JoinSet::new())), - - syn_map: Arc::new(DashMap::new()), - conn_map: Arc::new(DashMap::new()), - addr_conn_map: Arc::new(DashMap::new()), - - cidr_set: CidrSet::new(global_ctx), - - smoltcp_stack_sender: Some(smoltcp_stack_sender), - smoltcp_stack_receiver: Arc::new(Mutex::new(Some(smoltcp_stack_receiver))), - - #[cfg(feature = "smoltcp")] - smoltcp_net: Arc::new(Mutex::new(None)), - #[cfg(feature = "smoltcp")] - smoltcp_listener_tx: std::sync::Mutex::new(None), - - enable_smoltcp: Arc::new(AtomicBool::new(true)), - - connector, - }) - } - - pub fn get_peer_manager(&self) -> Option> { - self.peer_manager.upgrade() - } - - fn update_tcp_packet_checksum( - tcp_packet: &mut MutableTcpPacket, - ipv4_src: &Ipv4Addr, - ipv4_dst: &Ipv4Addr, - ) { - tcp_packet.set_checksum(ipv4_checksum( - &tcp_packet.to_immutable(), - ipv4_src, - ipv4_dst, - )); - } - - fn update_ip_packet_checksum(ip_packet: &mut MutableIpv4Packet) { - ip_packet.set_checksum(pnet::packet::ipv4::checksum(&ip_packet.to_immutable())); - } - - pub async fn start(self: &Arc, add_pipeline: bool) -> Result<()> { - self.run_syn_map_cleaner().await?; - self.run_listener().await?; - if add_pipeline { - let peer_manager = self - .get_peer_manager() - .ok_or_else(|| anyhow::anyhow!("peer manager is gone"))?; - peer_manager - .add_packet_process_pipeline(Box::new(self.clone())) - .await; - peer_manager - .add_nic_packet_process_pipeline(Box::new(self.clone())) - .await; - } - join_joinset_background(self.tasks.clone(), "TcpProxy".to_owned()); - - Ok(()) - } - - async fn run_syn_map_cleaner(&self) -> Result<()> { - let syn_map = self.syn_map.clone(); - let tasks = self.tasks.clone(); - let syn_map_cleaner_task = async move { - loop { - syn_map.retain(|_, entry| { - if entry.start_time.elapsed() > Duration::from_secs(30) { - tracing::warn!(entry = ?entry, "syn nat entry expired"); - entry.state.store(NatDstEntryState::Closed); - false - } else { - true - } - }); - syn_map.shrink_to_fit(); - tokio::time::sleep(Duration::from_secs(10)).await; - } - }; - tasks.lock().unwrap().spawn(syn_map_cleaner_task); - - Ok(()) - } - - async fn get_proxy_listener(&self) -> Result { - #[cfg(feature = "smoltcp")] - if self.global_ctx.get_flags().use_smoltcp - || self.global_ctx.no_tun() - || cfg!(any( - target_os = "android", - any( - target_os = "ios", - all(target_os = "macos", feature = "macos-ne") - ), - target_env = "ohos" - )) - { - // use smoltcp network stack - - use crate::gateway::tokio_smoltcp::BufferSize; - self.local_port - .store(8899, std::sync::atomic::Ordering::Relaxed); - - let mut cap = smoltcp::phy::DeviceCapabilities::default(); - cap.max_transmission_unit = 1280; - cap.medium = smoltcp::phy::Medium::Ip; - let (dev, stack_sink, mut stack_stream) = channel_device::ChannelDevice::new(cap); - - let mut smoltcp_stack_receiver = - self.smoltcp_stack_receiver.lock().await.take().unwrap(); - self.tasks.lock().unwrap().spawn(async move { - while let Some(packet) = smoltcp_stack_receiver.recv().await { - tracing::trace!(?packet, "receive from peer send to smoltcp packet"); - if let Err(e) = stack_sink.send(Ok(packet.payload().to_vec())).await { - tracing::error!("send to smoltcp stack failed: {:?}", e); - } - } - tracing::error!("smoltcp stack sink exited"); - }); - - let peer_mgr = self.peer_manager.clone(); - self.tasks.lock().unwrap().spawn(async move { - while let Some(data) = stack_stream.recv().await { - tracing::trace!( - ?data, - "receive from smoltcp stack and send to peer mgr packet" - ); - let Some(ipv4) = Ipv4Packet::new(&data) else { - tracing::error!(?data, "smoltcp stack stream get non ipv4 packet"); - continue; - }; - - let dst = ipv4.get_destination(); - let packet = ZCPacket::new_with_payload(&data); - let Some(peer_mgr) = peer_mgr.upgrade() else { - tracing::warn!("peer manager is gone, smoltcp sender exited"); - return; - }; - if let Err(e) = peer_mgr - .send_msg_by_ip(packet, IpAddr::V4(dst), false) - .await - { - tracing::error!("send to peer failed in smoltcp sender: {:?}", e); - } - } - tracing::error!("smoltcp stack stream exited"); - }); - - let interface_config = smoltcp::iface::Config::new(smoltcp::wire::HardwareAddress::Ip); - let net = Net::new( - dev, - NetConfig::new( - interface_config, - format!("{}/24", self.get_local_ip().unwrap()) - .parse() - .unwrap(), - vec![format!("{}", self.get_local_ip().unwrap()).parse().unwrap()], - Some(BufferSize { - tcp_rx_size: 1024 * 16, - tcp_tx_size: 1024 * 16, - ..Default::default() - }), - ), - ); - net.set_any_ip(true); - self.smoltcp_net.lock().await.replace(net); - let tcp = SmolTcpListener::new().await; - self.smoltcp_listener_tx - .lock() - .unwrap() - .replace(tcp.stream_tx()); - - self.enable_smoltcp - .store(true, std::sync::atomic::Ordering::Relaxed); - - return Ok(ProxyTcpListener::SmolTcpListener(tcp)); - } - - { - // use kernel network stack - let listen_addr = SocketAddr::new(Ipv4Addr::UNSPECIFIED.into(), 0); - let net_ns = self.global_ctx.net_ns.clone(); - let tcp_listener = net_ns - .run_async(|| async { TcpListener::bind(&listen_addr).await }) - .await?; - self.local_port.store( - tcp_listener.local_addr()?.port(), - std::sync::atomic::Ordering::Relaxed, - ); - - self.enable_smoltcp - .store(false, std::sync::atomic::Ordering::Relaxed); - - Ok(ProxyTcpListener::KernelTcpListener(tcp_listener)) - } - } - - async fn run_listener(&self) -> Result<()> { - // bind on both v4 & v6 - let mut tcp_listener = self.get_proxy_listener().await?; - - let global_ctx = self.global_ctx.clone(); - let tasks = Arc::downgrade(&self.tasks); - let syn_map = self.syn_map.clone(); - let conn_map = self.conn_map.clone(); - let addr_conn_map = self.addr_conn_map.clone(); - let connector = self.connector.clone(); - let accept_task = async move { - let conn_map = conn_map.clone(); - loop { - let accept_ret = tcp_listener.accept().await; - let Ok((tcp_stream, mut socket_addr)) = accept_ret else { - tracing::error!("nat tcp listener accept failed: {:?}", accept_ret.err()); - continue; - }; - - let my_ip_inet = global_ctx.get_ipv4(); - let my_ip = my_ip_inet - .as_ref() - .map(Ipv4Inet::address) - .unwrap_or(Ipv4Addr::UNSPECIFIED); - - if my_ip_inet.is_some() - && socket_addr.ip() == Self::get_fake_local_ipv4(&my_ip_inet.unwrap()) - { - socket_addr.set_ip(IpAddr::V4(my_ip)); - } - - let Some(entry) = syn_map.get(&socket_addr) else { - tracing::error!( - ?my_ip, - ?socket_addr, - "tcp connection from unknown source, ignore it" - ); - continue; - }; - tracing::info!( - ?socket_addr, - "tcp connection accepted for proxy, nat dst: {:?}", - entry.real_dst - ); - assert_eq!(entry.state.load(), NatDstEntryState::SynReceived); - - let entry_clone = entry.clone(); - drop(entry); - syn_map.remove_if(&socket_addr, |_, entry| entry.id == entry_clone.id); - - entry_clone.state.store(NatDstEntryState::ConnectingDst); - - let _ = addr_conn_map.insert(entry_clone.src, entry_clone.clone()); - let old_nat_val = conn_map.insert(entry_clone.id, entry_clone.clone()); - assert!(old_nat_val.is_none()); - - let Some(tasks) = tasks.upgrade() else { - tracing::error!("tcp proxy tasks is dropped, exit accept loop"); - break; - }; - - tasks.lock().unwrap().spawn(Self::connect_to_nat_dst( - connector.clone(), - global_ctx.clone(), - tcp_stream, - conn_map.clone(), - addr_conn_map.clone(), - entry_clone, - )); - } - }; - self.tasks - .lock() - .unwrap() - .spawn(accept_task.instrument(tracing::info_span!("tcp_proxy_listener"))); - - Ok(()) - } - - fn remove_entry_from_all_conn_map( - conn_map: ConnSockMap, - addr_conn_map: AddrConnSockMap, - nat_entry: ArcNatDstEntry, - ) { - conn_map.remove(&nat_entry.id); - addr_conn_map.remove_if(&nat_entry.src, |_, entry| entry.id == nat_entry.id); - if conn_map.capacity() - conn_map.len() > 16 { - conn_map.shrink_to_fit(); - } - if addr_conn_map.capacity() - addr_conn_map.len() > 16 { - addr_conn_map.shrink_to_fit(); - } - } - - async fn connect_to_nat_dst( - connector: C, - global_ctx: ArcGlobalCtx, - src_tcp_stream: ProxyTcpStream, - conn_map: ConnSockMap, - addr_conn_map: AddrConnSockMap, - nat_entry: ArcNatDstEntry, - ) { - if let Err(e) = src_tcp_stream.set_nodelay(true) { - tracing::warn!("set_nodelay failed, ignore it: {:?}", e); - } - - if global_ctx.should_deny_proxy(&nat_entry.real_dst, false) { - tracing::error!( - ?nat_entry, - "nat dst port {} is in running listeners, ignore it", - nat_entry.real_dst.port() - ); - nat_entry.state.store(NatDstEntryState::Closed); - Self::remove_entry_from_all_conn_map(conn_map, addr_conn_map, nat_entry); - return; - } - - let nat_dst = if global_ctx.is_ip_local_virtual_ip(&nat_entry.real_dst.ip()) { - format!("127.0.0.1:{}", nat_entry.real_dst.port()) - .parse() - .unwrap() - } else { - nat_entry.real_dst - }; - - global_ctx - .stats_manager() - .get_counter( - MetricName::TcpProxyConnect, - LabelSet::new() - .with_label_type(LabelType::Protocol( - connector.transport_type().as_str_name().to_string(), - )) - .with_label_type(LabelType::DstIp(nat_dst.ip().to_string())) - .with_label_type(LabelType::MappedDstIp( - nat_entry.mapped_dst.ip().to_string(), - )), - ) - .inc(); - - let _guard = global_ctx.net_ns.guard(); - let Ok(dst_tcp_stream) = connector.connect(nat_entry.src, nat_dst).await else { - tracing::error!("connect to dst failed: {:?}", nat_entry); - nat_entry.state.store(NatDstEntryState::Closed); - Self::remove_entry_from_all_conn_map(conn_map, addr_conn_map, nat_entry); - return; - }; - drop(_guard); - - tracing::info!(?nat_entry, ?nat_dst, "tcp connection to dst established"); - - assert_eq!(nat_entry.state.load(), NatDstEntryState::ConnectingDst); - nat_entry.state.store(NatDstEntryState::Connected); - - Self::handle_nat_connection( - src_tcp_stream, - dst_tcp_stream, - conn_map, - addr_conn_map, - nat_entry, - ) - .await; - } - - async fn handle_nat_connection( - mut src_tcp_stream: ProxyTcpStream, - mut dst_tcp_stream: C::DstStream, - conn_map: ConnSockMap, - addr_conn_map: AddrConnSockMap, - nat_entry: ArcNatDstEntry, - ) { - let nat_entry_clone = nat_entry.clone(); - nat_entry.tasks.lock().await.spawn(async move { - let ret = src_tcp_stream.copy_bidirectional(&mut dst_tcp_stream).await; - tracing::info!(nat_entry = ?nat_entry_clone, ret = ?ret, "nat tcp connection closed"); - - nat_entry_clone.state.store(NatDstEntryState::ClosingSrc); - let ret = timeout(Duration::from_secs(10), src_tcp_stream.shutdown()).await; - tracing::info!(nat_entry = ?nat_entry_clone, ret = ?ret, "src tcp stream shutdown"); - - nat_entry_clone.state.store(NatDstEntryState::ClosingDst); - let ret = timeout(Duration::from_secs(10), dst_tcp_stream.shutdown()).await; - tracing::info!(nat_entry = ?nat_entry_clone, ret = ?ret, "dst tcp stream shutdown"); - - drop(src_tcp_stream); - drop(dst_tcp_stream); - - nat_entry_clone.state.store(NatDstEntryState::Closed); - // sleep later so the fin packet can be processed - tokio::time::sleep(Duration::from_secs(10)).await; - - Self::remove_entry_from_all_conn_map(conn_map, addr_conn_map, nat_entry_clone); - }); - } - - pub fn get_local_port(&self) -> u16 { - self.local_port.load(std::sync::atomic::Ordering::Relaxed) - } - - pub fn get_my_peer_id(&self) -> u32 { - self.peer_manager - .upgrade() - .map(|pm| pm.my_peer_id()) - .unwrap_or_default() - } - - pub fn get_local_ip(&self) -> Option { - self.get_local_inet().map(|inet| inet.address()) - } - - pub fn get_local_inet(&self) -> Option { - if self.is_smoltcp_enabled() { - Some(Ipv4Inet::new(Ipv4Addr::new(192, 88, 99, 254), 24).unwrap()) - } else { - self.global_ctx.get_ipv4().as_ref().cloned() - } - } - - pub fn get_global_ctx(&self) -> &ArcGlobalCtx { - &self.global_ctx - } - - pub fn is_smoltcp_enabled(&self) -> bool { - self.enable_smoltcp - .load(std::sync::atomic::Ordering::Relaxed) - } - - pub fn get_fake_local_ipv4(local_ip: &Ipv4Inet) -> Ipv4Addr { - local_ip.first_address() - } - - async fn try_handle_peer_packet(&self, packet: &mut ZCPacket) -> Option<()> { - if !self - .connector - .check_packet_from_peer_fast(&self.cidr_set, &self.global_ctx) - { - return None; - } - - let ipv4_inet = self.get_local_inet()?; - let ipv4_addr = ipv4_inet.address(); - { - let hdr = packet.peer_manager_header().unwrap(); - if (hdr.packet_type != PacketType::Data as u8 - && hdr.packet_type != PacketType::DataWithKcpSrcModified as u8 - && hdr.packet_type != PacketType::DataWithQuicSrcModified as u8) - || hdr.is_no_proxy() - { - return None; - }; - } - - let origin_ip = { - let payload_bytes = packet.mut_payload(); - let ipv4 = Ipv4Packet::new(payload_bytes)?; - if ipv4.get_version() != 4 - || ipv4.get_next_level_protocol() != IpNextHeaderProtocols::Tcp - { - return None; - } - - ipv4.get_destination() - }; - let mut real_dst_ip = origin_ip; - let hdr = packet.mut_peer_manager_header().unwrap(); - - if !self.connector.check_packet_from_peer( - &self.cidr_set, - &self.global_ctx, - hdr, - &origin_ip, - &mut real_dst_ip, - ) { - return None; - } - - // restore to data packet - hdr.packet_type = PacketType::Data as u8; - - tracing::trace!(ipv4 = ?origin_ip, cidr_set = ?self.cidr_set, "proxy tcp packet received"); - - let payload_bytes = packet.mut_payload(); - let ip_packet = Ipv4Packet::new(payload_bytes).unwrap(); - let tcp_packet = TcpPacket::new(ip_packet.payload()).unwrap(); - - let source_ip = ip_packet.get_source(); - let source_port = tcp_packet.get_source(); - let src = SocketAddr::V4(SocketAddrV4::new(source_ip, source_port)); - - let is_tcp_syn = tcp_packet.get_flags() & pnet::packet::tcp::TcpFlags::SYN != 0; - let is_tcp_ack = tcp_packet.get_flags() & pnet::packet::tcp::TcpFlags::ACK != 0; - if is_tcp_syn && !is_tcp_ack { - let dest_ip = ip_packet.get_destination(); - let dest_port = tcp_packet.get_destination(); - let mapped_dst = SocketAddr::V4(SocketAddrV4::new(dest_ip, dest_port)); - let real_dst = SocketAddr::V4(SocketAddrV4::new(real_dst_ip, dest_port)); - - let old_val = self - .syn_map - .insert(src, Arc::new(NatDstEntry::new(src, real_dst, mapped_dst))); - tracing::info!(src = ?src, ?real_dst, ?mapped_dst, old_entry = ?old_val, "tcp syn received"); - - // if smoltcp is enabled, add the listener to the net - #[cfg(feature = "smoltcp")] - if self.is_smoltcp_enabled() { - let smoltcp_listener_tx = self.smoltcp_listener_tx.lock().unwrap().clone().unwrap(); - SmolTcpListener::add_listener( - smoltcp_listener_tx, - self.smoltcp_net.clone(), - self.tasks.clone(), - ) - .await; - tracing::info!("smol tcp listener added for src: {:?}", src); - } - } else if !self.addr_conn_map.contains_key(&src) && !self.syn_map.contains_key(&src) { - // if not in syn map and addr conn map, may forwarding n2n packet - return None; - } - - let mut ip_packet = MutableIpv4Packet::new(payload_bytes).unwrap(); - if !self.is_smoltcp_enabled() && source_ip == ipv4_addr { - // modify the source so the response packet can be handled by tun device - ip_packet.set_source(Self::get_fake_local_ipv4(&ipv4_inet)); - } - ip_packet.set_destination(ipv4_addr); - let source = ip_packet.get_source(); - - let mut tcp_packet = MutableTcpPacket::new(ip_packet.payload_mut()).unwrap(); - tcp_packet.set_destination(self.get_local_port()); - - Self::update_tcp_packet_checksum(&mut tcp_packet, &source, &ipv4_addr); - drop(tcp_packet); - Self::update_ip_packet_checksum(&mut ip_packet); - - tracing::trace!(?source, ?ipv4_addr, ?packet, "tcp packet after modified"); - - Some(()) - } - - pub fn is_tcp_proxy_connection(&self, src: SocketAddr) -> bool { - self.syn_map.contains_key(&src) || self.addr_conn_map.contains_key(&src) - } - - pub fn list_proxy_entries(&self) -> Vec { - let mut entries: Vec = Vec::new(); - let transport_type = self.connector.transport_type(); - for entry in self.syn_map.iter() { - entries.push(entry.value().as_ref().parse_as_pb(transport_type)); - } - for entry in self.conn_map.iter() { - entries.push(entry.value().as_ref().parse_as_pb(transport_type)); - } - entries - } - - pub fn get_transport_type(&self) -> TcpProxyEntryTransportType { - self.connector.transport_type() - } -} - -#[derive(Clone)] -pub struct TcpProxyRpcService { - tcp_proxy: Weak>, -} - -#[async_trait::async_trait] -impl TcpProxyRpc for TcpProxyRpcService { - type Controller = BaseController; - async fn list_tcp_proxy_entry( - &self, - _: BaseController, - _request: ListTcpProxyEntryRequest, // Accept request of type HelloRequest - ) -> std::result::Result { - let mut reply = ListTcpProxyEntryResponse::default(); - if let Some(tcp_proxy) = self.tcp_proxy.upgrade() { - reply.entries = tcp_proxy.list_proxy_entries(); - } - Ok(reply) - } -} - -impl TcpProxyRpcService { - pub fn new(tcp_proxy: Arc>) -> Self { - Self { - tcp_proxy: Arc::downgrade(&tcp_proxy), - } - } -} diff --git a/easytier/src/gateway/udp_proxy.rs b/easytier/src/gateway/udp_proxy.rs deleted file mode 100644 index f284fed0..00000000 --- a/easytier/src/gateway/udp_proxy.rs +++ /dev/null @@ -1,817 +0,0 @@ -use std::{ - net::{Ipv4Addr, SocketAddr, SocketAddrV4}, - sync::{Arc, Weak, atomic::AtomicBool}, - time::Duration, -}; - -use bytes::{BufMut, BytesMut}; -use cidr::Ipv4Inet; -use crossbeam::atomic::AtomicCell; -use dashmap::DashMap; -use pnet::packet::{ - Packet, - ip::IpNextHeaderProtocols, - ipv4::Ipv4Packet, - udp::{self, MutableUdpPacket}, -}; -use quanta::Instant; -use tokio::sync::mpsc::{Receiver, Sender, channel, error::TrySendError}; -use tokio::{ - net::UdpSocket, - sync::Mutex, - task::{JoinHandle, JoinSet}, - time::timeout, -}; -use tokio_util::task::AbortOnDropHandle; - -use tracing::Level; - -use super::{CidrSet, ip_reassembler::IpReassembler}; -use crate::tunnel::common::bind; -use crate::{ - common::{PeerId, error::Error, global_ctx::ArcGlobalCtx}, - gateway::ip_reassembler::{ComposeIpv4PacketArgs, compose_ipv4_packet}, - peers::{PeerPacketFilter, peer_manager::PeerManager}, - tunnel::{ - common::reserve_buf, - packet_def::{PacketType, ZCPacket}, - }, -}; - -#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] -struct UdpNatKey { - src_socket: SocketAddr, - dst_socket: SocketAddr, -} - -impl UdpNatKey { - fn new(src_socket: SocketAddr, dst_socket: SocketAddr) -> Self { - Self { - src_socket, - dst_socket, - } - } -} - -#[derive(Debug)] -struct UdpNatEntry { - src_peer_id: PeerId, - my_peer_id: PeerId, - src_socket: SocketAddr, - socket: Option, - forward_task: Mutex>>, - stopped: AtomicBool, - start_time: Instant, - last_active_time: AtomicCell, - denied: bool, -} - -impl UdpNatEntry { - #[tracing::instrument(err(level = Level::WARN))] - fn new( - src_peer_id: PeerId, - my_peer_id: PeerId, - src_socket: SocketAddr, - denied: bool, - ) -> Result { - // TODO: try use src port, so we will be ip restricted nat type - let socket = (!denied) - .then(|| bind().addr("0.0.0.0:0".parse().unwrap()).call()) - .transpose()?; - - Ok(Self { - src_peer_id, - my_peer_id, - src_socket, - socket, - forward_task: Mutex::new(None), - stopped: AtomicBool::new(false), - start_time: Instant::now(), - last_active_time: AtomicCell::new(Instant::now()), - denied, - }) - } - - pub fn stop(&self) { - self.stopped - .store(true, std::sync::atomic::Ordering::Relaxed); - } - - async fn compose_ipv4_packet( - self: &Arc, - packet_sender: &Sender, - buf: &mut [u8], - src_v4: &SocketAddrV4, - payload_len: usize, - payload_mtu: usize, - ip_id: u16, - ) -> Result<(), Error> { - let SocketAddr::V4(nat_src_v4) = self.src_socket else { - return Err(Error::Unknown); - }; - - assert_eq!(0, payload_mtu % 8); - - // udp payload is in buf[20 + 8..] - let mut udp_packet = MutableUdpPacket::new(&mut buf[20..28 + payload_len]).unwrap(); - udp_packet.set_source(src_v4.port()); - udp_packet.set_destination(self.src_socket.port()); - udp_packet.set_length(payload_len as u16 + 8); - udp_packet.set_checksum(udp::ipv4_checksum( - &udp_packet.to_immutable(), - src_v4.ip(), - nat_src_v4.ip(), - )); - - compose_ipv4_packet( - ComposeIpv4PacketArgs { - buf: &mut buf[..], - src_v4: src_v4.ip(), - dst_v4: nat_src_v4.ip(), - next_protocol: IpNextHeaderProtocols::Udp, - payload_len: payload_len + 8, // include udp header - payload_mtu, - ip_id, - }, - |buf| { - let mut p = ZCPacket::new_with_payload(buf); - p.fill_peer_manager_hdr(self.my_peer_id, self.src_peer_id, PacketType::Data as u8); - p.mut_peer_manager_header().unwrap().set_no_proxy(true); - - match packet_sender.try_send(p) { - Err(TrySendError::Closed(e)) => { - tracing::error!("send icmp packet to peer failed: {:?}, may exiting..", e); - Err(Error::Unknown) - } - _ => Ok(()), - } - }, - )?; - - Ok(()) - } - - async fn forward_task( - self: Arc, - packet_sender: Sender, - virtual_ipv4: Ipv4Addr, - real_ipv4: Ipv4Addr, - mapped_ipv4: Ipv4Addr, - ) { - let (s, mut r) = channel(128); - - let self_clone = self.clone(); - let recv_task = AbortOnDropHandle::new(tokio::spawn(async move { - let mut cur_buf = BytesMut::new(); - loop { - if self_clone - .stopped - .load(std::sync::atomic::Ordering::Relaxed) - { - break; - } - - reserve_buf(&mut cur_buf, 64 * 1024 + 28, 128 * 1024 + 28); - assert_eq!(cur_buf.len(), 0); - unsafe { - cur_buf.advance_mut(28); - } - - let (len, src_socket) = match timeout( - Duration::from_secs(120), - self_clone - .socket - .as_ref() - .unwrap() - .recv_buf_from(&mut cur_buf), - ) - .await - { - Ok(Ok(x)) => x, - Ok(Err(err)) => { - tracing::error!(?err, "udp nat recv failed"); - break; - } - Err(err) => { - tracing::error!(?err, "udp nat recv timeout"); - break; - } - }; - - tracing::trace!(?len, ?src_socket, "udp nat packet response received"); - - let ret_buf = cur_buf.split(); - s.send((ret_buf, len, src_socket)).await.unwrap(); - } - })); - - let self_clone = self.clone(); - let send_task = AbortOnDropHandle::new(tokio::spawn(async move { - let mut ip_id = 1; - while let Some((mut packet, len, src_socket)) = r.recv().await { - let SocketAddr::V4(mut src_v4) = src_socket else { - continue; - }; - - self_clone.mark_active(); - - let has_mapped_dst = real_ipv4 != mapped_ipv4; - let mut reply_src_ip = *src_v4.ip(); - - // Preserve the existing priority for proxy rules that expose a - // real loopback address as a mapped address. Other loopback - // replies come from local delivery to 127.0.0.1 for the local - // virtual IP and may need the mapped rewrite below. - if has_mapped_dst && reply_src_ip == real_ipv4 { - reply_src_ip = mapped_ipv4; - } else if reply_src_ip.is_loopback() { - reply_src_ip = virtual_ipv4; - } - - if has_mapped_dst && reply_src_ip == real_ipv4 { - reply_src_ip = mapped_ipv4; - } - src_v4.set_ip(reply_src_ip); - - let Ok(_) = Self::compose_ipv4_packet( - &self_clone, - &packet_sender, - &mut packet, - &src_v4, - len, - 1280, - ip_id, - ) - .await - else { - break; - }; - ip_id = ip_id.wrapping_add(1); - } - })); - - let _ = tokio::join!(recv_task, send_task); - - self.stop(); - } - - fn mark_active(&self) { - self.last_active_time.store(Instant::now()); - } - - fn is_active(&self) -> bool { - self.last_active_time.load().elapsed().as_secs() < 180 - } -} - -#[derive(Debug)] -pub struct UdpProxy { - global_ctx: ArcGlobalCtx, - peer_manager: Weak, - - cidr_set: CidrSet, - - nat_table: Arc>>, - - sender: Sender, - receiver: Mutex>>, - - tasks: Mutex>, - - ip_resemmbler: Arc, -} - -impl UdpProxy { - async fn try_handle_packet(&self, packet: &ZCPacket) -> Option<()> { - if self.cidr_set.is_empty() - && !self.global_ctx.enable_exit_node() - && !self.global_ctx.no_tun() - { - return None; - } - - let _ = self.global_ctx.get_ipv4()?; - let hdr = packet.peer_manager_header().unwrap(); - let is_exit_node = hdr.is_exit_node(); - if hdr.packet_type != PacketType::Data as u8 || hdr.is_no_proxy() { - return None; - }; - - let ipv4 = Ipv4Packet::new(packet.payload())?; - if ipv4.get_version() != 4 || ipv4.get_next_level_protocol() != IpNextHeaderProtocols::Udp { - return None; - } - - let mut real_dst_ip = ipv4.get_destination(); - - if !(self - .cidr_set - .contains_v4(ipv4.get_destination(), &mut real_dst_ip) - || is_exit_node - || self.global_ctx.no_tun() - && Some(ipv4.get_destination()) - == self.global_ctx.get_ipv4().as_ref().map(Ipv4Inet::address)) - { - return None; - } - - let resembled_buf: Option>; - let udp_packet = if IpReassembler::is_packet_fragmented(&ipv4) { - resembled_buf = - self.ip_resemmbler - .add_fragment(ipv4.get_source(), ipv4.get_destination(), &ipv4); - resembled_buf.as_ref()?; - udp::UdpPacket::new(resembled_buf.as_ref().unwrap())? - } else { - udp::UdpPacket::new(ipv4.payload())? - }; - - // TODO: should it be async. - let dst_socket = if self.global_ctx.is_ip_local_virtual_ip(&real_dst_ip.into()) { - format!("127.0.0.1:{}", udp_packet.get_destination()) - .parse() - .unwrap() - } else { - SocketAddr::new(real_dst_ip.into(), udp_packet.get_destination()) - }; - - tracing::trace!( - ?packet, - ?ipv4, - ?udp_packet, - "udp nat packet request received" - ); - - let nat_key = UdpNatKey::new( - SocketAddr::new(ipv4.get_source().into(), udp_packet.get_source()), - SocketAddr::new(ipv4.get_destination().into(), udp_packet.get_destination()), - ); - let nat_entry = self - .nat_table - .entry(nat_key) - .or_try_insert_with::(|| { - tracing::info!(?packet, ?ipv4, ?udp_packet, "udp nat table entry created"); - let denied = self.global_ctx.should_deny_proxy( - &SocketAddr::new(real_dst_ip.into(), udp_packet.get_destination()), - true, - ); - let _g = self.global_ctx.net_ns.guard(); - Ok(Arc::new(UdpNatEntry::new( - hdr.from_peer_id.get(), - hdr.to_peer_id.get(), - nat_key.src_socket, - denied, - )?)) - }) - .ok()? - .clone(); - - if nat_entry.denied { - tracing::debug!( - dst_port = udp_packet.get_destination(), - "dst socket is in running listeners, ignore it" - ); - return Some(()); - } - - if nat_entry.forward_task.lock().await.is_none() { - nat_entry - .forward_task - .lock() - .await - .replace(tokio::spawn(UdpNatEntry::forward_task( - nat_entry.clone(), - self.sender.clone(), - self.global_ctx.get_ipv4().map(|x| x.address())?, - real_dst_ip, - ipv4.get_destination(), - ))); - } - - nat_entry.mark_active(); - - let send_ret = { - let _g = self.global_ctx.net_ns.guard(); - nat_entry - .socket - .as_ref() - .unwrap() - .send_to(udp_packet.payload(), dst_socket) - .await - }; - - if let Err(send_err) = send_ret { - tracing::error!( - ?send_err, - ?nat_key, - ?nat_entry, - ?send_err, - "udp nat send failed" - ); - } - - Some(()) - } -} - -#[async_trait::async_trait] -impl PeerPacketFilter for UdpProxy { - async fn try_process_packet_from_peer(&self, packet: ZCPacket) -> Option { - self.try_handle_packet(&packet) - .await - .is_none() - .then_some(packet) - } -} - -impl UdpProxy { - pub fn new( - global_ctx: ArcGlobalCtx, - peer_manager: Arc, - ) -> Result, Error> { - let cidr_set = CidrSet::new(global_ctx.clone()); - let (sender, receiver) = channel(1024); - let ret = Self { - global_ctx, - peer_manager: Arc::downgrade(&peer_manager), - cidr_set, - nat_table: Arc::new(DashMap::new()), - sender, - receiver: Mutex::new(Some(receiver)), - tasks: Mutex::new(JoinSet::new()), - ip_resemmbler: Arc::new(IpReassembler::new(Duration::from_secs(10))), - }; - Ok(Arc::new(ret)) - } - - pub async fn start(self: &Arc) -> Result<(), Error> { - let Some(peer_manager) = self.peer_manager.upgrade() else { - return Err(anyhow::anyhow!("peer manager is gone").into()); - }; - peer_manager - .add_packet_process_pipeline(Box::new(self.clone())) - .await; - - // clean up nat table - let nat_table = self.nat_table.clone(); - self.tasks.lock().await.spawn(async move { - loop { - tokio::time::sleep(Duration::from_secs(15)).await; - nat_table.retain(|_, v| { - if !v.is_active() { - tracing::info!(?v, "udp nat table entry removed"); - v.stop(); - false - } else { - true - } - }); - nat_table.shrink_to_fit(); - } - }); - - let ip_resembler = self.ip_resemmbler.clone(); - self.tasks.lock().await.spawn(async move { - loop { - tokio::time::sleep(Duration::from_secs(1)).await; - ip_resembler.remove_expired_packets(); - } - }); - - // forward packets to peer manager - let mut receiver = self.receiver.lock().await.take().unwrap(); - let peer_manager = self.peer_manager.clone(); - let is_latency_first = self.global_ctx.latency_first(); - self.tasks.lock().await.spawn(async move { - while let Some(mut msg) = receiver.recv().await { - let hdr = msg.mut_peer_manager_header().unwrap(); - hdr.set_latency_first(is_latency_first); - let to_peer_id = hdr.to_peer_id.into(); - tracing::trace!(?msg, ?to_peer_id, "udp nat packet response send"); - let Some(pm) = peer_manager.upgrade() else { - tracing::warn!("peer manager is gone, udp proxy send loop exit"); - return; - }; - let ret = pm.send_msg_for_proxy(msg, to_peer_id).await; - if ret.is_err() { - tracing::error!("send icmp packet to peer failed: {:?}", ret); - } - } - }); - Ok(()) - } -} - -impl Drop for UdpProxy { - fn drop(&mut self) { - for v in self.nat_table.iter() { - v.stop(); - } - } -} - -#[cfg(test)] -mod tests { - use std::{ - net::{Ipv4Addr, SocketAddr}, - sync::Arc, - time::Duration, - }; - - use pnet::packet::{ - MutablePacket, Packet, - ip::IpNextHeaderProtocols, - ipv4::{self, Ipv4Packet, MutableIpv4Packet}, - udp::{self, MutableUdpPacket, UdpPacket}, - }; - use tokio::{net::UdpSocket, sync::mpsc::Receiver, time::timeout}; - - use crate::{ - common::{config::ConfigLoader, global_ctx::tests::get_mock_global_ctx}, - peers::{ - create_packet_recv_chan, - peer_manager::{PeerManager, RouteAlgoType}, - }, - tunnel::packet_def::{PacketType, ZCPacket}, - }; - - use super::UdpProxy; - - fn build_udp_proxy_packet( - src_ip: Ipv4Addr, - src_port: u16, - dst_socket: SocketAddr, - payload: &[u8], - ) -> ZCPacket { - let SocketAddr::V4(dst_socket) = dst_socket else { - panic!("test only builds IPv4 UDP packets"); - }; - let dst_ip = *dst_socket.ip(); - let mut packet = vec![0; 20 + 8 + payload.len()]; - let packet_len = packet.len() as u16; - - { - let mut ipv4_packet = MutableIpv4Packet::new(&mut packet).unwrap(); - ipv4_packet.set_version(4); - ipv4_packet.set_header_length(5); - ipv4_packet.set_total_length(packet_len); - ipv4_packet.set_ttl(64); - ipv4_packet.set_next_level_protocol(IpNextHeaderProtocols::Udp); - ipv4_packet.set_source(src_ip); - ipv4_packet.set_destination(dst_ip); - } - - { - let mut udp_packet = MutableUdpPacket::new(&mut packet[20..]).unwrap(); - udp_packet.set_source(src_port); - udp_packet.set_destination(dst_socket.port()); - udp_packet.set_length((8 + payload.len()) as u16); - udp_packet.payload_mut().copy_from_slice(payload); - udp_packet.set_checksum(udp::ipv4_checksum( - &udp_packet.to_immutable(), - &src_ip, - &dst_ip, - )); - } - - { - let mut ipv4_packet = MutableIpv4Packet::new(&mut packet).unwrap(); - ipv4_packet.set_checksum(ipv4::checksum(&ipv4_packet.to_immutable())); - } - - let mut packet = ZCPacket::new_with_payload(&packet); - packet.fill_peer_manager_hdr(1009867077, 3831440917, PacketType::Data as u8); - packet - } - - async fn wait_proxy_cidr_loaded(proxy: &UdpProxy) { - timeout(Duration::from_secs(1), async { - while proxy.cidr_set.is_empty() { - tokio::time::sleep(Duration::from_millis(10)).await; - } - }) - .await - .unwrap(); - } - - async fn recv_payload(socket: &UdpSocket) -> (Vec, SocketAddr) { - let mut buf = [0; 64]; - let (len, addr) = timeout(Duration::from_secs(1), socket.recv_from(&mut buf)) - .await - .unwrap() - .unwrap(); - (buf[..len].to_vec(), addr) - } - - async fn recv_response_packet(receiver: &mut Receiver) -> ZCPacket { - timeout(Duration::from_secs(1), receiver.recv()) - .await - .unwrap() - .unwrap() - } - - fn assert_udp_response( - packet: ZCPacket, - src_socket: SocketAddr, - dst_ip: Ipv4Addr, - dst_port: u16, - payload: &[u8], - ) { - let SocketAddr::V4(src_socket) = src_socket else { - panic!("test only checks IPv4 UDP packets"); - }; - let ipv4_packet = Ipv4Packet::new(packet.payload()).unwrap(); - assert_eq!(ipv4_packet.get_source(), *src_socket.ip()); - assert_eq!(ipv4_packet.get_destination(), dst_ip); - - let udp_packet = UdpPacket::new(ipv4_packet.payload()).unwrap(); - assert_eq!(udp_packet.get_source(), src_socket.port()); - assert_eq!(udp_packet.get_destination(), dst_port); - assert_eq!(udp_packet.payload(), payload); - } - - async fn stop_nat_entries(proxy: &UdpProxy) { - let nat_socket_addrs = proxy - .nat_table - .iter() - .filter_map(|entry| { - entry - .socket - .as_ref() - .and_then(|socket| socket.local_addr().ok()) - .map(|addr| SocketAddr::from((Ipv4Addr::LOCALHOST, addr.port()))) - }) - .collect::>(); - - for entry in proxy.nat_table.iter() { - entry.stop(); - } - - let wake_socket = UdpSocket::bind((Ipv4Addr::LOCALHOST, 0)).await.unwrap(); - for addr in nat_socket_addrs { - let _ = wake_socket.send_to(b"wake", addr).await; - } - } - - #[tokio::test] - async fn udp_proxy_rewrites_unmapped_loopback_reply_to_virtual_ip() { - let global_ctx = get_mock_global_ctx(); - global_ctx.set_ipv4(Some("10.144.144.204/24".parse().unwrap())); - global_ctx - .config - .add_proxy_cidr("127.0.0.1/32".parse().unwrap(), None) - .unwrap(); - - let (packet_sender, _packet_receiver) = create_packet_recv_chan(); - let peer_manager = Arc::new(PeerManager::new( - RouteAlgoType::Ospf, - global_ctx.clone(), - packet_sender, - )); - let proxy = UdpProxy::new(global_ctx, peer_manager).unwrap(); - wait_proxy_cidr_loaded(&proxy).await; - let mut response_receiver = proxy.receiver.lock().await.take().unwrap(); - - let real_dst = UdpSocket::bind((Ipv4Addr::LOCALHOST, 0)).await.unwrap(); - let real_dst_port = real_dst.local_addr().unwrap().port(); - let dst_socket = SocketAddr::from((Ipv4Addr::LOCALHOST, real_dst_port)); - let src_ip = Ipv4Addr::new(10, 144, 144, 206); - let src_port = 53864; - - let packet = build_udp_proxy_packet(src_ip, src_port, dst_socket, b"request"); - assert!(proxy.try_handle_packet(&packet).await.is_some()); - let (payload, nat_socket) = recv_payload(&real_dst).await; - assert_eq!(payload, b"request"); - - real_dst.send_to(b"reply", nat_socket).await.unwrap(); - assert_udp_response( - recv_response_packet(&mut response_receiver).await, - SocketAddr::from((Ipv4Addr::new(10, 144, 144, 204), real_dst_port)), - src_ip, - src_port, - b"reply", - ); - - stop_nat_entries(&proxy).await; - } - - #[tokio::test] - async fn udp_proxy_maps_local_virtual_destination_reply_to_mapped_source() { - let global_ctx = get_mock_global_ctx(); - global_ctx.set_ipv4(Some("10.144.144.204/24".parse().unwrap())); - global_ctx - .config - .add_proxy_cidr( - "10.144.144.204/32".parse().unwrap(), - Some("10.10.10.3/32".parse().unwrap()), - ) - .unwrap(); - - let (packet_sender, _packet_receiver) = create_packet_recv_chan(); - let peer_manager = Arc::new(PeerManager::new( - RouteAlgoType::Ospf, - global_ctx.clone(), - packet_sender, - )); - let proxy = UdpProxy::new(global_ctx, peer_manager).unwrap(); - wait_proxy_cidr_loaded(&proxy).await; - let mut response_receiver = proxy.receiver.lock().await.take().unwrap(); - - let real_dst = UdpSocket::bind((Ipv4Addr::LOCALHOST, 0)).await.unwrap(); - let real_dst_port = real_dst.local_addr().unwrap().port(); - let mapped_dst = SocketAddr::from((Ipv4Addr::new(10, 10, 10, 3), real_dst_port)); - let src_ip = Ipv4Addr::new(10, 144, 144, 206); - let src_port = 53864; - - let packet = build_udp_proxy_packet(src_ip, src_port, mapped_dst, b"request"); - assert!(proxy.try_handle_packet(&packet).await.is_some()); - let (payload, nat_socket) = recv_payload(&real_dst).await; - assert_eq!(payload, b"request"); - - real_dst.send_to(b"reply", nat_socket).await.unwrap(); - assert_udp_response( - recv_response_packet(&mut response_receiver).await, - mapped_dst, - src_ip, - src_port, - b"reply", - ); - - stop_nat_entries(&proxy).await; - } - - #[tokio::test] - async fn udp_proxy_separates_same_source_port_to_multiple_mapped_destinations() { - let global_ctx = get_mock_global_ctx(); - global_ctx.set_ipv4(Some("10.144.144.204/24".parse().unwrap())); - global_ctx - .config - .add_proxy_cidr( - "127.0.0.1/32".parse().unwrap(), - Some("10.10.10.1/32".parse().unwrap()), - ) - .unwrap(); - global_ctx - .config - .add_proxy_cidr( - "127.0.0.1/32".parse().unwrap(), - Some("10.10.10.2/32".parse().unwrap()), - ) - .unwrap(); - - let (packet_sender, _packet_receiver) = create_packet_recv_chan(); - let peer_manager = Arc::new(PeerManager::new( - RouteAlgoType::Ospf, - global_ctx.clone(), - packet_sender, - )); - let proxy = UdpProxy::new(global_ctx, peer_manager).unwrap(); - wait_proxy_cidr_loaded(&proxy).await; - let mut response_receiver = proxy.receiver.lock().await.take().unwrap(); - - let real_dst = UdpSocket::bind((Ipv4Addr::LOCALHOST, 0)).await.unwrap(); - let real_dst_port = real_dst.local_addr().unwrap().port(); - let first_mapped_dst = SocketAddr::from((Ipv4Addr::new(10, 10, 10, 1), real_dst_port)); - let second_mapped_dst = SocketAddr::from((Ipv4Addr::new(10, 10, 10, 2), real_dst_port)); - let src_ip = Ipv4Addr::new(10, 144, 144, 206); - let src_port = 53864; - - let first_packet = build_udp_proxy_packet(src_ip, src_port, first_mapped_dst, b"first"); - assert!(proxy.try_handle_packet(&first_packet).await.is_some()); - let (payload, first_nat_socket) = recv_payload(&real_dst).await; - assert_eq!(payload, b"first"); - - let second_packet = build_udp_proxy_packet(src_ip, src_port, second_mapped_dst, b"second"); - assert!(proxy.try_handle_packet(&second_packet).await.is_some()); - let (payload, second_nat_socket) = recv_payload(&real_dst).await; - assert_eq!(payload, b"second"); - - assert_eq!(proxy.nat_table.len(), 2); - - real_dst - .send_to(b"first-reply", first_nat_socket) - .await - .unwrap(); - assert_udp_response( - recv_response_packet(&mut response_receiver).await, - first_mapped_dst, - src_ip, - src_port, - b"first-reply", - ); - - real_dst - .send_to(b"second-reply", second_nat_socket) - .await - .unwrap(); - assert_udp_response( - recv_response_packet(&mut response_receiver).await, - second_mapped_dst, - src_ip, - src_port, - b"second-reply", - ); - - stop_nat_entries(&proxy).await; - } -} diff --git a/easytier/src/gateway/wrapped_proxy.rs b/easytier/src/gateway/wrapped_proxy.rs deleted file mode 100644 index 1e5aefd8..00000000 --- a/easytier/src/gateway/wrapped_proxy.rs +++ /dev/null @@ -1,153 +0,0 @@ -use std::{ - net::{IpAddr, Ipv4Addr, SocketAddr}, - sync::Arc, -}; - -use pnet::packet::{ - Packet as _, - ip::IpNextHeaderProtocols, - ipv4::Ipv4Packet, - tcp::{TcpFlags, TcpPacket}, -}; -use tokio::io::{AsyncRead, AsyncWrite, copy_bidirectional}; -use tokio_util::io::InspectReader; - -use crate::tunnel::packet_def::{PacketType, PeerManagerHeader}; -use crate::{ - common::{acl_processor::PacketInfo, error::Result}, - gateway::tcp_proxy::{NatDstConnector, TcpProxy}, - peers::{NicPacketFilter, acl_filter::AclFilter}, - proto::acl::{Action, ChainType}, - tunnel::packet_def::ZCPacket, -}; - -#[derive(Clone)] -pub struct ProxyAclHandler { - pub acl_filter: Arc, - pub packet_info: PacketInfo, - pub chain_type: ChainType, -} - -impl ProxyAclHandler { - pub fn handle_packet(&self, buf: &[u8]) -> Result<()> { - let mut packet_info = self.packet_info.clone(); - packet_info.packet_size = buf.len(); - let ret = self - .acl_filter - .get_processor() - .process_packet(&packet_info, self.chain_type); - self.acl_filter.handle_acl_result( - &ret, - &packet_info, - self.chain_type, - &self.acl_filter.get_processor(), - ); - if !matches!(ret.action, Action::Allow) { - return Err(anyhow::anyhow!("acl denied").into()); - } - - Ok(()) - } - - pub async fn copy_bidirection_with_acl( - &self, - src: impl AsyncRead + AsyncWrite + Unpin, - mut dst: impl AsyncRead + AsyncWrite + Unpin, - ) -> Result<()> { - let (src_reader, src_writer) = tokio::io::split(src); - let src_reader = InspectReader::new(src_reader, |buf| { - let _ = self.handle_packet(buf); - }); - let mut src = tokio::io::join(src_reader, src_writer); - - copy_bidirectional(&mut src, &mut dst).await?; - Ok(()) - } -} - -#[async_trait::async_trait] -pub(crate) trait TcpProxyForWrappedSrcTrait: Send + Sync + 'static { - type Connector: NatDstConnector; - fn get_tcp_proxy(&self) -> &Arc>; - fn mark_src_modified(hdr: &mut PeerManagerHeader) -> &mut PeerManagerHeader; - async fn check_dst_allow_wrapped_input(&self, dst_ip: &Ipv4Addr) -> bool; -} - -#[async_trait::async_trait] -impl> NicPacketFilter for T { - async fn try_process_packet_from_nic(&self, zc_packet: &mut ZCPacket) -> bool { - let ret = self - .get_tcp_proxy() - .try_process_packet_from_nic(zc_packet) - .await; - if ret { - return true; - } - - let hdr = zc_packet.mut_peer_manager_header().unwrap(); - if hdr.packet_type != PacketType::Data as u8 { - // already handled by other proxy - return false; - } - - let data = zc_packet.payload(); - let ip_packet = Ipv4Packet::new(data).unwrap(); - if ip_packet.get_version() != 4 - || ip_packet.get_next_level_protocol() != IpNextHeaderProtocols::Tcp - { - return false; - } - - // if no connection is established, only allow SYN packet - let tcp_packet = TcpPacket::new(ip_packet.payload()).unwrap(); - let is_syn = tcp_packet.get_flags() & TcpFlags::SYN != 0 - && tcp_packet.get_flags() & TcpFlags::ACK == 0; - if is_syn { - // only check dst feature flag when SYN packet - if !self - .check_dst_allow_wrapped_input(&ip_packet.get_destination()) - .await - { - tracing::warn!( - "{:?} proxy src: dst {} not allow wrapped input", - self.get_tcp_proxy().get_transport_type(), - ip_packet.get_destination() - ); - return false; - } - } else { - // if not syn packet, only allow established connection - if !self - .get_tcp_proxy() - .is_tcp_proxy_connection(SocketAddr::new( - IpAddr::V4(ip_packet.get_source()), - tcp_packet.get_source(), - )) - { - return false; - } - } - - if let Some(my_ipv4) = self.get_tcp_proxy().get_global_ctx().get_ipv4() { - // this is a net-to-net packet, only allow it when smoltcp is enabled - // because the syn-ack packet will not be through and handled by the tun device when - // the source ip is in the local network - if ip_packet.get_source() != my_ipv4.address() - && !self.get_tcp_proxy().is_smoltcp_enabled() - { - tracing::warn!( - "{:?} nat 2 nat packet, src: {} dst: {} not allow wrapped input", - self.get_tcp_proxy().get_transport_type(), - ip_packet.get_source(), - ip_packet.get_destination() - ); - return false; - } - }; - - let hdr = zc_packet.mut_peer_manager_header().unwrap(); - hdr.to_peer_id = self.get_tcp_proxy().get_my_peer_id().into(); - Self::mark_src_modified(hdr); - true - } -} diff --git a/easytier/src/host_runtime.rs b/easytier/src/host_runtime.rs new file mode 100644 index 00000000..f92ba2f8 --- /dev/null +++ b/easytier/src/host_runtime.rs @@ -0,0 +1,235 @@ +use std::{ + net::{IpAddr, SocketAddr, SocketAddrV4, SocketAddrV6}, + sync::{Arc, OnceLock}, +}; + +use async_trait::async_trait; +use easytier_core::{ + connectivity::{composite::ConnectorRuntime, transport::ConnectedByteStream}, + host::dns::{DnsQuery, DnsRecordResolver, DnsResolver, DnsSrvRecord}, + socket::{ + SocketContext, + tcp::{ + TcpConnectOptions, TcpListenOptions, TcpSocketPurpose, VirtualTcpListenerFactory, + VirtualTcpSocketFactory, + }, + udp::{PreferredIpv6Source, UdpBindOptions, VirtualUdpSocketFactory}, + }, +}; + +use crate::{ + common::{ + dns::RuntimeDnsResolver, + netns::NetNS, + network::{collect_interfaces, collect_local_ip_addrs}, + }, + proto::peer_rpc::GetIpListResponse, + socket::{ + tcp::{RuntimeTcpListener, RuntimeTcpSocket}, + udp::{RuntimeUdpSocket, RuntimeUdpSocketFactory}, + }, +}; + +/// Process-wide native implementation of the host capabilities consumed by core. +/// +/// Instance-specific policy is carried by each request's socket context. Keeping +/// this object stateless prevents a socket operation from capturing one +/// instance's namespace or mark. +#[derive(Debug)] +pub struct NativeHostRuntime { + udp_sockets: RuntimeUdpSocketFactory, + dns: RuntimeDnsResolver, +} + +static NATIVE_HOST_RUNTIME: OnceLock> = OnceLock::new(); + +pub(crate) fn native_host_runtime() -> Arc { + NATIVE_HOST_RUNTIME + .get_or_init(|| { + Arc::new(NativeHostRuntime { + udp_sockets: RuntimeUdpSocketFactory::new(), + dns: RuntimeDnsResolver::new(), + }) + }) + .clone() +} + +#[async_trait] +impl VirtualTcpSocketFactory for NativeHostRuntime { + type Socket = RuntimeTcpSocket; + + async fn connect_tcp(&self, options: TcpConnectOptions) -> anyhow::Result { + #[cfg(feature = "faketcp")] + if options.purpose == TcpSocketPurpose::FakeTcp { + let remote_addr = options.remote_addr; + let socket_mark = options.bind.context.socket_mark; + let net_ns = NetNS::from_socket_context(&options.bind.context); + let socket = + crate::socket::fake_tcp::connect_socket(remote_addr, socket_mark, net_ns).await?; + return Ok(RuntimeTcpSocket::from_fake_tcp(socket)); + } + + #[cfg(not(feature = "faketcp"))] + if options.purpose == TcpSocketPurpose::FakeTcp { + anyhow::bail!("FakeTCP socket support is disabled") + } + + crate::socket::tcp::connect_tcp(options) + .await + .map_err(anyhow::Error::from) + } +} + +#[async_trait] +impl ConnectorRuntime for NativeHostRuntime { + async fn connect_byte_stream( + &self, + url: &url::Url, + ) -> anyhow::Result> { + #[cfg(unix)] + if url.scheme() == "unix" { + let stream = tokio::net::UnixStream::connect(url.path()).await?; + let local_url = stream + .local_addr() + .ok() + .and_then(crate::socket::tcp::url_from_unix_socket_addr); + return Ok(ConnectedByteStream::new( + RuntimeTcpSocket::from_unix(stream), + local_url, + url.clone(), + Some(url.clone()), + )); + } + + anyhow::bail!("unsupported runtime byte stream: {url}") + } + + async fn local_addr_for_remote( + &self, + remote_addr: SocketAddr, + context: SocketContext, + ) -> anyhow::Result { + let socket = NetNS::from_socket_context(&context).run(|| -> anyhow::Result<_> { + let (domain, bind_addr) = match remote_addr { + SocketAddr::V4(_) => ( + socket2::Domain::IPV4, + SocketAddr::V4(SocketAddrV4::new(std::net::Ipv4Addr::UNSPECIFIED, 0)), + ), + SocketAddr::V6(_) => ( + socket2::Domain::IPV6, + SocketAddr::V6(SocketAddrV6::new(std::net::Ipv6Addr::UNSPECIFIED, 0, 0, 0)), + ), + }; + let socket = + socket2::Socket::new(domain, socket2::Type::DGRAM, Some(socket2::Protocol::UDP))?; + crate::tunnel::common::apply_socket_mark(&socket, context.socket_mark)?; + socket.set_nonblocking(true)?; + socket.bind(&socket2::SockAddr::from(bind_addr))?; + Ok(std::net::UdpSocket::from(socket)) + })?; + let socket = tokio::net::UdpSocket::from_std(socket)?; + socket.connect(remote_addr).await?; + Ok(socket.local_addr()?) + } + + async fn preferred_ipv6_source( + &self, + ip: std::net::Ipv6Addr, + context: SocketContext, + ) -> Option { + collect_interfaces(NetNS::from_socket_context(&context), false) + .await + .into_iter() + .find(|interface| { + interface + .ips + .iter() + .any(|local| matches!(local.ip(), IpAddr::V6(local_ip) if local_ip == ip)) + }) + .map(|interface| PreferredIpv6Source { + ip, + ifindex: interface.index, + }) + } + + async fn collect_ip_addrs(&self, context: &SocketContext) -> GetIpListResponse { + collect_local_ip_addrs(NetNS::from_socket_context(context)).await + } +} + +impl NativeHostRuntime { + pub(crate) fn is_local_ip(&self, ip: &IpAddr, context: &SocketContext) -> bool { + NetNS::from_socket_context(context) + .run(|| std::net::UdpSocket::bind(format!("{ip}:0")).is_ok()) + } +} + +#[async_trait] +impl VirtualTcpListenerFactory for NativeHostRuntime { + type Listener = RuntimeTcpListener; + + async fn bind_tcp(&self, options: TcpListenOptions) -> anyhow::Result> { + Ok(Arc::new(crate::socket::tcp::bind_tcp_listener(options)?)) + } +} + +#[async_trait] +impl VirtualUdpSocketFactory for NativeHostRuntime { + type Socket = RuntimeUdpSocket; + + async fn bind_udp(&self, options: UdpBindOptions) -> anyhow::Result> { + let socket = self.udp_sockets.bind_udp(options).await?; + #[cfg(target_os = "windows")] + crate::arch::windows::disable_connection_reset(socket.socket().as_ref())?; + Ok(socket) + } +} + +#[async_trait] +impl DnsResolver for NativeHostRuntime { + async fn resolve(&self, query: DnsQuery) -> anyhow::Result> { + self.dns.resolve(query).await + } +} + +#[async_trait] +impl DnsRecordResolver for NativeHostRuntime { + async fn resolve_txt(&self, query: DnsQuery) -> anyhow::Result { + self.dns.resolve_txt(query).await + } + + async fn resolve_srv(&self, query: DnsQuery) -> anyhow::Result> { + self.dns.resolve_srv(query).await + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn native_host_runtime_is_process_wide() { + assert!(Arc::ptr_eq(&native_host_runtime(), &native_host_runtime())); + } + + #[test] + fn native_local_ip_probe_uses_process_runtime() { + assert!(native_host_runtime().is_local_ip( + &IpAddr::V4(std::net::Ipv4Addr::LOCALHOST), + &SocketContext::default(), + )); + } + + #[tokio::test] + async fn native_route_probe_uses_remote_address_family() { + let local_addr = native_host_runtime() + .local_addr_for_remote( + SocketAddr::from(([127, 0, 0, 1], 9)), + SocketContext::default(), + ) + .await + .unwrap(); + + assert!(local_addr.is_ipv4()); + } +} diff --git a/easytier/src/instance/cli_event_logger.rs b/easytier/src/instance/cli_event_logger.rs new file mode 100644 index 00000000..f88c4df7 --- /dev/null +++ b/easytier/src/instance/cli_event_logger.rs @@ -0,0 +1,272 @@ +use std::fmt::{Display, Formatter}; + +use uuid::Uuid; + +use crate::{ + common::{ + global_ctx::{EventBusSubscriber, GlobalCtxEvent}, + log, + }, + proto, +}; + +struct DisplayPeerConnInfo<'a>(&'a proto::api::instance::PeerConnInfo); + +impl Display for DisplayPeerConnInfo<'_> { + fn fmt(&self, formatter: &mut Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("PeerConnInfo") + .field("my_peer_id", &self.0.my_peer_id) + .field("dst_peer_id", &self.0.peer_id) + .field("tunnel_info", &self.0.tunnel) + .finish() + } +} + +macro_rules! event { + ($lvl:ident, category: $cat:expr, $($args:tt)+) => { + event!(@impl $lvl, concat!("INSTANCE::", $cat), $($args)+) + }; + + ($lvl:ident, $($args:tt)+) => { + event!(@impl $lvl, "INSTANCE", $($args)+) + }; + + (@impl $lvl:ident, $cat:expr, $($args:tt)+) => { + log::$lvl!( + category: $cat, + $($args)+ + ); + }; +} + +pub(super) fn spawn(instance_id: Uuid, events: EventBusSubscriber) { + drop(tokio::spawn(log_events(instance_id, events))); +} + +async fn log_events(instance_id: Uuid, mut events: EventBusSubscriber) { + loop { + match events.recv().await { + Ok(event) => log_event(instance_id, event), + Err(tokio::sync::broadcast::error::RecvError::Closed) => return, + Err(tokio::sync::broadcast::error::RecvError::Lagged(_)) => continue, + } + } +} + +fn log_event(instance_id: Uuid, event: GlobalCtxEvent) { + match event { + GlobalCtxEvent::PeerAdded(peer_id) => { + event!(info, peer_id, "[{}] new peer added", instance_id); + } + GlobalCtxEvent::PeerRemoved(peer_id) => { + event!(info, peer_id, "[{}] peer removed", instance_id); + } + GlobalCtxEvent::PeerConnAdded(conn_info) => { + let conn_info = DisplayPeerConnInfo(&conn_info); + event!( + info, + category: "CONNECTION", + %conn_info, + "[{}] new peer connection added", + instance_id, + ); + } + GlobalCtxEvent::PeerConnRemoved(conn_info) => { + let conn_info = DisplayPeerConnInfo(&conn_info); + event!( + info, + category: "CONNECTION", + %conn_info, + "[{}] peer connection removed", + instance_id, + ); + } + GlobalCtxEvent::ListenerAddFailed(listener, msg) => { + event!(warn, %listener, msg, "[{}] listener add failed", instance_id); + } + GlobalCtxEvent::ListenerAcceptFailed(listener, msg) => { + event!(warn, %listener, msg, "[{}] listener accept failed", instance_id); + } + GlobalCtxEvent::ListenerAdded(listener) => { + if listener.scheme() == "ring" { + return; + } + event!( + info, + %listener, + "[{}] new listener added", + instance_id + ); + } + GlobalCtxEvent::ConnectionAccepted(local, remote) => { + event!( + info, + category: "CONNECTION", + local, + remote, + "[{}] new connection accepted", + instance_id + ); + } + GlobalCtxEvent::ConnectionError(local, remote, err) => { + event!( + info, + category: "CONNECTION", + local, + remote, + err, + "[{}] connection error", + instance_id + ); + } + GlobalCtxEvent::ListenerPortMappingEstablished { + local_listener, + mapped_listener, + backend, + } => { + event!( + info, + %local_listener, + %mapped_listener, + backend, + "[{}] listener port mapping established", + instance_id + ); + } + GlobalCtxEvent::TunDeviceReady(dev) => { + event!(info, dev, "[{}] tun device ready", instance_id); + } + GlobalCtxEvent::TunDeviceError(err) => { + event!(error, %err, "[{}] tun device error", instance_id); + } + GlobalCtxEvent::Connecting(dst) => { + event!( + info, + category: "CONNECTION", + %dst, + "[{}] connecting to peer", + instance_id + ); + } + GlobalCtxEvent::ConnectError(dst, ip_version, error) => { + event!( + info, + category: "CONNECTION", + dst, + ip_version, + %error, + "[{}] connect to peer error", + instance_id + ); + } + GlobalCtxEvent::VpnPortalStarted(portal) => { + event!(info, portal, "[{}] vpn portal started", instance_id); + } + GlobalCtxEvent::VpnPortalClientConnected(portal, client_addr) => { + event!( + info, + portal, + client_addr, + "[{}] vpn portal client connected", + instance_id + ); + } + GlobalCtxEvent::VpnPortalClientDisconnected(portal, client_addr) => { + event!( + info, + portal, + client_addr, + "[{}] vpn portal client disconnected", + instance_id + ); + } + GlobalCtxEvent::DhcpIpv4Changed(old, new) => { + event!(info, ?old, ?new, "[{}] dhcp ip changed", instance_id); + } + GlobalCtxEvent::DhcpIpv4Conflicted(ip) => { + event!(info, ?ip, "[{}] dhcp ip conflict", instance_id); + } + GlobalCtxEvent::PublicIpv6Changed(old, new) => { + event!(info, ?old, ?new, "[{}] public ipv6 changed", instance_id); + } + GlobalCtxEvent::PublicIpv6RoutesUpdated(added, removed) => { + event!( + info, + ?added, + ?removed, + "[{}] public ipv6 routes updated", + instance_id + ); + } + GlobalCtxEvent::PortForwardAdded(cfg) => { + event!( + info, + local = %cfg.bind_addr.unwrap(), + remote = %cfg.dst_addr.unwrap(), + proto = %cfg.socket_type().as_str_name(), + "[{}] port forward added", + instance_id, + ); + } + #[cfg(feature = "management")] + GlobalCtxEvent::ConfigPatched(patch) => { + event!(info, ?patch, "[{}] config patched", instance_id); + } + GlobalCtxEvent::ProxyCidrsUpdated(added, removed) => { + event!( + info, + ?added, + ?removed, + "[{}] proxy CIDRs updated", + instance_id + ); + } + GlobalCtxEvent::UdpBroadcastRelayStartResult { + capture_backend, + error, + } => { + if let Some(error) = error { + event!( + warn, + ?capture_backend, + %error, + "[{}] UDP broadcast relay start failed", + instance_id + ); + } else { + event!( + info, + ?capture_backend, + "[{}] UDP broadcast relay started", + instance_id + ); + } + } + GlobalCtxEvent::CredentialChanged => { + event!(info, "[{}] credential changed", instance_id); + } + } +} + +#[cfg(test)] +mod tests { + use std::time::Duration; + + use tokio::sync::broadcast; + + use super::*; + + #[tokio::test] + async fn event_loop_stops_when_the_instance_closes() { + let (sender, events) = broadcast::channel(1); + let task = tokio::spawn(log_events(Uuid::new_v4(), events)); + + sender.send(GlobalCtxEvent::CredentialChanged).unwrap(); + drop(sender); + tokio::time::timeout(Duration::from_secs(1), task) + .await + .unwrap() + .unwrap(); + } +} diff --git a/easytier/src/instance/composition.rs b/easytier/src/instance/composition.rs new file mode 100644 index 00000000..f30cc1c8 --- /dev/null +++ b/easytier/src/instance/composition.rs @@ -0,0 +1,529 @@ +use std::sync::Arc; + +#[cfg(feature = "wrapped-transport")] +use easytier_core::gateway::proxy::wrapped_transport::WrappedTransportEngines; +#[cfg(feature = "wireguard")] +use easytier_core::gateway::vpn_portal::VpnPortalHost; +#[cfg(feature = "management")] +use easytier_core::{ + connectivity::manual::ManualTunnelConnector, + host::dns::{DnsRecordResolver, DnsResolver}, + instance::CoreInstanceConfig, +}; +use easytier_core::{ + events::{CoreEvent, CoreEventSink}, + host::packet::PacketSink, + instance::{CoreHostAdapters, CoreInstance}, + process_runtime::CoreProcessRuntime, +}; + +use crate::common::global_ctx::GlobalCtxEvent; +#[cfg(feature = "public-ipv6-provider")] +use crate::instance::public_ipv6_provider::runtime_public_ipv6_provider_platform; +use crate::{ + common::global_ctx::ArcGlobalCtx, + common::{config::TomlConfig, global_ctx::GlobalCtx}, + host_runtime::native_host_runtime, + instance::config::{runtime_core_host_config, runtime_peer_credential_storage}, + instance::listeners::RuntimeExternalListenerFactory, + instance::runtime_host::NativeInstanceRuntimeHost, +}; + +use super::host::{NativeInstanceHost, native_instance_host}; +#[cfg(feature = "kcp")] +use crate::gateway::kcp_proxy::KcpProxyService; +#[cfg(feature = "quic")] +use crate::gateway::quic_proxy::QuicProxyService; +use crate::tunnel::protocol::{runtime_client_protocol_upgrader, runtime_server_protocol_upgrader}; +#[cfg(any(feature = "kcp", feature = "quic"))] +use easytier_core::gateway::proxy::wrapped_transport::WrappedTransportEngine; + +pub(crate) type NativeCoreInstance = CoreInstance; + +pub(crate) fn compose_native_core_instance( + config: TomlConfig, + process_runtime: Arc, +) -> anyhow::Result> { + let global_ctx = Arc::new(GlobalCtx::new(config.clone())); + let (packet_sender, packet_receiver) = tokio::sync::mpsc::channel(128); + let mut adapters = + runtime_core_host_adapters(global_ctx.clone(), process_runtime, Arc::new(packet_sender)); + adapters.instance_runtime = NativeInstanceRuntimeHost::new(global_ctx.clone(), packet_receiver); + NativeCoreInstance::from_toml(config, adapters) +} + +impl CoreEventSink for GlobalCtx { + fn emit(&self, event: CoreEvent) { + let event = match event { + CoreEvent::PeerAdded(peer_id) => GlobalCtxEvent::PeerAdded(peer_id), + CoreEvent::PeerRemoved(peer_id) => GlobalCtxEvent::PeerRemoved(peer_id), + CoreEvent::PeerConnAdded(info) => GlobalCtxEvent::PeerConnAdded(info.into()), + CoreEvent::PeerConnRemoved(info) => GlobalCtxEvent::PeerConnRemoved(info.into()), + CoreEvent::CredentialChanged => GlobalCtxEvent::CredentialChanged, + CoreEvent::ManualConnecting { url } => GlobalCtxEvent::Connecting(url), + CoreEvent::ManualConnectError { + url, + ip_version, + error, + } => GlobalCtxEvent::ConnectError(url.to_string(), format!("{ip_version:?}"), error), + CoreEvent::ListenerPlanFailed { url, error } => { + GlobalCtxEvent::ListenerAddFailed(url, error) + } + CoreEvent::ListenerAdded { url, .. } => GlobalCtxEvent::ListenerAdded(url), + CoreEvent::ListenerRemoved { .. } | CoreEvent::ListenerSocketAccepted { .. } => return, + CoreEvent::ListenerAddFailed { + url, + error, + will_retry, + .. + } => { + let message = if will_retry { + format!("error: {error}, retry listen later...") + } else { + format!("error: {error}") + }; + GlobalCtxEvent::ListenerAddFailed(url, message) + } + CoreEvent::ListenerAcceptFailed { url, error } => GlobalCtxEvent::ListenerAcceptFailed( + url, + format!("error: {error}, retry listen later..."), + ), + CoreEvent::ListenerAcceptedSocketHandleFailed { url, error } => { + tracing::error!(%url, %error, "accepted socket handler failed"); + return; + } + CoreEvent::TunnelAccepted { + local_url, + remote_url, + } => GlobalCtxEvent::ConnectionAccepted(local_url, remote_url), + CoreEvent::TunnelAdmissionFailed { + local_url, + remote_url, + error, + } => GlobalCtxEvent::ConnectionError(local_url, remote_url, error), + CoreEvent::UdpPortMappingEstablished { + local_listener, + mapped_listener, + backend, + } => GlobalCtxEvent::ListenerPortMappingEstablished { + local_listener, + mapped_listener, + backend, + }, + CoreEvent::ProxyCidrsUpdated { added, removed } => { + GlobalCtxEvent::ProxyCidrsUpdated(added, removed) + } + CoreEvent::PublicIpv6LeaseChanged { old, new } => { + GlobalCtxEvent::PublicIpv6Changed(old, new) + } + CoreEvent::PublicIpv6RoutesChanged { added, removed } => { + GlobalCtxEvent::PublicIpv6RoutesUpdated(added, removed) + } + CoreEvent::VpnPortalStarted(portal) => GlobalCtxEvent::VpnPortalStarted(portal), + CoreEvent::VpnPortalClientConnected { portal, client } => { + GlobalCtxEvent::VpnPortalClientConnected(portal, client) + } + CoreEvent::VpnPortalClientDisconnected { portal, client } => { + GlobalCtxEvent::VpnPortalClientDisconnected(portal, client) + } + CoreEvent::GatewayPortForwardAdded(config) => { + GlobalCtxEvent::PortForwardAdded(config.into()) + } + }; + self.issue_event(event); + } +} + +#[cfg(feature = "wrapped-transport")] +fn runtime_wrapped_transport_engines() -> WrappedTransportEngines { + #[cfg(feature = "kcp")] + let kcp = Some(Arc::new(KcpProxyService::new()) as Arc); + #[cfg(not(feature = "kcp"))] + let kcp = None; + #[cfg(feature = "quic")] + let quic = Some(Arc::new(QuicProxyService::new()) as Arc); + #[cfg(not(feature = "quic"))] + let quic = None; + + WrappedTransportEngines { kcp, quic } +} + +pub(crate) fn runtime_core_host_adapters( + global_ctx: ArcGlobalCtx, + process_runtime: Arc, + packet_sink: Arc, +) -> CoreHostAdapters { + let host = native_instance_host(global_ctx.clone()); + let runtime_dns = native_host_runtime(); + let mut adapters = CoreHostAdapters::new(host, runtime_dns, packet_sink, process_runtime); + #[cfg(test)] + adapters.replace_stun_provider(Arc::new(crate::common::stun::MockStunInfoCollector { + udp_nat_type: crate::proto::common::NatType::Unknown, + })); + adapters.config = runtime_core_host_config(); + adapters.credential_storage = runtime_peer_credential_storage(&global_ctx); + adapters.events = global_ctx.clone(); + #[cfg(feature = "wrapped-transport")] + { + adapters.wrapped_transports = runtime_wrapped_transport_engines(); + } + adapters.protocol = Some(runtime_client_protocol_upgrader(global_ctx.clone())); + adapters.external_listener_factory = Some(Arc::new(RuntimeExternalListenerFactory)); + adapters.server_protocol = Some(runtime_server_protocol_upgrader(global_ctx.clone())); + #[cfg(feature = "upnp")] + { + adapters.udp_hole_punch_platform = Some( + crate::instance::udp_hole_punch::runtime_udp_hole_punch_platform( + global_ctx.net_ns.clone(), + ), + ); + } + #[cfg(feature = "icmp-proxy")] + { + adapters.icmp_proxy_host = Some(Arc::new(crate::gateway::icmp_proxy::RuntimeIcmpProxyHost)); + } + #[cfg(feature = "proxy-cidr-monitor")] + { + adapters.proxy_cidr_monitor_enabled = true; + } + #[cfg(feature = "public-ipv6-provider")] + { + adapters.public_ipv6_host = Some(global_ctx.clone()); + adapters.public_ipv6_provider = Some(runtime_public_ipv6_provider_platform(&global_ctx)); + } + #[cfg(feature = "wireguard")] + { + use crate::common::config::ConfigLoader as _; + + adapters.vpn_portal = Some(crate::vpn_portal::wireguard::WireGuardPortalHost::new( + global_ctx.clone(), + global_ctx + .config + .get_vpn_portal_config() + .map(|config| config.wireguard_listen), + ) as Arc); + } + adapters +} + +#[cfg(feature = "management")] +pub(crate) fn runtime_one_shot_manual_connector( + global_ctx: ArcGlobalCtx, + config: &TomlConfig, + process_runtime: Arc, +) -> anyhow::Result> { + let normalized = CoreInstanceConfig::from_toml_with_host(config, &runtime_core_host_config())?; + let host = native_instance_host(global_ctx.clone()); + let runtime_dns = native_host_runtime(); + let dns: Arc = runtime_dns.clone(); + let dns_records: Arc = runtime_dns; + Ok(process_runtime.manual_connector( + host, + dns, + dns_records, + runtime_client_protocol_upgrader(global_ctx.clone()), + normalized.connectivity.endpoint_discovery, + normalized.connectivity.manual, + )) +} + +#[cfg(test)] +mod tests { + use std::sync::Arc; + + #[cfg(feature = "kcp")] + use easytier_core::gateway::proxy::wrapped_transport::{ + WrappedTransportConnect, WrappedTransportEngine, + }; + use easytier_core::listener::plan::ListenerRuntimeConfig; + use pnet::packet::{ + ip::IpNextHeaderProtocols, + ipv4::{self, MutableIpv4Packet}, + udp::{self, MutableUdpPacket}, + }; + #[cfg(feature = "kcp")] + use tokio::io::{AsyncReadExt, AsyncWriteExt}; + + #[cfg(feature = "kcp")] + use crate::gateway::kcp_proxy::KcpProxyService; + use crate::{ + common::{config::NetworkIdentity, global_ctx::tests::get_mock_global_ctx_with_network}, + instance::config::test_core_instance_config, + }; + + use super::*; + + fn create_host_packet_channel() -> ( + tokio::sync::mpsc::Sender>, + tokio::sync::mpsc::Receiver>, + ) { + tokio::sync::mpsc::channel(16) + } + + #[cfg(feature = "kcp")] + fn build_native_kcp_test_instance( + global_ctx: ArcGlobalCtx, + packet_sink: tokio::sync::mpsc::Sender>, + listeners: Option, + ) -> anyhow::Result<(Arc, Arc)> { + let mut adapters = runtime_core_host_adapters( + global_ctx.clone(), + CoreProcessRuntime::new(), + Arc::new(packet_sink), + ); + adapters.proxy_cidr_monitor_enabled = false; + let service = Arc::new(KcpProxyService::new()); + adapters.wrapped_transports = WrappedTransportEngines { + kcp: Some(service.clone()), + quic: None, + }; + + let mut config = test_core_instance_config(&global_ctx); + config.connectivity.listeners = listeners; + config.connectivity.startup_plan.gateway = false; + config.connectivity.stun.udp_servers.clear(); + config.connectivity.stun.tcp_servers.clear(); + config.connectivity.stun.udp_v6_servers.clear(); + config.connectivity.manual = Default::default(); + config.connectivity.direct.testing = true; + + let instance = NativeCoreInstance::new(config, adapters)?; + Ok((instance, service)) + } + + #[cfg(feature = "kcp")] + #[tokio::test] + async fn native_kcp_engine_round_trips_through_portable_cores() { + tokio::time::timeout(std::time::Duration::from_secs(20), async { + const REQUEST: &[u8] = b"native-kcp-request"; + const REPLY: &[u8] = b"native-kcp-reply"; + + let global_a = get_mock_global_ctx_with_network(Some(NetworkIdentity::new( + "native-kcp-round-trip".to_owned(), + "shared-secret".to_owned(), + ))); + let global_b = get_mock_global_ctx_with_network(Some(NetworkIdentity::new( + "native-kcp-round-trip".to_owned(), + "shared-secret".to_owned(), + ))); + global_a.set_ipv4(Some("10.250.0.1/24".parse().unwrap())); + global_b.set_ipv4(Some("10.250.0.2/24".parse().unwrap())); + + let mut flags_a = global_a.get_flags(); + flags_a.enable_kcp_proxy = true; + flags_a.disable_kcp_input = true; + flags_a.disable_tcp_hole_punching = true; + flags_a.disable_udp_hole_punching = true; + flags_a.disable_sym_hole_punching = true; + flags_a.disable_upnp = true; + global_a.set_flags(flags_a); + + let mut flags_b = global_b.get_flags(); + flags_b.enable_kcp_proxy = false; + flags_b.disable_kcp_input = false; + flags_b.disable_tcp_hole_punching = true; + flags_b.disable_udp_hole_punching = true; + flags_b.disable_sym_hole_punching = true; + flags_b.disable_upnp = true; + global_b.set_flags(flags_b); + + let (packet_sink_a, _packet_receiver_a) = create_host_packet_channel(); + let (packet_sink_b, _packet_receiver_b) = create_host_packet_channel(); + let (instance_a, kcp_a) = build_native_kcp_test_instance( + global_a.clone(), + packet_sink_a, + Some(ListenerRuntimeConfig::new( + vec!["tcp://127.0.0.1:0".parse().unwrap()], + false, + test_core_instance_config(&global_a) + .connectivity + .direct + .tcp_bind + .context, + )), + ) + .unwrap(); + let (instance_b, _kcp_b) = + build_native_kcp_test_instance(global_b, packet_sink_b, None).unwrap(); + + let (start_a, start_b) = tokio::join!(instance_a.start(), instance_b.start()); + start_a.unwrap(); + start_b.unwrap(); + + let listener = instance_a.running_listeners().pop().unwrap(); + instance_b.add_connector(listener).unwrap(); + let peer_a_id = instance_a.peer_id(); + let peer_b_id = instance_b.peer_id(); + loop { + let a_peers = instance_a.connected_peers().await; + let b_peers = instance_b.connected_peers().await; + if a_peers.contains(&peer_b_id) && b_peers.contains(&peer_a_id) { + break; + } + tokio::time::sleep(std::time::Duration::from_millis(20)).await; + } + + let echo_listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); + let echo_addr = echo_listener.local_addr().unwrap(); + let responder = tokio::spawn(async move { + let (mut socket, _) = echo_listener.accept().await.unwrap(); + let mut request = [0; REQUEST.len()]; + socket.read_exact(&mut request).await.unwrap(); + assert_eq!(&request, REQUEST); + socket.write_all(REPLY).await.unwrap(); + }); + + let mut stream = kcp_a + .connect_source(WrappedTransportConnect { + my_peer_id: peer_a_id, + dst_peer_id: peer_b_id, + src: "10.250.0.1:40000".parse().unwrap(), + dst: echo_addr, + }) + .await + .unwrap(); + stream.write_all(REQUEST).await.unwrap(); + let mut reply = [0; REPLY.len()]; + stream.read_exact(&mut reply).await.unwrap(); + assert_eq!(&reply, REPLY); + responder.await.unwrap(); + + instance_b.stop().await; + instance_a.stop().await; + }) + .await + .expect("native KCP round trip timed out"); + } + + #[tokio::test] + async fn portable_core_instances_connect_through_core_tcp_listener() { + let global_a = get_mock_global_ctx_with_network(Some(NetworkIdentity::new( + "portable-connect-listen".to_owned(), + "shared-secret".to_owned(), + ))); + let global_b = get_mock_global_ctx_with_network(Some(NetworkIdentity::new( + "portable-connect-listen".to_owned(), + "shared-secret".to_owned(), + ))); + global_a.set_ipv4(Some("10.250.0.1/24".parse().unwrap())); + global_b.set_ipv4(Some("10.250.0.2/24".parse().unwrap())); + let (packet_sink_a, _packet_receiver_a) = create_host_packet_channel(); + let (packet_sink_b, mut packet_receiver_b) = create_host_packet_channel(); + let mut config_a = test_core_instance_config(&global_a); + config_a.connectivity.initial_peers.clear(); + config_a.connectivity.listeners = Some(ListenerRuntimeConfig::new( + vec!["tcp://127.0.0.1:0".parse().unwrap()], + false, + config_a.connectivity.direct.tcp_bind.context.clone(), + )); + config_a.connectivity.runtime = Default::default(); + config_a.connectivity.stun.udp_servers.clear(); + config_a.connectivity.stun.tcp_servers.clear(); + config_a.connectivity.stun.udp_v6_servers.clear(); + config_a.connectivity.manual = Default::default(); + config_a.connectivity.direct.testing = true; + let instance_a = NativeCoreInstance::new( + config_a, + runtime_core_host_adapters( + global_a.clone(), + CoreProcessRuntime::new(), + Arc::new(packet_sink_a), + ), + ) + .unwrap(); + + let mut config_b = test_core_instance_config(&global_b); + config_b.connectivity.initial_peers.clear(); + config_b.connectivity.listeners = None; + config_b.connectivity.runtime = Default::default(); + config_b.connectivity.stun.udp_servers.clear(); + config_b.connectivity.stun.tcp_servers.clear(); + config_b.connectivity.stun.udp_v6_servers.clear(); + config_b.connectivity.manual = Default::default(); + config_b.connectivity.direct.testing = true; + let instance_b = NativeCoreInstance::new( + config_b, + runtime_core_host_adapters( + global_b.clone(), + CoreProcessRuntime::new(), + Arc::new(packet_sink_b), + ), + ) + .unwrap(); + + let (start_a, start_b) = tokio::join!(instance_a.start(), instance_b.start()); + start_a.unwrap(); + start_b.unwrap(); + let listener = instance_a.running_listeners().pop().unwrap(); + instance_b.add_connector(listener).unwrap(); + + let peer_a_id = instance_a.peer_id(); + let peer_b_id = instance_b.peer_id(); + tokio::time::timeout(std::time::Duration::from_secs(10), async { + loop { + let a_peers = instance_a.connected_peers().await; + let b_peers = instance_b.connected_peers().await; + if a_peers.contains(&peer_b_id) && b_peers.contains(&peer_a_id) { + break; + } + tokio::time::sleep(std::time::Duration::from_millis(20)).await; + } + }) + .await + .expect("portable core instances did not connect through the core listener"); + + let source_ip = "10.250.0.1".parse().unwrap(); + let destination_ip = "10.250.0.2".parse().unwrap(); + let mut ip_packet = vec![0u8; 28]; + { + let mut ipv4 = MutableIpv4Packet::new(&mut ip_packet).unwrap(); + ipv4.set_version(4); + ipv4.set_header_length(5); + ipv4.set_total_length(28); + ipv4.set_ttl(64); + ipv4.set_next_level_protocol(IpNextHeaderProtocols::Udp); + ipv4.set_source(source_ip); + ipv4.set_destination(destination_ip); + } + { + let mut udp = MutableUdpPacket::new(&mut ip_packet[20..]).unwrap(); + udp.set_source(10000); + udp.set_destination(10001); + udp.set_length(8); + udp.set_checksum(udp::ipv4_checksum( + &udp.to_immutable(), + &source_ip, + &destination_ip, + )); + } + { + let mut ipv4 = MutableIpv4Packet::new(&mut ip_packet).unwrap(); + ipv4.set_checksum(ipv4::checksum(&ipv4.to_immutable())); + } + let received = tokio::time::timeout(std::time::Duration::from_secs(10), async { + loop { + instance_a + .packet_plane() + .send_ip_packet(ip_packet.clone()) + .await + .unwrap(); + match tokio::time::timeout( + std::time::Duration::from_millis(100), + packet_receiver_b.recv(), + ) + .await + { + Ok(Some(packet)) => break packet, + Ok(None) => panic!("portable host packet sink closed"), + Err(_) => {} + } + } + }) + .await + .expect("portable host packet sink did not receive the IP packet"); + assert_eq!(received, ip_packet); + + instance_b.stop().await; + instance_a.stop().await; + } +} diff --git a/easytier/src/instance/config.rs b/easytier/src/instance/config.rs new file mode 100644 index 00000000..188ed159 --- /dev/null +++ b/easytier/src/instance/config.rs @@ -0,0 +1,184 @@ +use std::sync::Arc; + +use easytier_core::{ + config::peers::HostRoutingPolicy, instance::CoreInstanceHostConfig, + peers::credential_manager::CredentialStorage, +}; +use strum::VariantArray as _; + +#[cfg(feature = "management")] +use crate::common::credential_manager::runtime_credential_storage; +#[cfg(test)] +use crate::common::global_ctx::GlobalCtxEvent; +use crate::{ + common::{constants::EASYTIER_VERSION, global_ctx::ArcGlobalCtx}, + tunnel::IpScheme, +}; + +/// Projects only native Host policy and build capabilities. All TOML-derived +/// Instance configuration is normalized by `easytier-core`. +pub(crate) fn runtime_core_host_config() -> CoreInstanceHostConfig { + let hostname = gethostname::gethostname().to_string_lossy().to_string(); + CoreInstanceHostConfig { + hostname_fallback: (!hostname.is_empty()).then_some(hostname), + host_routing: HostRoutingPolicy { + local_exit_node_fallback: cfg!(target_env = "ohos"), + }, + force_exit_node: cfg!(target_env = "ohos"), + allow_interface_bind: !cfg!(any( + target_os = "android", + target_os = "ios", + all(target_os = "macos", feature = "macos-ne"), + target_env = "ohos" + )), + smoltcp_available: cfg!(feature = "smoltcp"), + requires_smoltcp: cfg!(any( + target_os = "android", + target_os = "ios", + all(target_os = "macos", feature = "macos-ne"), + target_env = "ohos" + )), + icmp_failure_is_fatal: cfg!(not(any( + target_os = "android", + target_os = "ios", + all(target_os = "macos", feature = "macos-ne"), + target_env = "ohos" + ))), + public_ipv6_provider_supported: cfg!(target_os = "linux"), + gateway_enabled: cfg!(feature = "socks5"), + easytier_version: EASYTIER_VERSION.to_owned(), + endpoint_protocols: IpScheme::VARIANTS.iter().map(ToString::to_string).collect(), + } +} + +pub(crate) fn runtime_peer_credential_storage( + global_ctx: &ArcGlobalCtx, +) -> Option> { + #[cfg(not(feature = "management"))] + { + let _ = global_ctx; + None + } + #[cfg(feature = "management")] + runtime_credential_storage(global_ctx.config.get_credential_file()) +} + +#[cfg(test)] +pub(crate) fn test_core_instance_config( + global_ctx: &ArcGlobalCtx, +) -> easytier_core::instance::CoreInstanceConfig { + use easytier_core::config::toml::{ConfigLoader as _, TomlConfig}; + + let config = TomlConfig::new_from_str(&global_ctx.config.dump()) + .expect("test configuration should round-trip through TOML"); + let mut host = runtime_core_host_config(); + let hostname = global_ctx.get_hostname(); + host.hostname_fallback = (!hostname.is_empty()).then_some(hostname); + easytier_core::instance::CoreInstanceConfig::from_toml_with_host(&config, &host) + .expect("test configuration should normalize") +} + +#[cfg(test)] +pub(crate) fn test_runtime_instance_config( + global_ctx: &ArcGlobalCtx, +) -> easytier_core::config::runtime::CoreInstanceRuntimeConfig { + let config = test_core_instance_config(global_ctx); + easytier_core::config::runtime::CoreInstanceRuntimeConfig { + services: config.connectivity.runtime, + peer: Arc::new(config.peer.snapshot), + } +} + +#[cfg(test)] +mod tests { + use easytier_core::{ + config::toml::TomlConfig, + events::{CoreEvent, CoreEventSink}, + }; + + use crate::common::global_ctx::tests::get_mock_global_ctx; + + use super::*; + + #[test] + fn native_host_config_contains_only_platform_policy() { + let config = runtime_core_host_config(); + let hostname = gethostname::gethostname().to_string_lossy().to_string(); + + assert_eq!( + config.hostname_fallback, + (!hostname.is_empty()).then_some(hostname) + ); + assert_eq!(config.gateway_enabled, cfg!(feature = "socks5")); + assert_eq!(config.smoltcp_available, cfg!(feature = "smoltcp")); + assert_eq!( + config.public_ipv6_provider_supported, + cfg!(target_os = "linux") + ); + assert_eq!(config.easytier_version, EASYTIER_VERSION); + } + + #[test] + fn clearing_configured_hostname_does_not_reuse_the_old_value() { + let host = runtime_core_host_config(); + let initial = TomlConfig::new_from_str("hostname = \"configured-host\"").unwrap(); + let cleared = TomlConfig::new_from_str("hostname = \"\"").unwrap(); + + let initial = + easytier_core::instance::CoreInstanceConfig::from_toml_with_host(&initial, &host) + .unwrap(); + let cleared = + easytier_core::instance::CoreInstanceConfig::from_toml_with_host(&cleared, &host) + .unwrap(); + + assert_eq!( + initial.peer.snapshot.runtime.core.node.hostname.as_deref(), + Some("configured-host") + ); + assert_ne!( + cleared.peer.snapshot.runtime.core.node.hostname.as_deref(), + Some("configured-host") + ); + assert_eq!( + cleared.peer.snapshot.runtime.core.node.hostname.as_deref(), + host.hostname_fallback.as_deref() + ); + } + + #[test] + fn test_config_uses_current_global_context_hostname_as_fallback() { + let global_ctx = get_mock_global_ctx(); + global_ctx.set_hostname("test-hostname".to_owned()); + + let config = test_core_instance_config(&global_ctx); + + assert_eq!( + config.peer.snapshot.runtime.core.node.hostname.as_deref(), + Some("test-hostname") + ); + } + + #[tokio::test] + async fn core_event_sink_projects_peer_events_to_global_context() { + let global_ctx = get_mock_global_ctx(); + let mut events = global_ctx.subscribe(); + CoreEventSink::emit(global_ctx.as_ref(), CoreEvent::PeerAdded(7)); + + assert!(matches!( + events.recv().await.unwrap(), + GlobalCtxEvent::PeerAdded(7) + )); + } + + #[tokio::test] + async fn core_event_sink_projects_credential_changes_to_global_context() { + let global_ctx = get_mock_global_ctx(); + let mut events = global_ctx.subscribe(); + CoreEventSink::emit(global_ctx.as_ref(), CoreEvent::CredentialChanged); + + assert!(matches!( + events.recv().await.unwrap(), + GlobalCtxEvent::CredentialChanged + )); + } +} diff --git a/easytier/src/instance/config_storage.rs b/easytier/src/instance/config_storage.rs new file mode 100644 index 00000000..d27e4034 --- /dev/null +++ b/easytier/src/instance/config_storage.rs @@ -0,0 +1,42 @@ +use std::path::Path; + +use easytier_core::management::{ConfigFileControl, ConfigFilePermission, ConfigFileStorage}; + +#[derive(Default)] +pub(crate) struct NativeConfigFileStorage; + +#[async_trait::async_trait] +impl ConfigFileStorage for NativeConfigFileStorage { + async fn inspect(&self, path: &Path) -> ConfigFileControl { + let read_only = tokio::fs::metadata(path) + .await + .map(|metadata| metadata.permissions().readonly()) + .unwrap_or(true); + ConfigFileControl::new( + Some(path.to_owned()), + if read_only { + ConfigFilePermission::from(ConfigFilePermission::READ_ONLY) + } else { + ConfigFilePermission::default() + }, + ) + } + + async fn read(&self, path: &Path) -> anyhow::Result>> { + match tokio::fs::read(path).await { + Ok(contents) => Ok(Some(contents)), + Err(error) if error.kind() == std::io::ErrorKind::NotFound => Ok(None), + Err(error) => Err(error.into()), + } + } + + async fn write(&self, path: &Path, contents: &[u8]) -> anyhow::Result<()> { + tokio::fs::write(path, contents).await?; + Ok(()) + } + + async fn remove(&self, path: &Path) -> anyhow::Result<()> { + tokio::fs::remove_file(path).await?; + Ok(()) + } +} diff --git a/easytier/src/instance/dns_server/client_instance.rs b/easytier/src/instance/dns_server/client_instance.rs index f1d5f35b..a9c7c11c 100644 --- a/easytier/src/instance/dns_server/client_instance.rs +++ b/easytier/src/instance/dns_server/client_instance.rs @@ -1,99 +1,106 @@ use std::{sync::Arc, time::Duration}; +use easytier_core::gateway::magic_dns::{ + MagicDnsRoutePublisher, MagicDnsRouteSnapshot, run_magic_dns_route_publisher, +}; +use easytier_core::instance::CorePacketPlane; use tokio::task::JoinSet; -use crate::{ - peers::peer_manager::PeerManager, - proto::{ - api::instance::Route, - common::Void, - magic_dns::{ - HandshakeRequest, MagicDnsServerRpc, MagicDnsServerRpcClientFactory, - UpdateDnsRecordRequest, - }, - rpc_impl::standalone::StandAloneClient, - rpc_types::controller::BaseController, +use crate::proto::{ + api::instance::Route, + common::Void, + magic_dns::{ + HandshakeRequest, MagicDnsServerRpc, MagicDnsServerRpcClientFactory, UpdateDnsRecordRequest, }, - tunnel::tcp::TcpTunnelConnector, + rpc::standalone::{RuntimeRpcClient, runtime_rpc_client}, + rpc_types::controller::BaseController, }; use super::MAGIC_DNS_INSTANCE_ADDR; pub struct MagicDnsClientInstance { - rpc_client: StandAloneClient, + rpc_client: RuntimeRpcClient, rpc_stub: Option + Send>>, - peer_mgr: Arc, + route_source: Arc, tasks: JoinSet<()>, } +struct RpcMagicDnsRoutePublisher { + rpc_stub: Box + Send>, +} + +#[async_trait::async_trait] +impl MagicDnsRoutePublisher for RpcMagicDnsRoutePublisher { + async fn handshake(&mut self) -> anyhow::Result<()> { + self.rpc_stub + .handshake(BaseController::default(), HandshakeRequest::default()) + .await?; + Ok(()) + } + + async fn heartbeat(&mut self) -> anyhow::Result<()> { + self.rpc_stub + .heartbeat(BaseController::default(), Void::default()) + .await?; + Ok(()) + } + + async fn publish(&mut self, snapshot: &MagicDnsRouteSnapshot) -> anyhow::Result<()> { + let request = UpdateDnsRecordRequest { + routes: snapshot + .routes + .iter() + .map(|route| Route { + hostname: route.hostname.clone(), + ipv4_addr: route.ipv4_addr, + ..Default::default() + }) + .collect(), + zone: snapshot.zone.clone(), + }; + tracing::debug!( + "MagicDnsClientInstance::update_dns_task: update dns records: {:?}", + request + ); + self.rpc_stub + .update_dns_record(BaseController::default(), request) + .await?; + Ok(()) + } +} + impl MagicDnsClientInstance { - pub async fn new(peer_mgr: Arc) -> Result { - let tcp_connector = TcpTunnelConnector::new(MAGIC_DNS_INSTANCE_ADDR.parse().unwrap()); - let mut rpc_client = StandAloneClient::new(tcp_connector); + pub(crate) async fn new(route_source: Arc) -> Result { + let mut rpc_client = runtime_rpc_client(MAGIC_DNS_INSTANCE_ADDR.parse().unwrap()); let rpc_stub = rpc_client .scoped_client::>("".to_string()) .await?; Ok(MagicDnsClientInstance { rpc_client, rpc_stub: Some(rpc_stub), - peer_mgr, + route_source, tasks: JoinSet::new(), }) } async fn update_dns_task( - peer_mgr: Arc, + route_source: Arc, rpc_stub: Box + Send>, ) -> Result<(), anyhow::Error> { - let mut prev_last_update = None; - rpc_stub - .handshake(BaseController::default(), HandshakeRequest::default()) - .await?; - loop { - rpc_stub - .heartbeat(BaseController::default(), Void::default()) - .await?; - - let last_update = peer_mgr.get_route_peer_info_last_update_time().await; - if Some(last_update) == prev_last_update { - tokio::time::sleep(Duration::from_millis(500)).await; - continue; - } - - let mut routes = peer_mgr.list_routes().await; - // add self as a route - let ctx = peer_mgr.get_global_ctx(); - routes.push(Route { - hostname: ctx.get_hostname(), - ipv4_addr: ctx.get_ipv4().map(Into::into), - ..Default::default() - }); - // Use configured tld_dns_zone (always set by default) - let flags = ctx.config.get_flags(); - let req = UpdateDnsRecordRequest { - routes, - zone: flags.tld_dns_zone.clone(), - }; - tracing::debug!( - "MagicDnsClientInstance::update_dns_task: update dns records: {:?}", - req - ); - rpc_stub - .update_dns_record(BaseController::default(), req) - .await?; - - let last_update_after_rpc = peer_mgr.get_route_peer_info_last_update_time().await; - if last_update_after_rpc == last_update { - prev_last_update = Some(last_update); - } - } + let mut publisher = RpcMagicDnsRoutePublisher { rpc_stub }; + run_magic_dns_route_publisher( + route_source.as_ref(), + &mut publisher, + Duration::from_millis(500), + ) + .await } pub async fn run_and_wait(&mut self) { let rpc_stub = self.rpc_stub.take().unwrap(); - let peer_mgr = self.peer_mgr.clone(); + let route_source = self.route_source.clone(); self.tasks.spawn(async move { - let ret = Self::update_dns_task(peer_mgr, rpc_stub).await; + let ret = Self::update_dns_task(route_source, rpc_stub).await; if let Err(e) = ret { tracing::error!("MagicDnsServerInstanceData::run_and_wait: {:?}", e); } diff --git a/easytier/src/instance/dns_server/mod.rs b/easytier/src/instance/dns_server/mod.rs index 1bec8c06..4719a1d0 100644 --- a/easytier/src/instance/dns_server/mod.rs +++ b/easytier/src/instance/dns_server/mod.rs @@ -1,21 +1,22 @@ // This module is copy and modified from https://github.com/fanyang89/libdns -#[cfg(feature = "magic-dns")] +#[cfg(all(feature = "magic-dns", feature = "tun"))] pub(crate) mod config; -#[cfg(feature = "magic-dns")] +#[cfg(all(feature = "magic-dns", feature = "tun"))] pub(crate) mod server; -#[cfg(feature = "magic-dns")] +#[cfg(all(feature = "magic-dns", feature = "tun"))] pub mod client_instance; -#[cfg(feature = "magic-dns")] +#[cfg(all(feature = "magic-dns", feature = "tun"))] pub mod runner; -#[cfg(feature = "magic-dns")] +#[cfg(all(feature = "magic-dns", feature = "tun"))] pub mod server_instance; -#[cfg(feature = "magic-dns")] +#[cfg(all(feature = "magic-dns", feature = "tun"))] pub mod system_config; #[cfg(all(test, feature = "tun", feature = "magic-dns"))] mod tests; pub static MAGIC_DNS_INSTANCE_ADDR: &str = "tcp://127.0.0.1:49813"; +pub static MAGIC_DNS_INSTANCE_SOCKET_ADDR: &str = "127.0.0.1:49813"; pub static MAGIC_DNS_FAKE_IP: &str = "100.100.100.101"; -pub static DEFAULT_ET_DNS_ZONE: &str = "et.net."; +pub use easytier_core::config::toml::DEFAULT_ET_DNS_ZONE; diff --git a/easytier/src/instance/dns_server/runner.rs b/easytier/src/instance/dns_server/runner.rs index e777ca04..76791c68 100644 --- a/easytier/src/instance/dns_server/runner.rs +++ b/easytier/src/instance/dns_server/runner.rs @@ -1,25 +1,28 @@ use cidr::Ipv4Inet; use tokio_util::sync::CancellationToken; -use crate::peers::peer_manager::PeerManager; use std::{net::Ipv4Addr, sync::Arc, time::Duration}; -use super::{client_instance::MagicDnsClientInstance, server_instance::MagicDnsServerInstance}; +use easytier_core::instance::CorePacketPlane; -static DEFAULT_ET_DNS_ZONE: &str = "et.net."; +use crate::common::global_ctx::ArcGlobalCtx; + +use super::{client_instance::MagicDnsClientInstance, server_instance::MagicDnsServerInstance}; pub struct DnsRunner { client: Option, server: Option, - peer_mgr: Arc, + packet_plane: Arc, + global_ctx: ArcGlobalCtx, tun_dev: Option, tun_inet: Ipv4Inet, fake_ip: Ipv4Addr, } impl DnsRunner { - pub fn new( - peer_mgr: Arc, + pub(crate) fn new( + packet_plane: Arc, + global_ctx: ArcGlobalCtx, tun_dev: Option, tun_inet: Ipv4Inet, fake_ip: Ipv4Addr, @@ -27,7 +30,8 @@ impl DnsRunner { Self { client: None, server: None, - peer_mgr, + packet_plane, + global_ctx, tun_dev, tun_inet, fake_ip, @@ -44,7 +48,8 @@ impl DnsRunner { async fn run_once(&mut self) -> anyhow::Result<()> { // try server first match MagicDnsServerInstance::new( - self.peer_mgr.clone(), + self.packet_plane.clone(), + self.global_ctx.clone(), self.tun_dev.clone(), self.tun_inet, self.fake_ip, @@ -61,7 +66,7 @@ impl DnsRunner { } // every runner must run a client - let client = MagicDnsClientInstance::new(self.peer_mgr.clone()).await?; + let client = MagicDnsClientInstance::new(self.packet_plane.clone()).await?; self.client = Some(client); self.client.as_mut().unwrap().run_and_wait().await; diff --git a/easytier/src/instance/dns_server/server.rs b/easytier/src/instance/dns_server/server.rs index f9b434bb..e9cd5ad3 100644 --- a/easytier/src/instance/dns_server/server.rs +++ b/easytier/src/instance/dns_server/server.rs @@ -1,5 +1,4 @@ use anyhow::{Context, Result}; -use hickory_proto::op::Edns; use hickory_proto::rr; use hickory_proto::rr::LowerName; use hickory_resolver::config::ResolverOpts; @@ -10,14 +9,12 @@ use hickory_server::authority::{AuthorityObject, Catalog, ZoneType}; use hickory_server::server::{Request, RequestHandler, ResponseHandler, ResponseInfo}; use hickory_server::store::forwarder::ForwardConfig; use hickory_server::store::{forwarder::ForwardAuthority, in_memory::InMemoryAuthority}; -use std::io; use std::net::SocketAddr; use std::str::FromStr; use std::sync::Arc; use std::time::Duration; use tokio::net::{TcpListener, UdpSocket}; -use tokio::sync::{RwLock, RwLockReadGuard, RwLockWriteGuard}; -use tokio::task::JoinSet; +use tokio::sync::{RwLock, RwLockReadGuard}; use crate::common::dns::get_default_resolver_config; @@ -28,8 +25,6 @@ pub struct Server { catalog: Arc>, general_config: GeneralConfig, udp_local_addr: Option, - tcp_local_addr: Option, - tasks: JoinSet<()>, } struct CatalogRequestHandler { @@ -124,19 +119,14 @@ impl Server { catalog, general_config: config.general().clone(), udp_local_addr: None, - tcp_local_addr: None, - tasks: JoinSet::new(), }) } + #[cfg(test)] pub fn udp_local_addr(&self) -> Option { self.udp_local_addr } - pub fn tcp_local_addr(&self) -> Option { - self.tcp_local_addr - } - pub async fn register_udp_socket(&mut self, address: String) -> Result { let bind_addr = SocketAddr::from_str(&address) .with_context(|| format!("DNS Server failed to parse address {}", address))?; @@ -185,7 +175,6 @@ impl Server { let tcp_listener = TcpListener::bind(address.clone()) .await .with_context(|| format!("DNS Server failed to bind TCP address {}", address))?; - self.tcp_local_addr = Some(tcp_listener.local_addr()?); self.server .register_listener(tcp_listener, Duration::from_secs(5)); } @@ -198,6 +187,7 @@ impl Server { Ok(()) } + #[cfg(test)] pub async fn shutdown(&mut self) -> Result<()> { self.server.shutdown_gracefully().await?; Ok(()) @@ -207,47 +197,9 @@ impl Server { self.catalog.write().await.upsert(name, vec![authority]); } - pub async fn remove(&self, name: &LowerName) -> Option>> { - self.catalog.write().await.remove(name) - } - - pub async fn update( - &self, - update: &Request, - response_edns: Option, - response_handle: R, - ) -> io::Result { - self.catalog - .write() - .await - .update(update, response_edns, response_handle) - .await - } - - pub async fn contains(&self, name: &LowerName) -> bool { - self.catalog.read().await.contains(name) - } - - pub async fn lookup( - &self, - request: &Request, - response_edns: Option, - response_handle: R, - ) -> ResponseInfo { - self.catalog - .read() - .await - .lookup(request, response_edns, response_handle) - .await - } - pub async fn read_catalog(&self) -> RwLockReadGuard<'_, Catalog> { self.catalog.read().await } - - pub async fn write_catalog(&self) -> RwLockWriteGuard<'_, Catalog> { - self.catalog.write().await - } } #[cfg(test)] diff --git a/easytier/src/instance/dns_server/server_instance.rs b/easytier/src/instance/dns_server/server_instance.rs index 91d21514..494056e4 100644 --- a/easytier/src/instance/dns_server/server_instance.rs +++ b/easytier/src/instance/dns_server/server_instance.rs @@ -7,72 +7,59 @@ // all the clients will exit and let the easytier instance to launch a new server instance. use super::{ - MAGIC_DNS_INSTANCE_ADDR, + MAGIC_DNS_INSTANCE_SOCKET_ADDR, config::{GeneralConfigBuilder, RunConfigBuilder}, server::Server, system_config::{OSConfig, SystemConfig}, }; use crate::{ common::{ - PeerId, + global_ctx::ArcGlobalCtx, ifcfg::{IfConfiger, IfConfiguerTrait}, }, instance::dns_server::{ config::{Record, RecordBuilder, RecordType}, server::build_authority, }, - peers::{NicPacketFilter, peer_manager::PeerManager}, proto::{ - api::instance::Route, common::{TunnelInfo, Void}, magic_dns::{ DnsRecord, DnsRecordA, DnsRecordList, GetDnsRecordResponse, HandshakeRequest, HandshakeResponse, MagicDnsServerRpc, MagicDnsServerRpcServer, UpdateDnsRecordRequest, dns_record::{self}, }, - rpc_impl::standalone::{RpcServerHook, StandAloneServer}, + rpc::standalone::{ + RpcServerHook, RuntimeRpcListener, StandAloneServer, runtime_rpc_listener, + }, rpc_types::controller::{BaseController, Controller}, }, - tunnel::{packet_def::ZCPacket, tcp::TcpTunnelListener}, }; use anyhow::Context; use cidr::Ipv4Inet; -use dashmap::DashMap; +use easytier_core::gateway::magic_dns::{ + MagicDnsQuery, MagicDnsQueryResolver, MagicDnsRecordStore, MagicDnsResolverRegistration, + MagicDnsRoute, +}; +use easytier_core::instance::CorePacketPlane; use hickory_proto::rr::LowerName; use hickory_proto::serialize::binary::{BinDecodable, BinEncoder}; use hickory_server::authority::{MessageRequest, MessageResponse}; use hickory_server::server::{Request, RequestHandler, ResponseHandler, ResponseInfo}; -use multimap::MultiMap; -use pnet::packet::icmp::{IcmpTypes, MutableIcmpPacket}; -use pnet::packet::ipv4::Ipv4Packet; -use pnet::packet::udp::UdpPacket; -use pnet::packet::{ - MutablePacket, Packet, icmp, - ip::IpNextHeaderProtocols, - ipv4::{self, MutableIpv4Packet}, - udp::{self, MutableUdpPacket}, -}; -use std::net::{SocketAddr, SocketAddrV4}; use std::sync::Mutex; use std::{collections::BTreeMap, io, net::Ipv4Addr, str::FromStr, sync::Arc, time::Duration}; -static NIC_PIPELINE_NAME: &str = "magic_dns_server"; - pub(super) struct MagicDnsServerInstanceData { dns_server: Server, tun_dev: Option, - tun_ip: Ipv4Addr, fake_ip: Ipv4Addr, - my_peer_id: PeerId, - - // zone -> (tunnel remote addr -> route) - route_infos: DashMap>, + route_store: MagicDnsRecordStore, + record_apply: tokio::sync::Mutex<()>, system_config: Option>, } impl MagicDnsServerInstanceData { - pub async fn update_dns_records<'a, T: Iterator>( + pub async fn update_dns_records<'a, T: Iterator>( &self, routes: T, zone: &str, @@ -83,7 +70,7 @@ impl MagicDnsServerInstanceData { continue; } - let Some(ipv4_addr) = route.ipv4_addr.unwrap_or_default().address else { + let Some(ipv4_addr) = route.ipv4_addr else { continue; }; @@ -130,10 +117,9 @@ impl MagicDnsServerInstanceData { } pub async fn update(&self) { - for item in self.route_infos.iter() { - let zone = item.key(); - let route_iter = item.value().flat_iter().map(|x| x.1); - if let Err(e) = self.update_dns_records(route_iter, zone).await { + let snapshot = self.route_store.snapshot(); + for (zone, routes) in &snapshot.zones { + if let Err(e) = self.update_dns_records(routes.iter(), zone).await { tracing::error!("Failed to update DNS records for zone {}: {:?}", zone, e); } } @@ -141,7 +127,7 @@ impl MagicDnsServerInstanceData { async fn keep_zone_authoritative(&self, zone: &str) { if let Err(e) = self - .update_dns_records(std::iter::empty::<&Route>(), zone) + .update_dns_records(std::iter::empty::<&MagicDnsRoute>(), zone) .await { tracing::error!( @@ -194,24 +180,22 @@ impl MagicDnsServerRpc for MagicDnsServerInstanceData { let Some(remote_addr) = &tunnel_info.remote_addr else { return Err(anyhow::anyhow!("No remote addr").into()); }; + let _apply = self.record_apply.lock().await; let zone = input.zone.clone(); let remote_addr: url::Url = remote_addr.clone().into(); - let mut zone_removed = false; - - if let Some(mut routes_by_addr) = self.route_infos.get_mut(&zone) { - routes_by_addr.remove(&remote_addr); - if !input.routes.is_empty() { - routes_by_addr.insert_many(remote_addr, input.routes); - } - zone_removed = routes_by_addr.is_empty(); - } else if !input.routes.is_empty() { - let mut routes_by_addr = MultiMap::new(); - routes_by_addr.insert_many(remote_addr, input.routes); - self.route_infos.insert(zone.clone(), routes_by_addr); - } + let routes = input + .routes + .into_iter() + .map(|route| MagicDnsRoute { + hostname: route.hostname, + ipv4_addr: route.ipv4_addr.unwrap_or_default().address.map(Into::into), + }) + .collect(); + let zone_removed = + self.route_store + .replace_client_routes(zone.clone(), remote_addr.to_string(), routes); if zone_removed { - self.route_infos.remove(&zone); self.keep_zone_authoritative(&zone).await; } @@ -225,20 +209,18 @@ impl MagicDnsServerRpc for MagicDnsServerInstanceData { _input: Void, ) -> crate::proto::rpc_types::error::Result { let mut ret = BTreeMap::new(); - for item in self.route_infos.iter() { - let zone = item.key(); - let routes = item.value(); + for (zone, routes) in self.route_store.snapshot().zones { let mut dns_records = DnsRecordList::default(); - for route in routes.iter().map(|x| x.1) { + for route in routes { dns_records.records.push(DnsRecord { record: Some(dns_record::Record::A(DnsRecordA { name: format!("{}.{}", route.hostname, zone), - value: route.ipv4_addr.unwrap_or_default().address, + value: route.ipv4_addr.map(Into::into), ttl: 1, })), }); } - ret.insert(zone.clone(), dns_records); + ret.insert(zone, dns_records); } Ok(GetDnsRecordResponse { records: ret }) } @@ -289,160 +271,33 @@ impl ResponseHandler for ResponseWrapper { } impl MagicDnsServerInstanceData { - /// Replace content of incoming UDP DNS request and ICMP echo request packet with reply data, - /// and swap source and destination IP addresses to send it back. - async fn handle_ip_packet(&self, zc_packet: &mut ZCPacket) -> Option<()> { - let (ip_header_length, ip_protocol, src_ip, dst_ip) = { - let ip_packet = Ipv4Packet::new(zc_packet.payload())?; + async fn resolve_query_inner(&self, query: MagicDnsQuery) -> Option> { + let request = Request::new( + MessageRequest::from_bytes(&query.payload).ok()?, + query.source, + hickory_proto::xfer::Protocol::Udp, + ); + let response = Arc::new(Mutex::new(Vec::with_capacity(512))); - if ip_packet.get_version() != 4 { - return None; - } - - ( - ip_packet.get_header_length() as usize * 4, - ip_packet.get_next_level_protocol(), - ip_packet.get_source(), - ip_packet.get_destination(), + self.dns_server + .read_catalog() + .await + .handle_request( + &request, + ResponseWrapper { + response: response.clone(), + }, ) - }; + .await; - if dst_ip != self.fake_ip { - return None; - } - - match ip_protocol { - IpNextHeaderProtocols::Udp => { - self.handle_udp_packet(zc_packet, ip_header_length, src_ip, dst_ip) - .await?; - } - IpNextHeaderProtocols::Icmp => { - self.handle_icmp_packet(zc_packet, ip_header_length)?; - } - _ => { - return None; - } - } - - let mut ip_packet = MutableIpv4Packet::new(zc_packet.mut_payload())?; - ip_packet.set_source(dst_ip); - ip_packet.set_destination(src_ip); - - ip_packet.set_checksum(ipv4::checksum(&ip_packet.to_immutable())); - - zc_packet.mut_peer_manager_header().unwrap().to_peer_id = self.my_peer_id.into(); - - Some(()) - } - - /// Extract the DNS request message and send it to the hickory-dns server instance. - /// Replace the content of the UDP packet with the response message. - async fn handle_udp_packet( - &self, - zc_packet: &mut ZCPacket, - ip_header_length: usize, - src_ip: Ipv4Addr, - dst_ip: Ipv4Addr, - ) -> Option<()> { - let (src_port, dst_port, request, request_length) = { - let udp_packet = UdpPacket::new(&zc_packet.payload()[ip_header_length..])?; - - let src_port = udp_packet.get_source(); - let dst_port = udp_packet.get_destination(); - - // Remove this to support any UDP port - if dst_port != 53 { - return None; - } - - let request_payload = udp_packet.payload(); - - ( - src_port, - dst_port, - Request::new( - MessageRequest::from_bytes(request_payload).ok()?, - SocketAddr::from(SocketAddrV4::new(src_ip, src_port)), - hickory_proto::xfer::Protocol::Udp, - ), - request_payload.len(), - ) - }; - - let response_payload = { - let response_payload_arc = Arc::new(Mutex::new(Vec::with_capacity(512))); - - self.dns_server - .read_catalog() - .await - .handle_request( - &request, - ResponseWrapper { - response: response_payload_arc.clone(), - }, - ) - .await; - - Arc::into_inner(response_payload_arc)?.into_inner().ok()? - }; - - let response_length = response_payload.len(); - let delta_length = response_length as isize - request_length as isize; - - let inner_length = (zc_packet.buf_len() as isize + delta_length) as usize; - if zc_packet.mut_inner().capacity() < inner_length { - let header_length = inner_length - response_length; - zc_packet.mut_inner().truncate(header_length); - } - zc_packet.mut_inner().resize(inner_length, 0); - - let mut ip_packet = MutableIpv4Packet::new(zc_packet.mut_payload())?; - - let ip_length = (ip_packet.get_total_length() as isize + delta_length) as u16; - ip_packet.set_total_length(ip_length); - - let mut udp_packet = MutableUdpPacket::new(ip_packet.payload_mut())?; - - let udp_length = (udp_packet.get_length() as isize + delta_length) as u16; - udp_packet.set_length(udp_length); - - udp_packet.set_source(dst_port); - udp_packet.set_destination(src_port); - - udp_packet.payload_mut().copy_from_slice(&response_payload); - - udp_packet.set_checksum(udp::ipv4_checksum( - &udp_packet.to_immutable(), - &dst_ip, - &src_ip, - )); - - Some(()) - } - - fn handle_icmp_packet(&self, zc_packet: &mut ZCPacket, ip_header_length: usize) -> Option<()> { - let mut icmp_packet = - MutableIcmpPacket::new(&mut zc_packet.mut_payload()[ip_header_length..])?; - - if icmp_packet.get_icmp_type() != IcmpTypes::EchoRequest { - return None; - } - - icmp_packet.set_icmp_type(IcmpTypes::EchoReply); - icmp_packet.set_checksum(icmp::checksum(&icmp_packet.to_immutable())); - - Some(()) + Arc::into_inner(response)?.into_inner().ok() } } #[async_trait::async_trait] -impl NicPacketFilter for MagicDnsServerInstanceData { - async fn try_process_packet_from_nic(&self, zc_packet: &mut ZCPacket) -> bool { - self.handle_ip_packet(zc_packet).await.is_some() - } - - fn id(&self) -> String { - NIC_PIPELINE_NAME.to_string() +impl MagicDnsQueryResolver for MagicDnsServerInstanceData { + async fn resolve(&self, query: MagicDnsQuery) -> Option> { + self.resolve_query_inner(query).await } } @@ -464,18 +319,9 @@ impl RpcServerHook for MagicDnsServerInstanceData { let Some(remote_addr) = tunnel_info.remote_addr else { return; }; - let remote_addr = remote_addr.into(); - let mut removed_zones = vec![]; - for mut item in self.route_infos.iter_mut() { - item.value_mut().remove(&remote_addr); - if item.value().is_empty() { - removed_zones.push(item.key().clone()); - } - } - for zone in &removed_zones { - self.route_infos.remove(zone); - } - for zone in removed_zones { + let _apply = self.record_apply.lock().await; + let remote_addr: url::Url = remote_addr.into(); + for zone in self.route_store.remove_client(remote_addr.as_ref()) { self.keep_zone_authoritative(&zone).await; } self.update().await; @@ -483,9 +329,9 @@ impl RpcServerHook for MagicDnsServerInstanceData { } pub struct MagicDnsServerInstance { - rpc_server: StandAloneServer, + _rpc_server: StandAloneServer, pub(super) data: Arc, - peer_mgr: Arc, + packet_filter: MagicDnsResolverRegistration, tun_inet: Ipv4Inet, } @@ -510,13 +356,14 @@ fn get_system_config( } impl MagicDnsServerInstance { - pub async fn new( - peer_mgr: Arc, + pub(crate) async fn new( + packet_plane: Arc, + global_ctx: ArcGlobalCtx, tun_dev: Option, tun_inet: Ipv4Inet, fake_ip: Ipv4Addr, ) -> Result { - let tcp_listener = TcpTunnelListener::new(MAGIC_DNS_INSTANCE_ADDR.parse()?); + let tcp_listener = runtime_rpc_listener(MAGIC_DNS_INSTANCE_SOCKET_ADDR.parse()?); let mut rpc_server = StandAloneServer::new(tcp_listener); rpc_server.serve().await?; @@ -544,10 +391,9 @@ impl MagicDnsServerInstance { let data = Arc::new(MagicDnsServerInstanceData { dns_server, tun_dev: tun_dev.clone(), - tun_ip: tun_inet.address(), fake_ip, - my_peer_id: peer_mgr.my_peer_id(), - route_infos: DashMap::new(), + route_store: MagicDnsRecordStore::default(), + record_apply: tokio::sync::Mutex::new(()), system_config: get_system_config(tun_dev.as_deref())?, }); @@ -556,11 +402,8 @@ impl MagicDnsServerInstance { .register(MagicDnsServerRpcServer::new_arc(data.clone()), ""); rpc_server.set_hook(data.clone()); - peer_mgr - .add_nic_packet_process_pipeline(Box::new(data.clone())) - .await; // Use configured tld_dns_zone or fall back to DEFAULT_ET_DNS_ZONE if empty - let flags = peer_mgr.get_global_ctx().config.get_flags(); + let flags = global_ctx.config.get_flags(); let tld_dns_zone_clone = flags.tld_dns_zone.clone(); data.update_dns_records(std::iter::empty(), &tld_dns_zone_clone) @@ -572,10 +415,17 @@ impl MagicDnsServerInstance { .await .context("Failed to configure system")??; + // Install the resolver only after all fallible initialization has + // completed, so construction failure cannot leave a managed pipeline + // registration that never reaches async cleanup. + let packet_filter = packet_plane + .register_magic_dns_resolver(fake_ip, data.clone()) + .await; + Ok(Self { - rpc_server, + _rpc_server: rpc_server, data, - peer_mgr, + packet_filter, tun_inet, }) } @@ -596,10 +446,7 @@ impl MagicDnsServerInstance { } } - let _ = self - .peer_mgr - .remove_nic_packet_process_pipeline(NIC_PIPELINE_NAME.to_string()) - .await; + self.packet_filter.close().await; } } diff --git a/easytier/src/instance/dns_server/system_config/linux.rs b/easytier/src/instance/dns_server/system_config/linux.rs deleted file mode 100644 index 2b406195..00000000 --- a/easytier/src/instance/dns_server/system_config/linux.rs +++ /dev/null @@ -1,362 +0,0 @@ -// translated from tailscale #32ce1bdb48078ec4cedaeeb5b1b2ff9c0ef61a49 - -use anyhow::{Context, Result}; -use dbus::blocking::stdintf::org_freedesktop_dbus::Properties as _; -use std::fs; -use std::net::Ipv4Addr; -use std::path::Path; -use std::process::Command; -use std::time::Duration; -use version_compare::Cmp; - -// 声明依赖项(需要添加到Cargo.toml) -// use dbus::blocking::Connection; -// use nix::unistd::AccessFlags; -// use resolv_conf::Resolver; - -// 常量定义 -const RESOLV_CONF: &str = "/etc/resolv.conf"; -const PING_TIMEOUT: Duration = Duration::from_secs(1); - -// 错误类型定义 -#[derive(Debug)] -struct DNSConfigError { - message: String, - source: Option, -} - -type DbusPingFn = dyn Fn(&str, &str) -> Result<()>; -type DbusReadStringFn = dyn Fn(&str, &str, &str, &str) -> Result; -type NmIsUsingResolvedFn = dyn Fn() -> Result<()>; -type NmVersionBetweenFn = dyn Fn(&str, &str) -> Result; -type ResolvconfStyleFn = dyn Fn() -> String; - -// 配置环境结构体 -struct OSConfigEnv { - fs: Box, - dbus_ping: Box, - dbus_read_string: Box, - nm_is_using_resolved: Box, - nm_version_between: Box, - resolvconf_style: Box String>, -} - -// DNS管理器trait -trait OSConfigurator: Send + Sync { - // 实现相关方法 -} - -// 文件系统操作trait -trait FileSystem { - fn read_file(&self, path: &str) -> Result>; - fn exists(&self, path: &str) -> bool; -} - -// 直接文件系统实现 -struct DirectFS; - -impl FileSystem for DirectFS { - fn read_file(&self, path: &str) -> Result> { - fs::read(path).context("Failed to read file") - } - - fn exists(&self, path: &str) -> bool { - Path::new(path).exists() - } -} - -/// 检查 NetworkManager 是否使用 systemd-resolved 作为 DNS 管理器 -pub fn nm_is_using_resolved() -> Result<()> { - // 连接系统 D-Bus - let conn = dbus::blocking::Connection::new_system().context("Failed to connect to D-Bus")?; - - // 创建 NetworkManager DnsManager 对象代理 - let proxy = conn.with_proxy( - "org.freedesktop.NetworkManager", - "/org/freedesktop/NetworkManager/DnsManager", - std::time::Duration::from_secs(1), - ); - - // 获取 Mode 属性 - let (value,): (dbus::arg::Variant>,) = proxy - .method_call( - "org.freedesktop.DBus.Properties", - "Get", - ("org.freedesktop.NetworkManager.DnsManager", "Mode"), - ) - .context("Failed to get NM mode property")?; - - // 检查 Mode 是否为 "systemd-resolved" - if value.0.as_str() != Some("systemd-resolved") { - return Err(anyhow::anyhow!( - "NetworkManager is not using systemd-resolved, found: {:?}", - value - )); - } - - Ok(()) -} - -/// 返回系统中使用的 resolvconf 实现类型("debian" 或 "openresolv") -pub fn resolvconf_style() -> String { - // 检查 resolvconf 命令是否存在 - if which::which("resolvconf").is_err() { - return String::new(); - } - - // 执行 resolvconf --version 命令 - let output = match Command::new("resolvconf").arg("--version").output() { - Ok(output) => output, - Err(e) => { - // 处理命令执行错误 - if let Some(code) = e.raw_os_error() { - // Debian 版本的 resolvconf 不支持 --version,返回特定错误码 99 - if code == 99 { - return "debian".to_string(); - } - } - return String::new(); // 其他错误返回空字符串 - } - }; - - // 检查输出是否以 "Debian resolvconf" 开头 - if output.stdout.starts_with(b"Debian resolvconf") { - return "debian".to_string(); - } - - // 默认视为 openresolv - "openresolv".to_string() -} - -// 构建配置环境 -fn new_os_config_env() -> OSConfigEnv { - OSConfigEnv { - fs: Box::new(DirectFS), - dbus_ping: Box::new(dbus_ping), - dbus_read_string: Box::new(dbus_read_string), - nm_is_using_resolved: Box::new(nm_is_using_resolved), - nm_version_between: Box::new(nm_version_between), - resolvconf_style: Box::new(resolvconf_style), - } -} - -// 创建DNS配置器 -fn new_os_configurator(_interface_name: String) -> Result<()> { - let env = new_os_config_env(); - - let mode = dns_mode(&env).context("Failed to detect DNS mode")?; - - tracing::info!("dns: using {} mode", mode); - - // match mode.as_str() { - // "direct" => Ok(Box::new(DirectManager::new(env.fs)?)), - // // "systemd-resolved" => Ok(Box::new(ResolvedManager::new( - // // &logf, - // // health, - // // interface_name, - // // )?)), - // // "network-manager" => Ok(Box::new(NMManager::new(interface_name)?)), - // // "debian-resolvconf" => Ok(Box::new(DebianResolvconfManager::new(&logf)?)), - // // "openresolv" => Ok(Box::new(OpenresolvManager::new(&logf)?)), - // _ => { - // tracing::warn!("Unexpected DNS mode {}, using direct manager", mode); - // Ok(Box::new(DirectManager::new(env.fs)?)) - // } - // } - Ok(()) -} - -use guarden::defer; -use std::io::{self, BufRead, Cursor}; - -/// 返回 `resolv.conf` 内容的拥有者("systemd-resolved"、"NetworkManager"、"resolvconf" 或空字符串) -pub fn resolv_owner(bs: &[u8]) -> String { - let mut likely = String::new(); - let cursor = Cursor::new(bs); - let reader = io::BufReader::new(cursor); - - for line_result in reader.lines() { - match line_result { - Ok(line) => { - let line = line.trim(); - if line.is_empty() { - continue; - } - - if !line.starts_with('#') { - // 第一个非注释且非空的行,直接返回当前结果 - return likely; - } - - // 检查注释行中的关键字 - if line.contains("systemd-resolved") { - likely = "systemd-resolved".to_string(); - } else if line.contains("NetworkManager") { - likely = "NetworkManager".to_string(); - } else if line.contains("resolvconf") { - likely = "resolvconf".to_string(); - } - } - Err(_) => { - // 读取错误(如无效 UTF-8),直接返回当前结果 - return likely; - } - } - } - - likely -} - -// 检测DNS模式 -fn dns_mode(env: &OSConfigEnv) -> Result { - let debug = std::cell::RefCell::new(Vec::new()); - let dbg = |k: &str, v: &str| debug.borrow_mut().push((k.to_string(), v.to_string())); - - // defer 日志记录 - defer! { - if !debug.borrow().is_empty() { - let log_entries: Vec = - debug.borrow().iter().map(|(k, v)| format!("{}={}", k, v)).collect(); - tracing::info!("dns: [{}]", log_entries.join(" ")); - } - }; - - // 检查systemd-resolved状态 - let resolved_up = - (env.dbus_ping)("org.freedesktop.resolve1", "/org/freedesktop/resolve1").is_ok(); - if resolved_up { - dbg("resolved-ping", "yes"); - } - - // 读取resolv.conf - let content = match env.fs.read_file(RESOLV_CONF) { - Ok(content) => content, - Err(e) if e.to_string().contains("NotFound") => { - dbg("rc", "missing"); - return Ok("direct".to_string()); - } - Err(e) => return Err(e).context("reading /etc/resolv.conf"), - }; - - // 检查resolv.conf所有者 - match resolv_owner(&content).as_str() { - "systemd-resolved" => { - dbg("rc", "resolved"); - // 检查是否实际使用resolved - if let Err(e) = resolved_is_actually_resolver(env, &dbg, &content) { - tracing::warn!("resolvedIsActuallyResolver error: {}", e); - dbg("resolved", "not-in-use"); - return Ok("direct".to_string()); - } - - // NetworkManager检查逻辑... - - Ok("systemd-resolved".to_string()) - } - "resolvconf" => { - // resolvconf处理逻辑... - Ok("debian-resolvconf".to_string()) - } - "NetworkManager" => { - // NetworkManager处理逻辑... - Ok("systemd-resolved".to_string()) - } - _ => Ok("direct".to_string()), - } -} - -// D-Bus ping实现 -fn dbus_ping(name: &str, object_path: &str) -> Result<()> { - let conn = dbus::blocking::Connection::new_system()?; - let proxy = conn.with_proxy(name, object_path, PING_TIMEOUT); - let _: () = proxy.method_call("org.freedesktop.DBus.Peer", "Ping", ())?; - Ok(()) -} - -// D-Bus读取字符串实现 -fn dbus_read_string(name: &str, object_path: &str, iface: &str, member: &str) -> Result { - let conn = dbus::blocking::Connection::new_system()?; - let proxy = conn.with_proxy(name, object_path, PING_TIMEOUT); - let (value,): (String,) = - proxy.method_call("org.freedesktop.DBus.Properties", "Get", (iface, member))?; - Ok(value) -} - -// NetworkManager版本检查 -fn nm_version_between(first: &str, last: &str) -> Result { - let conn = dbus::blocking::Connection::new_system()?; - let proxy = conn.with_proxy( - "org.freedesktop.NetworkManager", - "/org/freedesktop/NetworkManager", - PING_TIMEOUT, - ); - - let version: String = proxy.get("org.freedesktop.NetworkManager", "Version")?; - let cmp_first = version_compare::compare(&version, first).unwrap_or(Cmp::Lt); - let cmp_last = version_compare::compare(&version, last).unwrap_or(Cmp::Gt); - Ok(cmp_first == Cmp::Ge && cmp_last == Cmp::Le) -} - -// 检查是否实际使用systemd-resolved -fn resolved_is_actually_resolver( - env: &OSConfigEnv, - dbg: &dyn Fn(&str, &str), - content: &[u8], -) -> Result<()> { - if is_libnss_resolve_used(env).is_ok() { - dbg("resolved", "nss"); - return Ok(()); - } - - // 解析resolv.conf内容 - let resolver = resolv_conf::Config::parse(content)?; - - // 检查nameserver配置 - if resolver.nameservers.is_empty() { - return Err(anyhow::anyhow!("resolv.conf has no nameservers")); - } - - for ns in resolver.nameservers { - if ns != Ipv4Addr::new(127, 0, 0, 53).into() { - return Err(anyhow::anyhow!( - "resolv.conf doesn't point to systemd-resolved" - )); - } - } - - dbg("resolved", "file"); - Ok(()) -} - -// 检查是否使用libnss_resolve -fn is_libnss_resolve_used(env: &OSConfigEnv) -> Result<()> { - let content = env.fs.read_file("/etc/nsswitch.conf")?; - - for line in String::from_utf8_lossy(&content).lines() { - let parts: Vec<&str> = line.split_whitespace().collect(); - if parts.first() == Some(&"hosts:") { - for module in parts.iter().skip(1) { - if *module == "dns" { - return Err(anyhow::anyhow!("dns module has higher priority")); - } - if *module == "resolve" { - return Ok(()); - } - } - } - } - - Err(anyhow::anyhow!("libnss_resolve not used")) -} - -#[cfg(test)] -mod tests { - use super::*; - - #[test] - fn dns_mode_test() { - let env = new_os_config_env(); - let mode = dns_mode(&env).unwrap(); - println!("Detected DNS mode: {}", mode); - } -} diff --git a/easytier/src/instance/dns_server/system_config/mod.rs b/easytier/src/instance/dns_server/system_config/mod.rs index 388ea3e8..6ec0d361 100644 --- a/easytier/src/instance/dns_server/system_config/mod.rs +++ b/easytier/src/instance/dns_server/system_config/mod.rs @@ -1,6 +1,3 @@ -#[cfg(target_os = "linux")] -pub mod linux; - #[cfg(target_os = "windows")] pub mod windows; diff --git a/easytier/src/instance/dns_server/system_config/windows.rs b/easytier/src/instance/dns_server/system_config/windows.rs index 3c298cf9..030b60ec 100644 --- a/easytier/src/instance/dns_server/system_config/windows.rs +++ b/easytier/src/instance/dns_server/system_config/windows.rs @@ -126,7 +126,6 @@ impl InterfaceControl { } pub struct WindowsDNSManager { - tun_dev_name: String, interface_control: InterfaceControl, } @@ -134,7 +133,6 @@ impl WindowsDNSManager { pub fn new(tun_dev_name: &str) -> io::Result { let interface_guid = RegistryManager::find_interface_guid(tun_dev_name)?; Ok(WindowsDNSManager { - tun_dev_name: tun_dev_name.to_string(), interface_control: InterfaceControl::new(&interface_guid), }) } @@ -180,12 +178,18 @@ mod tests { }; let tun_ip = Ipv4Inet::from_str("10.144.144.10/24").unwrap(); - let (peer_mgr, virtual_nic) = prepare_env("test1", tun_ip).await; + let (global_ctx, core_instance, virtual_nic) = prepare_env("test1", tun_ip).await; let tun_name = virtual_nic.ifname().await.unwrap(); println!("dev_name: {}", tun_name); let fake_ip = Ipv4Addr::from_str("100.100.100.101").unwrap(); - let mut dns_runner = DnsRunner::new(peer_mgr, Some(tun_name.clone()), tun_ip, fake_ip); + let mut dns_runner = DnsRunner::new( + core_instance.packet_plane(), + global_ctx, + Some(tun_name.clone()), + tun_ip, + fake_ip, + ); let cancel_token = CancellationToken::new(); let cancel_token_clone = cancel_token.clone(); diff --git a/easytier/src/instance/dns_server/tests.rs b/easytier/src/instance/dns_server/tests.rs index 2e9bb2e7..de097601 100644 --- a/easytier/src/instance/dns_server/tests.rs +++ b/easytier/src/instance/dns_server/tests.rs @@ -4,6 +4,7 @@ use std::sync::Arc; use std::time::Duration; use cidr::Ipv4Inet; +use easytier_core::{gateway::magic_dns::MagicDnsRoute, process_runtime::CoreProcessRuntime}; use hickory_client::client::{Client, ClientHandle as _}; use hickory_proto::rr; use hickory_proto::runtime::TokioRuntimeProvider; @@ -11,30 +12,49 @@ use hickory_proto::udp::UdpClientStream; use tokio::sync::Notify; use tokio_util::sync::CancellationToken; -use crate::common::global_ctx::tests::get_mock_global_ctx; -use crate::connector::udp_hole_punch::tests::replace_stun_info_collector; +use crate::common::global_ctx::{ArcGlobalCtx, tests::get_mock_global_ctx}; +use crate::instance::{ + composition::{NativeCoreInstance, runtime_core_host_adapters}, + config::test_core_instance_config, +}; use crate::instance::dns_server::runner::DnsRunner; use crate::instance::dns_server::server_instance::MagicDnsServerInstance; use crate::instance::dns_server::{DEFAULT_ET_DNS_ZONE, MAGIC_DNS_FAKE_IP}; use crate::instance::virtual_nic::NicCtx; -use crate::peers::peer_manager::{PeerManager, RouteAlgoType}; - -use crate::peers::create_packet_recv_chan; use crate::proto::api::instance::Route; -use crate::proto::common::NatType; use crate::proto::magic_dns::{MagicDnsServerRpc as _, UpdateDnsRecordRequest}; use crate::proto::rpc_types::controller::{BaseController, Controller as _}; -pub async fn prepare_env(dns_name: &str, tun_ip: Ipv4Inet) -> (Arc, NicCtx) { +pub async fn prepare_env( + dns_name: &str, + tun_ip: Ipv4Inet, +) -> (ArcGlobalCtx, Arc, NicCtx) { prepare_env_with_tld_dns_zone(dns_name, tun_ip, None).await } +async fn build_test_core( + ctx: ArcGlobalCtx, +) -> ( + Arc, + tokio::sync::mpsc::Receiver>, +) { + let (packet_sink, packet_receiver) = tokio::sync::mpsc::channel(128); + let adapters = runtime_core_host_adapters( + ctx.clone(), + CoreProcessRuntime::new(), + Arc::new(packet_sink), + ); + let core_instance = NativeCoreInstance::new(test_core_instance_config(&ctx), adapters).unwrap(); + core_instance.start().await.unwrap(); + (core_instance, packet_receiver) +} + pub async fn prepare_env_with_tld_dns_zone( dns_name: &str, tun_ip: Ipv4Inet, tld_dns_zone: Option<&str>, -) -> (Arc, NicCtx) { +) -> (ArcGlobalCtx, Arc, NicCtx) { let ctx = get_mock_global_ctx(); ctx.set_hostname(dns_name.to_owned()); ctx.set_ipv4(Some(tun_ip)); @@ -48,21 +68,17 @@ pub async fn prepare_env_with_tld_dns_zone( ctx.set_flags(flags); } - let (s, r) = create_packet_recv_chan(); - let peer_mgr = Arc::new(PeerManager::new(RouteAlgoType::Ospf, ctx, s)); - peer_mgr.run().await.unwrap(); - replace_stun_info_collector(peer_mgr.clone(), NatType::PortRestricted); - - let r = Arc::new(tokio::sync::Mutex::new(r)); + let (core_instance, host_packet_rx) = build_test_core(ctx.clone()).await; + let host_packet_rx = Arc::new(tokio::sync::Mutex::new(host_packet_rx)); let mut virtual_nic = NicCtx::new( - peer_mgr.get_global_ctx(), - &peer_mgr, - r, + ctx.clone(), + core_instance.packet_plane(), + host_packet_rx, Arc::new(Notify::new()), ); virtual_nic.run(Some(tun_ip), None).await.unwrap(); - (peer_mgr, virtual_nic) + (ctx, core_instance, virtual_nic) } pub async fn check_dns_record(fake_ip: &Ipv4Addr, domain: &str, expected_ip: &str) { @@ -94,55 +110,34 @@ pub async fn check_dns_record(fake_ip: &Ipv4Addr, domain: &str, expected_ip: &st ); } -pub async fn check_dns_record_missing(fake_ip: &Ipv4Addr, domain: &str) { - let stream = UdpClientStream::builder( - SocketAddr::new((*fake_ip).into(), 53), - TokioRuntimeProvider::default(), - ) - .build(); - let (mut client, background) = Client::connect(stream).await.unwrap(); - let background_task = tokio::spawn(background); - let response = client - .query( - rr::Name::from_str(domain).unwrap(), - rr::DNSClass::IN, - rr::RecordType::A, - ) - .await - .unwrap_or_else(|e| { - panic!("DNS query for missing record failed unexpectedly for domain '{domain}': {e}") - }); - background_task.abort(); - let _ = background_task.await; - assert!(response.answers().is_empty(), "{:?}", response.answers()); -} - #[tokio::test] async fn test_magic_dns_server_instance() { let tun_ip = Ipv4Inet::from_str("10.144.144.10/24").unwrap(); - let (peer_mgr, virtual_nic) = prepare_env("test1", tun_ip).await; + let (global_ctx, core_instance, virtual_nic) = prepare_env("test1", tun_ip).await; let tun_name = virtual_nic.ifname().await.unwrap(); let fake_ip = Ipv4Addr::from_str("100.100.100.101").unwrap(); - let dns_server_inst = - MagicDnsServerInstance::new(peer_mgr.clone(), Some(tun_name), tun_ip, fake_ip) - .await - .unwrap(); + let dns_server_inst = MagicDnsServerInstance::new( + core_instance.packet_plane(), + global_ctx, + Some(tun_name), + tun_ip, + fake_ip, + ) + .await + .unwrap(); let routes = [ - Route { + MagicDnsRoute { hostname: "test1".to_string(), - ipv4_addr: Some(Ipv4Inet::from_str("8.8.8.8/24").unwrap().into()), - ..Default::default() + ipv4_addr: Some("8.8.8.8".parse().unwrap()), }, - Route { + MagicDnsRoute { hostname: "中文".to_string(), - ipv4_addr: Some(Ipv4Inet::from_str("8.8.8.8/24").unwrap().into()), - ..Default::default() + ipv4_addr: Some("8.8.8.8".parse().unwrap()), }, - Route { + MagicDnsRoute { hostname: ".invalid".to_string(), - ipv4_addr: Some(Ipv4Inet::from_str("8.8.8.8/24").unwrap().into()), - ..Default::default() + ipv4_addr: Some("8.8.8.8".parse().unwrap()), }, ]; dns_server_inst @@ -160,10 +155,16 @@ async fn test_magic_dns_runner() { // Test first runner with default DNS settings { let tun_ip = Ipv4Inet::from_str("10.144.144.10/24").unwrap(); - let (peer_mgr, virtual_nic) = prepare_env("test1", tun_ip).await; + let (global_ctx, core_instance, virtual_nic) = prepare_env("test1", tun_ip).await; let tun_name = virtual_nic.ifname().await.unwrap(); let fake_ip = Ipv4Addr::from_str(MAGIC_DNS_FAKE_IP).unwrap(); - let mut dns_runner = DnsRunner::new(peer_mgr, Some(tun_name), tun_ip, fake_ip); + let mut dns_runner = DnsRunner::new( + core_instance.packet_plane(), + global_ctx, + Some(tun_name), + tun_ip, + fake_ip, + ); let cancel_token = CancellationToken::new(); let cancel_token_clone = cancel_token.clone(); @@ -187,11 +188,17 @@ async fn test_magic_dns_runner() { let tun_ip = Ipv4Inet::from_str("10.144.144.20/24").unwrap(); // NOTE: Using same fake IP to avoid system DNS configuration conflicts let custom_tld_zone = "custom.local."; // Different TLD zone is safer - let (peer_mgr, virtual_nic) = + let (global_ctx, core_instance, virtual_nic) = prepare_env_with_tld_dns_zone("test2", tun_ip, Some(custom_tld_zone)).await; let tun_name = virtual_nic.ifname().await.unwrap(); let fake_ip = Ipv4Addr::from_str(MAGIC_DNS_FAKE_IP).unwrap(); - let mut dns_runner = DnsRunner::new(peer_mgr, Some(tun_name), tun_ip, fake_ip); + let mut dns_runner = DnsRunner::new( + core_instance.packet_plane(), + global_ctx, + Some(tun_name), + tun_ip, + fake_ip, + ); let cancel_token = CancellationToken::new(); let cancel_token_clone = cancel_token.clone(); @@ -215,15 +222,13 @@ async fn test_magic_dns_update_replaces_records_for_same_client() { ctx.set_hostname("test1".to_string()); ctx.set_ipv4(Some(tun_ip)); - let (s, _r) = create_packet_recv_chan(); - let peer_mgr = Arc::new(PeerManager::new(RouteAlgoType::Ospf, ctx, s)); - peer_mgr.run().await.unwrap(); - replace_stun_info_collector(peer_mgr.clone(), NatType::PortRestricted); + let (core_instance, _packet_receiver) = build_test_core(ctx.clone()).await; let fake_ip = Ipv4Addr::from_str(MAGIC_DNS_FAKE_IP).unwrap(); - let dns_server_inst = MagicDnsServerInstance::new(peer_mgr.clone(), None, tun_ip, fake_ip) - .await - .unwrap(); + let dns_server_inst = + MagicDnsServerInstance::new(core_instance.packet_plane(), ctx, None, tun_ip, fake_ip) + .await + .unwrap(); let mut ctrl = BaseController::default(); ctrl.set_tunnel_info(Some(crate::proto::common::TunnelInfo { diff --git a/easytier/src/instance/factory.rs b/easytier/src/instance/factory.rs new file mode 100644 index 00000000..31e97e48 --- /dev/null +++ b/easytier/src/instance/factory.rs @@ -0,0 +1,200 @@ +use std::sync::Arc; + +#[cfg(any(feature = "management-rpc", test))] +use easytier_core::instance::manager::InstanceManager; +#[cfg(feature = "management-rpc")] +use easytier_core::management::ProcessRuntimeProvider; +use easytier_core::{ + config::toml::TomlConfig, + instance::{CoreInstance, manager::InstanceFactory}, + process_runtime::CoreProcessRuntime, +}; + +use crate::common::global_ctx::EventBusSubscriber; + +use super::{ + composition::compose_native_core_instance, host::NativeInstanceHost, + runtime_host::NativeInstanceRuntimeHost, +}; + +pub type NativeCoreInstance = CoreInstance; +#[cfg(feature = "management-rpc")] +pub type NativeInstanceManager = InstanceManager; +#[cfg(feature = "management")] +pub type NativeProcessManagement = + easytier_core::management::ProcessManagement; + +#[cfg(feature = "management-rpc")] +pub fn native_instance_manager() -> NativeInstanceManager { + native_instance_manager_with_optional_runtime(None) +} + +#[cfg(feature = "management")] +pub fn native_cli_instance_manager() -> NativeInstanceManager { + let process_runtime = CoreProcessRuntime::new(); + InstanceManager::new( + NativeInstanceFactory::new(process_runtime).with_cli_event_logging(), + None, + ) +} + +pub fn create_native_instance(config: TomlConfig) -> anyhow::Result> { + NativeInstanceFactory::new(CoreProcessRuntime::new()).create(config, ()) +} + +/// Subscribes to native presentation events owned by this instance's runtime. +pub fn subscribe_native_instance_event( + instance: &NativeCoreInstance, +) -> Option { + instance + .runtime_host::() + .map(NativeInstanceRuntimeHost::subscribe_event) +} + +#[cfg(feature = "management-rpc")] +pub fn native_instance_manager_with_runtime( + runtime_handle: tokio::runtime::Handle, +) -> NativeInstanceManager { + native_instance_manager_with_optional_runtime(Some(runtime_handle)) +} + +#[cfg(feature = "management")] +pub fn native_process_management( + instances: Arc, + hooks: Arc, +) -> NativeProcessManagement { + NativeProcessManagement::new( + instances, + hooks, + Arc::new(easytier_core::management::UnsupportedConfigFileStorage), + ) +} + +#[cfg(feature = "management-rpc")] +fn native_instance_manager_with_optional_runtime( + runtime_handle: Option, +) -> NativeInstanceManager { + let process_runtime = CoreProcessRuntime::new(); + InstanceManager::new( + NativeInstanceFactory::new(process_runtime).with_runtime_handle(runtime_handle.clone()), + runtime_handle, + ) +} + +/// Native construction Adapter for the canonical core InstanceManager. +pub struct NativeInstanceFactory { + process_runtime: Arc, + runtime_handle: Option, + #[cfg(feature = "management")] + log_cli_events: bool, +} + +impl NativeInstanceFactory { + pub fn new(process_runtime: Arc) -> Self { + Self { + process_runtime, + runtime_handle: None, + #[cfg(feature = "management")] + log_cli_events: false, + } + } + + #[cfg(feature = "management")] + fn with_cli_event_logging(mut self) -> Self { + self.log_cli_events = true; + self + } + + #[cfg(feature = "management-rpc")] + fn with_runtime_handle(mut self, runtime_handle: Option) -> Self { + self.runtime_handle = runtime_handle; + self + } +} + +impl InstanceFactory for NativeInstanceFactory { + type Instance = NativeCoreInstance; + type CreateContext = (); + type Error = anyhow::Error; + + fn create( + &self, + config: TomlConfig, + (): Self::CreateContext, + ) -> Result, Self::Error> { + let _runtime = self + .runtime_handle + .as_ref() + .map(tokio::runtime::Handle::enter); + let instance = compose_native_core_instance(config, self.process_runtime.clone())?; + #[cfg(feature = "management")] + if self.log_cli_events { + let events = subscribe_native_instance_event(&instance) + .ok_or_else(|| anyhow::anyhow!("native instance runtime host is unavailable"))?; + super::cli_event_logger::spawn(instance.instance_id(), events); + } + Ok(instance) + } +} + +#[cfg(feature = "management-rpc")] +impl ProcessRuntimeProvider for NativeInstanceFactory { + fn process_runtime(&self) -> Arc { + self.process_runtime.clone() + } +} + +#[cfg(test)] +mod tests { + use easytier_core::{config::toml::ConfigLoader as _, instance::CoreInstanceState}; + + use super::*; + + #[tokio::test] + async fn core_manager_stores_and_runs_native_core_instance_directly() { + let factory = NativeInstanceFactory::new(CoreProcessRuntime::new()); + let manager = InstanceManager::new(factory, None); + let config = TomlConfig::default(); + let mut flags = config.get_flags(); + flags.no_tun = true; + config.set_flags(flags); + config.set_listeners(Vec::new()); + + let instance = manager.create(config, ()).unwrap(); + instance.start().await.unwrap(); + assert_eq!(instance.state(), CoreInstanceState::Running); + + manager + .delete_network_instances([instance.instance_id()]) + .await + .unwrap(); + assert_eq!(instance.state(), CoreInstanceState::Stopped); + } + + #[test] + fn configured_runtime_supports_synchronous_instance_construction() { + let runtime = tokio::runtime::Runtime::new().unwrap(); + let factory = NativeInstanceFactory::new(CoreProcessRuntime::new()) + .with_runtime_handle(Some(runtime.handle().clone())); + let manager = InstanceManager::new(factory, Some(runtime.handle().clone())); + let config = TomlConfig::default(); + config.set_listeners(Vec::new()); + + let instance = manager.create(config, ()).unwrap(); + + drop(instance); + } + + #[test] + fn event_subscription_is_recovered_from_the_native_runtime() { + let runtime = tokio::runtime::Runtime::new().unwrap(); + let factory = NativeInstanceFactory::new(CoreProcessRuntime::new()) + .with_runtime_handle(Some(runtime.handle().clone())); + let manager = InstanceManager::new(factory, Some(runtime.handle().clone())); + let config = TomlConfig::default(); + config.set_listeners(Vec::new()); + + let instance = manager.create(config, ()).unwrap(); + assert!(subscribe_native_instance_event(&instance).is_some()); + } +} diff --git a/easytier/src/instance/host.rs b/easytier/src/instance/host.rs new file mode 100644 index 00000000..f07ce094 --- /dev/null +++ b/easytier/src/instance/host.rs @@ -0,0 +1,59 @@ +use std::{net::IpAddr, sync::Arc}; + +use easytier_core::{ + connectivity::composite::{ConnectorEnvironment, ConnectorHostAdapter}, + socket::{NetNamespace, SocketContext}, +}; + +use crate::{ + common::global_ctx::ArcGlobalCtx, + host_runtime::{NativeHostRuntime, native_host_runtime}, +}; + +pub type NativeInstanceHost = ConnectorHostAdapter; + +/// Instance facts queried by portable connector policy. +/// +/// This Adapter never creates or operates sockets. Mechanical network I/O is +/// owned by the process-wide [`NativeHostRuntime`] composed beside it. +pub struct NativeInstanceEnvironment { + global_ctx: ArcGlobalCtx, + runtime: Arc, + socket_context: SocketContext, +} + +impl NativeInstanceEnvironment { + fn new(global_ctx: ArcGlobalCtx, runtime: Arc) -> Self { + let socket_context = SocketContext::default() + .with_socket_mark(global_ctx.config.get_flags().socket_mark) + .with_netns(global_ctx.net_ns.name().map(NetNamespace::new)); + Self { + global_ctx, + runtime, + socket_context, + } + } +} + +pub(crate) fn native_instance_host(global_ctx: ArcGlobalCtx) -> Arc { + let runtime = native_host_runtime(); + Arc::new(ConnectorHostAdapter::new( + runtime.clone(), + Arc::new(NativeInstanceEnvironment::new(global_ctx, runtime)), + )) +} + +impl ConnectorEnvironment for NativeInstanceEnvironment { + fn socket_context(&self) -> SocketContext { + self.socket_context.clone() + } + + fn mapped_listeners(&self) -> Vec { + self.global_ctx.config.get_mapped_listeners() + } + + fn is_local_ip(&self, ip: &IpAddr) -> bool { + self.global_ctx.is_ip_local_virtual_ip(ip) + || self.runtime.is_local_ip(ip, &self.socket_context) + } +} diff --git a/easytier/src/instance/instance.rs b/easytier/src/instance/instance.rs deleted file mode 100644 index 6a8ee3a6..00000000 --- a/easytier/src/instance/instance.rs +++ /dev/null @@ -1,1880 +0,0 @@ -#[cfg(feature = "tun")] -use std::any::Any; -use std::collections::HashSet; -use std::net::{IpAddr, Ipv4Addr}; -use std::sync::atomic::{AtomicBool, Ordering}; -use std::sync::{Arc, Weak}; -#[cfg(feature = "tun")] -use std::time::Duration; - -use anyhow::Context; -use cidr::{IpCidr, Ipv4Inet}; -use futures::FutureExt; -use tokio::sync::{Mutex, Notify}; -#[cfg(feature = "tun")] -use tokio::{sync::oneshot, task::JoinSet}; -#[cfg(feature = "magic-dns")] -use tokio_util::sync::CancellationToken; -use tokio_util::task::AbortOnDropHandle; - -use crate::common::PeerId; -use crate::common::acl_processor::AclRuleBuilder; -use crate::common::config::ConfigLoader; -use crate::common::error::Error; -use crate::common::global_ctx::{ArcGlobalCtx, GlobalCtx, GlobalCtxEvent}; -use crate::connector::direct::DirectConnectorManager; -use crate::connector::manual::{ConnectorManagerRpcService, ManualConnectorManager}; -use crate::connector::tcp_hole_punch::TcpHolePunchConnector; -use crate::connector::udp_hole_punch::UdpHolePunchConnector; -use crate::gateway::icmp_proxy::IcmpProxy; -#[cfg(feature = "kcp")] -use crate::gateway::kcp_proxy::{KcpProxyDst, KcpProxyDstRpcService, KcpProxySrc}; -#[cfg(feature = "quic")] -use crate::gateway::quic_proxy::{QuicProxy, QuicProxyDstRpcService}; -use crate::gateway::tcp_proxy::{NatDstTcpConnector, TcpProxy, TcpProxyRpcService}; -use crate::gateway::udp_proxy::UdpProxy; -use crate::peer_center::instance::{PeerCenterInstance, PeerCenterInstanceService}; -use crate::peers::peer_conn::PeerConnId; -use crate::peers::peer_manager::{PeerManager, RouteAlgoType}; -#[cfg(feature = "tun")] -use crate::peers::recv_packet_from_chan; -use crate::peers::rpc_service::PeerManagerRpcService; -use crate::peers::{PacketRecvChanReceiver, create_packet_recv_chan}; -use crate::proto::api::config::{ - ConfigPatchAction, ConfigRpc, GetConfigRequest, GetConfigResponse, PatchConfigRequest, - PatchConfigResponse, PortForwardPatch, -}; -use crate::proto::api::instance::{ - GetPrometheusStatsRequest, GetPrometheusStatsResponse, GetStatsRequest, GetStatsResponse, - GetVpnPortalInfoRequest, GetVpnPortalInfoResponse, ListMappedListenerRequest, - ListMappedListenerResponse, ListPortForwardRequest, ListPortForwardResponse, MappedListener, - MappedListenerManageRpc, MetricSnapshot, PortForwardManageRpc, StatsRpc, VpnPortalInfo, - VpnPortalRpc, -}; -use crate::proto::api::manage::NetworkConfig; -use crate::proto::common::{PortForwardConfigPb, TunnelInfo}; -use crate::proto::peer_rpc::PeerCenterRpc; -use crate::proto::rpc_impl::standalone::RpcServerHook; -use crate::proto::rpc_types; -use crate::proto::rpc_types::controller::BaseController; -use crate::rpc_service::InstanceRpcService; -use crate::utils::weak_upgrade; -use crate::vpn_portal::{self, VpnPortal}; - -#[cfg(feature = "magic-dns")] -use super::dns_server::{MAGIC_DNS_FAKE_IP, runner::DnsRunner}; -use super::listeners::ListenerManager; -use super::public_ipv6_provider::{ - PublicIpv6ProviderReconcileTask, reconcile_public_ipv6_provider_runtime, - run_public_ipv6_provider_reconcile_task, should_run_public_ipv6_provider_reconcile, - validate_public_ipv6_config, validate_public_ipv6_config_values, -}; - -#[cfg(feature = "socks5")] -use crate::gateway::socks5::Socks5Server; - -#[derive(Clone)] -struct IpProxy { - tcp_proxy: Arc>, - icmp_proxy: Arc, - udp_proxy: Arc, - global_ctx: ArcGlobalCtx, - started: Arc, -} - -impl IpProxy { - fn new(global_ctx: ArcGlobalCtx, peer_manager: Arc) -> Result { - let tcp_proxy = TcpProxy::new(peer_manager.clone(), NatDstTcpConnector {}); - let icmp_proxy = IcmpProxy::new(global_ctx.clone(), peer_manager.clone()) - .with_context(|| "create icmp proxy failed")?; - let udp_proxy = UdpProxy::new(global_ctx.clone(), peer_manager) - .with_context(|| "create udp proxy failed")?; - Ok(IpProxy { - tcp_proxy, - icmp_proxy, - udp_proxy, - global_ctx, - started: Arc::new(AtomicBool::new(false)), - }) - } - - async fn start(&self) -> Result<(), Error> { - if (self.global_ctx.config.get_proxy_cidrs().is_empty() - || self.started.load(Ordering::Relaxed)) - && !self.global_ctx.enable_exit_node() - && !self.global_ctx.no_tun() - { - return Ok(()); - } - - // Actually, if this node is enabled as an exit node, - // we still can use the system stack to forward packets. - if self.global_ctx.proxy_forward_by_system() && !self.global_ctx.no_tun() { - return Ok(()); - } - - self.started.store(true, Ordering::Relaxed); - self.tcp_proxy.start(true).await?; - if let Err(e) = self.icmp_proxy.start().await { - tracing::error!("start icmp proxy failed: {:?}", e); - if cfg!(not(any( - target_os = "android", - any( - target_os = "ios", - all(target_os = "macos", feature = "macos-ne") - ), - target_env = "ohos" - ))) { - // android, ios and ohos not support icmp proxy - return Err(e); - } - } - self.udp_proxy.start().await?; - Ok(()) - } -} - -#[cfg(feature = "tun")] -type NicCtx = super::virtual_nic::NicCtx; - -#[cfg(feature = "magic-dns")] -struct MagicDnsContainer { - dns_runner_task: AbortOnDropHandle<()>, - dns_runner_cancel_token: CancellationToken, -} - -// nic container will be cleared when dhcp ip changed -#[cfg(feature = "tun")] -pub struct NicCtxContainer { - nic_ctx: Option>, - #[cfg(feature = "magic-dns")] - magic_dns: Option, -} - -#[cfg(feature = "tun")] -impl NicCtxContainer { - #[cfg(not(feature = "magic-dns"))] - fn new(nic_ctx: NicCtx) -> Self { - Self { - nic_ctx: Some(Box::new(nic_ctx)), - } - } - - #[cfg(feature = "magic-dns")] - fn new(nic_ctx: NicCtx, dns_runner: Option) -> Self { - if let Some(mut dns_runner) = dns_runner { - let token = CancellationToken::new(); - let token_clone = token.clone(); - let task = tokio::spawn(async move { - let _ = dns_runner.run(token_clone).await; - }); - Self { - nic_ctx: Some(Box::new(nic_ctx)), - magic_dns: Some(MagicDnsContainer { - dns_runner_task: AbortOnDropHandle::new(task), - dns_runner_cancel_token: token, - }), - } - } else { - Self { - nic_ctx: Some(Box::new(nic_ctx)), - magic_dns: None, - } - } - } - - fn new_with_any(ctx: T) -> Self { - Self { - nic_ctx: Some(Box::new(ctx)), - #[cfg(feature = "magic-dns")] - magic_dns: None, - } - } -} - -#[cfg(feature = "tun")] -type ArcNicCtx = Arc>>; -type ArcPublicIpv6ProviderTaskSlot = Arc; - -struct PublicIpv6ProviderTaskSlot { - task: Mutex>, - closing: AtomicBool, -} - -impl PublicIpv6ProviderTaskSlot { - fn new() -> Self { - Self { - task: Mutex::new(None), - closing: AtomicBool::new(false), - } - } - - async fn ensure_started(&self, global_ctx: &ArcGlobalCtx) { - let mut task = self.task.lock().await; - if self.closing.load(Ordering::Acquire) || task.is_some() { - return; - } - *task = run_public_ipv6_provider_reconcile_task(global_ctx); - } - - async fn shutdown(&self) { - self.closing.store(true, Ordering::Release); - let task = self.task.lock().await.take(); - if let Some(task) = task { - task.shutdown().await; - } - } -} - -async fn ensure_public_ipv6_provider_reconcile_task( - global_ctx: &ArcGlobalCtx, - task_slot: &ArcPublicIpv6ProviderTaskSlot, -) { - task_slot.ensure_started(global_ctx).await; -} - -pub struct InstanceRpcServerHook { - rpc_portal_whitelist: Vec, -} - -impl InstanceRpcServerHook { - pub fn new(rpc_portal_whitelist: Option>) -> Self { - let rpc_portal_whitelist = rpc_portal_whitelist - .unwrap_or_else(|| vec!["127.0.0.0/8".parse().unwrap(), "::1/128".parse().unwrap()]); - InstanceRpcServerHook { - rpc_portal_whitelist, - } - } -} - -#[async_trait::async_trait] -impl RpcServerHook for InstanceRpcServerHook { - async fn on_new_client( - &self, - tunnel_info: Option, - ) -> Result, anyhow::Error> { - let tunnel_info = tunnel_info.ok_or_else(|| anyhow::anyhow!("tunnel info is None"))?; - - let remote_url = tunnel_info - .remote_addr - .clone() - .ok_or_else(|| anyhow::anyhow!("remote_addr is None"))?; - - let url_str = &remote_url.url; - let url = url::Url::parse(url_str) - .map_err(|e| anyhow::anyhow!("Failed to parse remote URL '{}': {}", url_str, e))?; - - let host = url - .host_str() - .ok_or_else(|| anyhow::anyhow!("No host found in remote URL '{}'", url_str))?; - - let ip_addr: IpAddr = host - .parse() - .map_err(|e| anyhow::anyhow!("Failed to parse IP address '{}': {}", host, e))?; - - for cidr in &self.rpc_portal_whitelist { - if cidr.contains(&ip_addr) { - return Ok(Some(tunnel_info)); - } - } - return Err(anyhow::anyhow!( - "Rpc portal client IP {} not in whitelist: {:?}, ignoring client.", - ip_addr, - self.rpc_portal_whitelist - )); - } -} - -#[derive(Clone)] -pub struct InstanceConfigPatcher { - global_ctx: Weak, - #[cfg(feature = "socks5")] - socks5_server: Weak, - peer_manager: Weak, - conn_manager: Weak, - public_ipv6_provider_task: ArcPublicIpv6ProviderTaskSlot, -} - -impl InstanceConfigPatcher { - fn parse_ipv6_public_addr_prefix_patch( - prefix: Option<&str>, - ) -> Result>, anyhow::Error> { - let Some(prefix) = prefix else { - return Ok(None); - }; - - let prefix = prefix.trim(); - if prefix.is_empty() { - return Ok(Some(None)); - } - - let parsed = prefix - .parse() - .with_context(|| format!("failed to parse ipv6 public address prefix: {prefix}"))?; - Ok(Some(Some(parsed))) - } - - fn effective_ipv6_for_public_ipv6_validation( - global_ctx: &ArcGlobalCtx, - patch: &crate::proto::api::config::InstanceConfigPatch, - _auto_enabled: bool, - ) -> Option { - if let Some(ipv6) = patch.ipv6 { - return Some(ipv6.into()); - } - - global_ctx.get_ipv6() - } - - fn validate_public_ipv6_patch( - global_ctx: &ArcGlobalCtx, - patch: &crate::proto::api::config::InstanceConfigPatch, - ) -> Result>, anyhow::Error> { - let parsed_prefix = - Self::parse_ipv6_public_addr_prefix_patch(patch.ipv6_public_addr_prefix.as_deref())?; - - let auto_enabled = patch - .ipv6_public_addr_auto - .unwrap_or(global_ctx.config.get_ipv6_public_addr_auto()); - let provider_enabled = patch - .ipv6_public_addr_provider - .unwrap_or(global_ctx.config.get_ipv6_public_addr_provider()); - let prefix = - parsed_prefix.unwrap_or_else(|| global_ctx.config.get_ipv6_public_addr_prefix()); - let ipv6 = Self::effective_ipv6_for_public_ipv6_validation(global_ctx, patch, auto_enabled); - - validate_public_ipv6_config_values(ipv6, provider_enabled, auto_enabled, prefix)?; - Ok(parsed_prefix) - } - - pub async fn apply_patch( - &self, - patch: crate::proto::api::config::InstanceConfigPatch, - ) -> Result<(), anyhow::Error> { - let patch_for_event = patch.clone(); - let global_ctx = weak_upgrade(&self.global_ctx)?; - let parsed_ipv6_public_addr_prefix = Self::validate_public_ipv6_patch(&global_ctx, &patch)?; - - self.patch_port_forwards(patch.port_forwards).await?; - self.patch_acl(patch.acl).await?; - self.patch_proxy_networks(patch.proxy_networks).await?; - self.patch_routes(patch.routes).await?; - self.patch_exit_nodes(patch.exit_nodes).await?; - self.patch_mapped_listeners(patch.mapped_listeners).await?; - self.patch_connector(patch.connectors).await?; - - let mut provider_config_changed = false; - if let Some(hostname) = patch.hostname { - global_ctx.set_hostname(hostname.clone()); - global_ctx.config.set_hostname(Some(hostname)); - } - if let Some(ipv4) = patch.ipv4 - && !global_ctx.config.get_dhcp() - { - global_ctx.set_ipv4(Some(ipv4.into())); - global_ctx.config.set_ipv4(Some(ipv4.into())); - } - if let Some(ipv6) = patch.ipv6 { - global_ctx.set_ipv6(Some(ipv6.into())); - global_ctx.config.set_ipv6(Some(ipv6.into())); - } - if let Some(disable_relay_data) = patch.disable_relay_data { - let mut flags = global_ctx.get_flags(); - flags.disable_relay_data = disable_relay_data; - global_ctx.set_flags(flags); - } - if let Some(enabled) = patch.ipv6_public_addr_provider { - global_ctx.config.set_ipv6_public_addr_provider(enabled); - provider_config_changed = true; - } - if let Some(enabled) = patch.ipv6_public_addr_auto { - global_ctx.config.set_ipv6_public_addr_auto(enabled); - } - if let Some(prefix) = parsed_ipv6_public_addr_prefix { - global_ctx.config.set_ipv6_public_addr_prefix(prefix); - provider_config_changed = true; - } - - global_ctx.issue_event(GlobalCtxEvent::ConfigPatched(patch_for_event)); - - if provider_config_changed { - reconcile_public_ipv6_provider_runtime(&global_ctx).await; - - if should_run_public_ipv6_provider_reconcile(&global_ctx) { - ensure_public_ipv6_provider_reconcile_task( - &global_ctx, - &self.public_ipv6_provider_task, - ) - .await; - } - } - - Ok(()) - } - - fn trace_patchables( - patches: &Vec>, - ) { - for patch in patches { - match patch.action { - Some(ConfigPatchAction::Add) | Some(ConfigPatchAction::Remove) => { - if let Some(value) = &patch.value { - tracing::info!("{:?} {:?}", patch.action, value); - } else { - tracing::warn!( - "Ignored {:?} patch with no value for type '{}'. Please ensure the patch value is provided.", - patch.action, - std::any::type_name::() - ); - } - } - Some(ConfigPatchAction::Clear) => { - tracing::info!("Clear all for type '{}'", std::any::type_name::()); - } - None => { - tracing::warn!( - "Invalid patch action for type '{}'", - std::any::type_name::() - ); - } - } - } - } - - async fn patch_port_forwards( - &self, - port_forwards: Vec, - ) -> Result<(), anyhow::Error> { - if port_forwards.is_empty() { - return Ok(()); - } - #[cfg(feature = "socks5")] - let Some(socks5_server) = self.socks5_server.upgrade() else { - return Err(anyhow::anyhow!("socks5 server not available")); - }; - let global_ctx = weak_upgrade(&self.global_ctx)?; - - let mut current_forwards = global_ctx.config.get_port_forwards(); - let patches = port_forwards.into_iter().map(Into::into).collect(); - InstanceConfigPatcher::trace_patchables(&patches); - crate::proto::api::config::patch_vec(&mut current_forwards, patches); - - global_ctx - .config - .set_port_forwards(current_forwards.clone()); - #[cfg(feature = "socks5")] - socks5_server - .reload_port_forwards(¤t_forwards) - .await - .with_context(|| "Failed to reload port forwards")?; - - Ok(()) - } - - async fn patch_acl( - &self, - acl_patch: Option, - ) -> Result<(), anyhow::Error> { - let Some(acl_patch) = acl_patch else { - return Ok(()); - }; - let global_ctx = weak_upgrade(&self.global_ctx)?; - if let Some(acl) = acl_patch.acl { - global_ctx.config.set_acl(Some(acl)); - } - if !acl_patch.tcp_whitelist.is_empty() { - let mut current_whitelist = global_ctx.config.get_tcp_whitelist(); - let patches = acl_patch - .tcp_whitelist - .into_iter() - .map(Into::into) - .collect(); - InstanceConfigPatcher::trace_patchables(&patches); - crate::proto::api::config::patch_vec(&mut current_whitelist, patches); - global_ctx.config.set_tcp_whitelist(current_whitelist); - } - if !acl_patch.udp_whitelist.is_empty() { - let mut current_whitelist = global_ctx.config.get_udp_whitelist(); - let patches = acl_patch - .udp_whitelist - .into_iter() - .map(Into::into) - .collect(); - InstanceConfigPatcher::trace_patchables(&patches); - crate::proto::api::config::patch_vec(&mut current_whitelist, patches); - global_ctx.config.set_udp_whitelist(current_whitelist); - } - global_ctx - .get_acl_filter() - .reload_rules(AclRuleBuilder::build(&global_ctx)?.as_ref()); - weak_upgrade(&self.peer_manager)? - .get_route() - .refresh_acl_groups() - .await; - Ok(()) - } - - async fn patch_proxy_networks( - &self, - proxy_networks: Vec, - ) -> Result<(), anyhow::Error> { - if proxy_networks.is_empty() { - return Ok(()); - } - let global_ctx = weak_upgrade(&self.global_ctx)?; - for proxy_network_patch in proxy_networks { - match ConfigPatchAction::try_from(proxy_network_patch.action) { - Ok(ConfigPatchAction::Add) => { - let Some(cidr) = proxy_network_patch.cidr.map(|c| c.into()) else { - tracing::warn!("Proxy network cidr is None, skipping add."); - continue; - }; - let mapped_cidr: Option = - proxy_network_patch.mapped_cidr.map(|s| s.into()); - tracing::info!("Proxy network added: {}", cidr); - global_ctx.config.add_proxy_cidr(cidr, mapped_cidr)?; - } - Ok(ConfigPatchAction::Remove) => { - let Some(cidr) = proxy_network_patch.cidr.map(|c| c.into()) else { - tracing::warn!("Proxy network cidr is None, skipping remove."); - continue; - }; - tracing::info!("Proxy network removed: {}", cidr); - global_ctx.config.remove_proxy_cidr(cidr); - } - Ok(ConfigPatchAction::Clear) => { - tracing::info!("Proxy networks cleared."); - global_ctx.config.clear_proxy_cidrs(); - } - Err(_) => { - tracing::warn!( - "Invalid proxy network action: {}", - proxy_network_patch.action - ); - } - } - } - Ok(()) - } - - async fn patch_routes( - &self, - routes: Vec, - ) -> Result<(), anyhow::Error> { - if routes.is_empty() { - return Ok(()); - } - let global_ctx = weak_upgrade(&self.global_ctx)?; - let mut current_routes = global_ctx.config.get_routes().unwrap_or_default(); - let patches = routes.into_iter().map(Into::into).collect(); - InstanceConfigPatcher::trace_patchables(&patches); - crate::proto::api::config::patch_vec(&mut current_routes, patches); - if current_routes.is_empty() { - global_ctx.config.set_routes(None); - } else { - global_ctx.config.set_routes(Some(current_routes)); - } - Ok(()) - } - - async fn patch_exit_nodes( - &self, - exit_nodes: Vec, - ) -> Result<(), anyhow::Error> { - if exit_nodes.is_empty() { - return Ok(()); - } - let global_ctx = weak_upgrade(&self.global_ctx)?; - let peer_manager = weak_upgrade(&self.peer_manager)?; - let mut current_exit_nodes = global_ctx.config.get_exit_nodes(); - let patches = exit_nodes.into_iter().map(Into::into).collect(); - InstanceConfigPatcher::trace_patchables(&patches); - crate::proto::api::config::patch_vec(&mut current_exit_nodes, patches); - global_ctx.config.set_exit_nodes(current_exit_nodes); - peer_manager.update_exit_nodes().await; - - Ok(()) - } - - async fn patch_mapped_listeners( - &self, - mapped_listeners: Vec, - ) -> Result<(), anyhow::Error> { - if mapped_listeners.is_empty() { - return Ok(()); - } - let global_ctx = weak_upgrade(&self.global_ctx)?; - let mut current_mapped_listeners = global_ctx.config.get_mapped_listeners(); - let patches = mapped_listeners.into_iter().map(Into::into).collect(); - InstanceConfigPatcher::trace_patchables(&patches); - crate::proto::api::config::patch_vec(&mut current_mapped_listeners, patches); - if current_mapped_listeners.is_empty() { - global_ctx.config.set_mapped_listeners(None); - } else { - global_ctx - .config - .set_mapped_listeners(Some(current_mapped_listeners)); - } - Ok(()) - } - - async fn patch_connector( - &self, - connectors: Vec, - ) -> Result<(), anyhow::Error> { - if connectors.is_empty() { - return Ok(()); - } - let conn_manager = weak_upgrade(&self.conn_manager)?; - for connector in connectors { - let Some(url) = connector.url.map(Into::::into) else { - tracing::warn!("Connector url is None, skipping."); - return Ok(()); - }; - match ConfigPatchAction::try_from(connector.action) { - Ok(ConfigPatchAction::Add) => { - tracing::info!("Connector added: {}", url); - conn_manager.add_connector_by_url(url).await?; - } - Ok(ConfigPatchAction::Remove) => { - tracing::info!("Connector removed: {}", url); - conn_manager.remove_connector(url).await?; - } - Ok(ConfigPatchAction::Clear) => { - tracing::info!("Connectors cleared."); - conn_manager.clear_connectors().await; - } - Err(_) => { - tracing::warn!("Invalid connector action: {}", connector.action); - } - } - } - Ok(()) - } -} - -pub struct Instance { - inst_name: String, - - id: uuid::Uuid, - - #[cfg(feature = "tun")] - nic_ctx: ArcNicCtx, - - peer_packet_receiver: Arc>, - peer_manager: Arc, - listener_manager: Arc>>, - conn_manager: Arc, - direct_conn_manager: Arc, - udp_hole_puncher: Arc>, - tcp_hole_puncher: Arc>, - - ip_proxy: Option, - - #[cfg(feature = "kcp")] - kcp_proxy_src: Option, - #[cfg(feature = "kcp")] - kcp_proxy_dst: Option, - - #[cfg(feature = "quic")] - quic_proxy: Option, - - peer_center: Arc, - - vpn_portal: Arc>>, - - #[cfg(feature = "socks5")] - socks5_server: Arc, - - proxy_cidrs_monitor: Option>, - public_ipv6_provider_task: ArcPublicIpv6ProviderTaskSlot, - - global_ctx: ArcGlobalCtx, -} - -impl Instance { - pub fn new(config: impl ConfigLoader + 'static) -> Self { - let global_ctx = Arc::new(GlobalCtx::new(config)); - - tracing::info!( - "[INIT] instance creating. config: {}", - global_ctx.config.dump() - ); - - let (peer_packet_sender, peer_packet_receiver) = create_packet_recv_chan(); - - let id = global_ctx.get_id(); - - let peer_manager = Arc::new(PeerManager::new( - RouteAlgoType::Ospf, - global_ctx.clone(), - peer_packet_sender, - )); - - peer_manager.set_allow_loopback_tunnel(false); - - let listener_manager = Arc::new(Mutex::new(ListenerManager::new( - global_ctx.clone(), - peer_manager.clone(), - ))); - - let conn_manager = Arc::new(ManualConnectorManager::new( - global_ctx.clone(), - peer_manager.clone(), - )); - - let mut direct_conn_manager = - DirectConnectorManager::new(global_ctx.clone(), peer_manager.clone()); - direct_conn_manager.run(); - let direct_conn_manager = Arc::new(direct_conn_manager); - - let udp_hole_puncher = - Arc::new(Mutex::new(UdpHolePunchConnector::new(peer_manager.clone()))); - let tcp_hole_puncher = - Arc::new(Mutex::new(TcpHolePunchConnector::new(peer_manager.clone()))); - - let peer_center = Arc::new(PeerCenterInstance::new(peer_manager.clone())); - - #[cfg(feature = "wireguard")] - let vpn_portal_inst = vpn_portal::wireguard::WireGuard::default(); - #[cfg(not(feature = "wireguard"))] - let vpn_portal_inst = vpn_portal::NullVpnPortal; - - #[cfg(feature = "socks5")] - let socks5_server = Socks5Server::new(global_ctx.clone(), peer_manager.clone(), None); - - Instance { - inst_name: global_ctx.inst_name.clone(), - id, - - peer_packet_receiver: Arc::new(Mutex::new(peer_packet_receiver)), - #[cfg(feature = "tun")] - nic_ctx: Arc::new(Mutex::new(None)), - - peer_manager, - listener_manager, - conn_manager, - direct_conn_manager, - udp_hole_puncher, - tcp_hole_puncher, - - ip_proxy: None, - #[cfg(feature = "kcp")] - kcp_proxy_src: None, - #[cfg(feature = "kcp")] - kcp_proxy_dst: None, - - #[cfg(feature = "quic")] - quic_proxy: None, - - peer_center, - - vpn_portal: Arc::new(Mutex::new(Box::new(vpn_portal_inst))), - - #[cfg(feature = "socks5")] - socks5_server, - - proxy_cidrs_monitor: None, - public_ipv6_provider_task: Arc::new(PublicIpv6ProviderTaskSlot::new()), - - global_ctx, - } - } - - pub fn get_conn_manager(&self) -> Arc { - self.conn_manager.clone() - } - - async fn add_initial_peers(&self) -> Result<(), Error> { - for peer in self.global_ctx.config.get_peers().iter() { - self.get_conn_manager() - .add_connector_by_url(peer.uri.clone()) - .await?; - } - Ok(()) - } - - async fn prepare_public_ipv6_config(&self) -> Result<(), Error> { - validate_public_ipv6_config(&self.global_ctx)?; - reconcile_public_ipv6_provider_runtime(&self.global_ctx).await; - Ok(()) - } - - // use a mock nic ctx to consume packets. - #[cfg(feature = "tun")] - async fn clear_nic_ctx( - arc_nic_ctx: ArcNicCtx, - packet_recv: Arc>, - ) { - #[cfg(feature = "magic-dns")] - if let Some(old_ctx) = arc_nic_ctx.lock().await.take() - && let Some(dns_runner) = old_ctx.magic_dns - { - dns_runner.dns_runner_cancel_token.cancel(); - tracing::debug!("cancelling dns runner task"); - let ret = dns_runner.dns_runner_task.await; - tracing::debug!("dns runner task cancelled, ret: {:?}", ret); - }; - - let mut tasks = JoinSet::new(); - tasks.spawn(async move { - let mut packet_recv = packet_recv.lock().await; - while let Ok(packet) = recv_packet_from_chan(&mut packet_recv).await { - tracing::trace!("packet consumed by mock nic ctx: {:?}", packet); - } - }); - arc_nic_ctx - .lock() - .await - .replace(NicCtxContainer::new_with_any(tasks)); - - tracing::debug!("nic ctx cleared."); - } - - #[cfg(feature = "magic-dns")] - fn create_magic_dns_runner( - peer_mgr: Arc, - tun_dev: Option, - tun_ip: Ipv4Inet, - ) -> Option { - let ctx = peer_mgr.get_global_ctx(); - if !ctx.config.get_flags().accept_dns { - return None; - } - - let runner = DnsRunner::new( - peer_mgr, - tun_dev, - tun_ip, - MAGIC_DNS_FAKE_IP.parse().unwrap(), - ); - Some(runner) - } - - #[cfg(feature = "tun")] - async fn use_new_nic_ctx( - arc_nic_ctx: ArcNicCtx, - nic_ctx: NicCtx, - #[cfg(feature = "magic-dns")] magic_dns: Option, - ) { - let mut g = arc_nic_ctx.lock().await; - *g = Some(NicCtxContainer::new( - nic_ctx, - #[cfg(feature = "magic-dns")] - magic_dns, - )); - tracing::debug!("nic ctx updated."); - } - - // Warning, if there is an IP conflict in the network when using DHCP, the IP will be automatically changed. - fn check_dhcp_ip_conflict(&self) { - use rand::Rng; - let peer_manager_c = Arc::downgrade(&self.peer_manager.clone()); - let global_ctx_c = self.get_global_ctx(); - #[cfg(feature = "tun")] - let nic_ctx = self.nic_ctx.clone(); - let _peer_packet_receiver = self.peer_packet_receiver.clone(); - tokio::spawn(async move { - let default_ipv4_addr = Ipv4Inet::new(Ipv4Addr::new(10, 126, 126, 0), 24).unwrap(); - let mut current_dhcp_ip: Option = None; - let mut next_sleep_time = 0; - let nic_closed_notifier = Arc::new(Notify::new()); - loop { - tokio::time::sleep(std::time::Duration::from_secs(next_sleep_time)).await; - - let Some(peer_manager_c) = peer_manager_c.upgrade() else { - tracing::warn!("peer manager is dropped, stop dhcp check."); - return; - }; - - if nic_closed_notifier.notified().now_or_never().is_some() { - tracing::debug!("nic ctx is closed, try recreate it"); - current_dhcp_ip = None; - } - - // do not allocate ip if no peer connected - let routes = peer_manager_c.list_routes().await; - if routes.is_empty() { - next_sleep_time = 1; - continue; - } else { - next_sleep_time = rand::thread_rng().gen_range(5..10); - } - - let mut used_ipv4 = HashSet::new(); - for route in routes { - let Some(peer_ipv4_addr) = route.ipv4_addr else { - continue; - }; - - used_ipv4.insert(peer_ipv4_addr.into()); - } - - let dhcp_inet = used_ipv4.iter().next().unwrap_or(&default_ipv4_addr); - // if old ip is already in this subnet and not conflicted, use it - if let Some(ip) = current_dhcp_ip - && ip.network() == dhcp_inet.network() - && !used_ipv4.contains(&ip) - { - continue; - } - - // find an available ip in the subnet - let candidate_ipv4_addr = dhcp_inet.network().iter().find(|ip| { - ip.address() != dhcp_inet.first_address() - && ip.address() != dhcp_inet.last_address() - && !used_ipv4.contains(ip) - }); - - if current_dhcp_ip == candidate_ipv4_addr { - continue; - } - - let last_ip = current_dhcp_ip; - tracing::debug!( - ?current_dhcp_ip, - ?candidate_ipv4_addr, - "dhcp start changing ip" - ); - - #[cfg(feature = "tun")] - Self::clear_nic_ctx(nic_ctx.clone(), _peer_packet_receiver.clone()).await; - - if let Some(ip) = candidate_ipv4_addr { - if global_ctx_c.no_tun() { - current_dhcp_ip = Some(ip); - global_ctx_c.set_ipv4(Some(ip)); - global_ctx_c - .issue_event(GlobalCtxEvent::DhcpIpv4Changed(last_ip, Some(ip))); - continue; - } - - #[cfg(all(not(mobile), feature = "tun"))] - { - let mut new_nic_ctx = NicCtx::new( - global_ctx_c.clone(), - &peer_manager_c, - _peer_packet_receiver.clone(), - nic_closed_notifier.clone(), - ); - if let Err(e) = new_nic_ctx.run(Some(ip), global_ctx_c.get_ipv6()).await { - tracing::error!( - ?current_dhcp_ip, - ?candidate_ipv4_addr, - ?e, - "add ip failed" - ); - global_ctx_c.set_ipv4(None); - continue; - } - #[cfg(feature = "magic-dns")] - let ifname = new_nic_ctx.ifname().await; - Self::use_new_nic_ctx( - nic_ctx.clone(), - new_nic_ctx, - #[cfg(feature = "magic-dns")] - Self::create_magic_dns_runner(peer_manager_c.clone(), ifname, ip), - ) - .await; - } - - current_dhcp_ip = Some(ip); - global_ctx_c.set_ipv4(Some(ip)); - global_ctx_c.issue_event(GlobalCtxEvent::DhcpIpv4Changed(last_ip, Some(ip))); - } else { - current_dhcp_ip = None; - global_ctx_c.set_ipv4(None); - global_ctx_c.issue_event(GlobalCtxEvent::DhcpIpv4Conflicted(last_ip)); - } - } - }); - } - - #[cfg(all(not(mobile), feature = "tun"))] - fn check_for_static_ip(&self, first_round_output: oneshot::Sender>) { - let ipv4_addr = self.global_ctx.get_ipv4(); - let ipv6_addr = self.global_ctx.get_ipv6(); - - // Only run if we have at least one IP address (IPv4 or IPv6) - if ipv4_addr.is_none() && ipv6_addr.is_none() { - let _ = first_round_output.send(Ok(())); - return; - } - - let nic_ctx = self.nic_ctx.clone(); - let peer_mgr = Arc::downgrade(&self.peer_manager); - let peer_packet_receiver = self.peer_packet_receiver.clone(); - - tokio::spawn(async move { - let mut output_tx = Some(first_round_output); - loop { - let close_notifier = Arc::new(Notify::new()); - { - let Some(peer_mgr) = peer_mgr.upgrade() else { - tracing::warn!("peer manager is dropped, stop static ip check."); - if let Some(output_tx) = output_tx.take() { - let _ = output_tx.send(Err(Error::Unknown)); - return; - } - return; - }; - - let mut new_nic_ctx = NicCtx::new( - peer_mgr.get_global_ctx(), - &peer_mgr, - peer_packet_receiver.clone(), - close_notifier.clone(), - ); - - if let Err(e) = new_nic_ctx.run(ipv4_addr, ipv6_addr).await { - if let Some(output_tx) = output_tx.take() { - let _ = output_tx.send(Err(e)); - return; - } - tracing::error!("failed to create new nic ctx, err: {:?}", e); - tokio::time::sleep(Duration::from_secs(1)).await; - continue; - } - - // Create Magic DNS runner only if we have IPv4 - #[cfg(feature = "magic-dns")] - { - let ifname = new_nic_ctx.ifname().await; - let dns_runner = if let Some(ipv4) = ipv4_addr { - Self::create_magic_dns_runner(peer_mgr, ifname, ipv4) - } else { - None - }; - Self::use_new_nic_ctx(nic_ctx.clone(), new_nic_ctx, dns_runner).await; - } - #[cfg(not(feature = "magic-dns"))] - Self::use_new_nic_ctx(nic_ctx.clone(), new_nic_ctx).await; - } - - if let Some(output_tx) = output_tx.take() { - let _ = output_tx.send(Ok(())); - } - - // NOTICE: make sure we do not hold the peer manager here, - while close_notifier.notified().now_or_never().is_none() { - tokio::time::sleep(Duration::from_secs(1)).await; - if peer_mgr.strong_count() == 0 { - tracing::warn!("peer manager is dropped, stop static ip check."); - return; - } - } - } - }); - } - - pub async fn run(&mut self) -> Result<(), Error> { - self.prepare_public_ipv6_config().await?; - self.listener_manager - .lock() - .await - .prepare_listeners() - .await?; - self.listener_manager.lock().await.run().await?; - self.peer_manager.run().await?; - ensure_public_ipv6_provider_reconcile_task( - &self.global_ctx, - &self.public_ipv6_provider_task, - ) - .await; - - #[cfg(feature = "tun")] - { - Self::clear_nic_ctx(self.nic_ctx.clone(), self.peer_packet_receiver.clone()).await; - - #[cfg(not(mobile))] - if !self.global_ctx.config.get_flags().no_tun { - let (output_tx, output_rx) = oneshot::channel(); - self.check_for_static_ip(output_tx); - output_rx.await.unwrap()?; - } - } - - if self.global_ctx.config.get_dhcp() { - self.check_dhcp_ip_conflict(); - } - - #[cfg(feature = "kcp")] - if self.global_ctx.get_flags().enable_kcp_proxy { - let src_proxy = KcpProxySrc::new(self.get_peer_manager()).await; - src_proxy.start().await; - self.kcp_proxy_src = Some(src_proxy); - } - - #[cfg(feature = "kcp")] - if !self.global_ctx.get_flags().disable_kcp_input { - let mut dst_proxy = KcpProxyDst::new(self.get_peer_manager()).await; - dst_proxy.start().await; - self.kcp_proxy_dst = Some(dst_proxy); - } - - #[cfg(feature = "quic")] - { - let quic_src = self.global_ctx.get_flags().enable_quic_proxy; - let quic_dst = !self.global_ctx.get_flags().disable_quic_input; - if quic_src || quic_dst { - let mut quic_proxy = QuicProxy::new(self.get_peer_manager()); - quic_proxy.run(quic_src, quic_dst).await; - self.quic_proxy = Some(quic_proxy); - } - } - - self.global_ctx - .get_acl_filter() - .reload_rules(AclRuleBuilder::build(&self.global_ctx)?.as_ref()); - - // run after tun device created, so listener can bind to tun device, which may be required by win 10 - self.ip_proxy = Some(IpProxy::new( - self.get_global_ctx(), - self.get_peer_manager(), - )?); - self.run_ip_proxy().await?; - - self.udp_hole_puncher.lock().await.run().await?; - self.tcp_hole_puncher.lock().await.run().await?; - - self.peer_center.init().await; - let route_calc = self.peer_center.get_cost_calculator(); - self.peer_manager - .get_route() - .set_route_cost_fn(route_calc) - .await; - - self.add_initial_peers().await?; - - let monitor = super::proxy_cidrs_monitor::ProxyCidrsMonitor::new( - self.peer_manager.clone(), - self.global_ctx.clone(), - ); - self.proxy_cidrs_monitor = Some(monitor.start()); - - if self.global_ctx.get_vpn_portal_cidr().is_some() { - self.run_vpn_portal().await?; - } - - #[cfg(feature = "socks5")] - self.socks5_server - .run( - #[cfg(feature = "kcp")] - self.kcp_proxy_src - .as_ref() - .map(|x| Arc::downgrade(&x.get_kcp_endpoint())), - ) - .await?; - - Ok(()) - } - - pub async fn run_ip_proxy(&mut self) -> Result<(), Error> { - if self.ip_proxy.is_none() { - return Err(anyhow::anyhow!("ip proxy not enabled.").into()); - } - self.ip_proxy.as_ref().unwrap().start().await?; - Ok(()) - } - - pub async fn run_vpn_portal(&mut self) -> Result<(), Error> { - if self.global_ctx.get_vpn_portal_cidr().is_none() { - return Err(anyhow::anyhow!("vpn portal cidr not set.").into()); - } - self.vpn_portal - .lock() - .await - .start(self.get_global_ctx(), self.get_peer_manager()) - .await?; - Ok(()) - } - - pub fn get_peer_manager(&self) -> Arc { - self.peer_manager.clone() - } - - #[cfg(feature = "ffi-dataplane")] - pub fn get_socks5_server(&self) -> Arc { - self.socks5_server.clone() - } - - pub async fn close_peer_conn( - &mut self, - peer_id: PeerId, - conn_id: &PeerConnId, - ) -> Result<(), Error> { - self.peer_manager - .get_peer_map() - .close_peer_conn(peer_id, conn_id) - .await?; - Ok(()) - } - - pub async fn wait(&self) { - self.peer_manager.wait().await; - } - - pub fn id(&self) -> uuid::Uuid { - self.id - } - - pub fn peer_id(&self) -> PeerId { - self.peer_manager.my_peer_id() - } - - fn get_vpn_portal_rpc_service( - &self, - ) -> impl VpnPortalRpc + Clone + use<> { - #[derive(Clone)] - struct VpnPortalRpcService { - peer_mgr: Weak, - vpn_portal: Weak>>, - } - - #[async_trait::async_trait] - impl VpnPortalRpc for VpnPortalRpcService { - type Controller = BaseController; - - async fn get_vpn_portal_info( - &self, - _: BaseController, - _request: GetVpnPortalInfoRequest, - ) -> Result { - let Some(vpn_portal) = self.vpn_portal.upgrade() else { - return Err(anyhow::anyhow!("vpn portal not available").into()); - }; - - let Some(peer_mgr) = self.peer_mgr.upgrade() else { - return Err(anyhow::anyhow!("peer manager not available").into()); - }; - - let vpn_portal = vpn_portal.lock().await; - let ret = GetVpnPortalInfoResponse { - vpn_portal_info: Some(VpnPortalInfo { - vpn_type: vpn_portal.name(), - client_config: vpn_portal.dump_client_config(peer_mgr).await, - connected_clients: vpn_portal.list_clients().await, - }), - }; - - Ok(ret) - } - } - - VpnPortalRpcService { - peer_mgr: Arc::downgrade(&self.peer_manager), - vpn_portal: Arc::downgrade(&self.vpn_portal), - } - } - - fn get_mapped_listener_manager_rpc_service( - &self, - ) -> impl MappedListenerManageRpc + Clone + use<> { - #[derive(Clone)] - pub struct MappedListenerManagerRpcService(Weak); - - #[async_trait::async_trait] - impl MappedListenerManageRpc for MappedListenerManagerRpcService { - type Controller = BaseController; - - async fn list_mapped_listener( - &self, - _: BaseController, - _request: ListMappedListenerRequest, - ) -> Result { - let mut ret = ListMappedListenerResponse::default(); - let urls = weak_upgrade(&self.0)?.config.get_mapped_listeners(); - let mapped_listeners: Vec = urls - .into_iter() - .map(|u| MappedListener { - url: Some(u.into()), - }) - .collect(); - ret.mappedlisteners = mapped_listeners; - Ok(ret) - } - } - - MappedListenerManagerRpcService(Arc::downgrade(&self.global_ctx)) - } - - fn get_port_forward_manager_rpc_service( - &self, - ) -> impl PortForwardManageRpc + Clone + use<> { - #[derive(Clone)] - pub struct PortForwardManagerRpcService { - global_ctx: Weak, - #[cfg(feature = "socks5")] - socks5_server: Weak, - } - - #[async_trait::async_trait] - impl PortForwardManageRpc for PortForwardManagerRpcService { - type Controller = BaseController; - - async fn list_port_forward( - &self, - _: BaseController, - _request: ListPortForwardRequest, - ) -> Result { - let forwards = weak_upgrade(&self.global_ctx)?.config.get_port_forwards(); - let cfgs: Vec = forwards.into_iter().map(Into::into).collect(); - Ok(ListPortForwardResponse { cfgs }) - } - } - - PortForwardManagerRpcService { - global_ctx: Arc::downgrade(&self.global_ctx), - #[cfg(feature = "socks5")] - socks5_server: Arc::downgrade(&self.socks5_server), - } - } - - fn get_stats_rpc_service(&self) -> impl StatsRpc + Clone + use<> { - #[derive(Clone)] - pub struct StatsRpcService { - global_ctx: Weak, - } - - #[async_trait::async_trait] - impl StatsRpc for StatsRpcService { - type Controller = BaseController; - - async fn get_stats( - &self, - _: BaseController, - _request: GetStatsRequest, - ) -> Result { - let snapshots = weak_upgrade(&self.global_ctx)? - .stats_manager() - .get_all_metrics(); - - let metrics = snapshots - .into_iter() - .map(|snapshot| { - let mut labels = std::collections::BTreeMap::new(); - for label in snapshot.labels.labels() { - labels.insert(label.key.clone(), label.value.clone()); - } - - MetricSnapshot { - name: snapshot.name_str(), - value: snapshot.value, - labels, - } - }) - .collect(); - - Ok(GetStatsResponse { metrics }) - } - - async fn get_prometheus_stats( - &self, - _: BaseController, - _request: GetPrometheusStatsRequest, - ) -> Result { - let prometheus_text = weak_upgrade(&self.global_ctx)? - .stats_manager() - .export_prometheus(); - - Ok(GetPrometheusStatsResponse { prometheus_text }) - } - } - - StatsRpcService { - global_ctx: Arc::downgrade(&self.global_ctx), - } - } - - pub fn get_config_patcher(&self) -> InstanceConfigPatcher { - InstanceConfigPatcher { - global_ctx: Arc::downgrade(&self.global_ctx), - #[cfg(feature = "socks5")] - socks5_server: Arc::downgrade(&self.socks5_server), - peer_manager: Arc::downgrade(&self.peer_manager), - conn_manager: Arc::downgrade(&self.conn_manager), - public_ipv6_provider_task: self.public_ipv6_provider_task.clone(), - } - } - - fn get_config_service(&self) -> impl ConfigRpc + Clone + use<> { - #[derive(Clone)] - pub struct ConfigRpcService { - patcher: InstanceConfigPatcher, - global_ctx: Weak, - } - - #[async_trait::async_trait] - impl ConfigRpc for ConfigRpcService { - type Controller = BaseController; - - async fn patch_config( - &self, - _: Self::Controller, - request: PatchConfigRequest, - ) -> crate::proto::rpc_types::error::Result { - let Some(patch) = request.patch else { - return Ok(PatchConfigResponse::default()); - }; - - self.patcher.apply_patch(patch).await?; - Ok(PatchConfigResponse::default()) - } - - async fn get_config( - &self, - _: Self::Controller, - _request: GetConfigRequest, - ) -> crate::proto::rpc_types::error::Result { - let global_ctx = weak_upgrade(&self.global_ctx)?; - let config = NetworkConfig::new_from_config(&global_ctx.config)?; - Ok(GetConfigResponse { - config: Some(config), - }) - } - } - - ConfigRpcService { - patcher: self.get_config_patcher(), - global_ctx: Arc::downgrade(&self.global_ctx), - } - } - - pub fn get_api_rpc_service(&self) -> impl InstanceRpcService + use<> { - use crate::proto::api::instance::*; - - #[derive(Clone)] - struct ApiRpcServiceImpl { - peer_mgr_rpc_service: A, - connector_mgr_rpc_service: B, - mapped_listener_mgr_rpc_service: C, - vpn_portal_rpc_service: D, - tcp_proxy_rpc_services: dashmap::DashMap< - String, - Arc + Send + Sync>, - >, - acl_manage_rpc_service: E, - port_forward_manage_rpc_service: F, - stats_rpc_service: G, - config_rpc_service: H, - peer_center_rpc_service: Arc, - credential_manage_rpc_service: PeerManagerRpcService, - } - - #[async_trait::async_trait] - impl< - A: PeerManageRpc + Send + Sync, - B: ConnectorManageRpc + Send + Sync, - C: MappedListenerManageRpc + Send + Sync, - D: VpnPortalRpc + Send + Sync, - E: AclManageRpc + Send + Sync, - F: PortForwardManageRpc + Send + Sync, - G: StatsRpc + Send + Sync, - H: ConfigRpc + Send + Sync, - > InstanceRpcService for ApiRpcServiceImpl - { - fn get_peer_manage_service(&self) -> &dyn PeerManageRpc { - &self.peer_mgr_rpc_service - } - - fn get_connector_manage_service( - &self, - ) -> &dyn ConnectorManageRpc { - &self.connector_mgr_rpc_service - } - - fn get_mapped_listener_manage_service( - &self, - ) -> &dyn MappedListenerManageRpc { - &self.mapped_listener_mgr_rpc_service - } - - fn get_vpn_portal_service(&self) -> &dyn VpnPortalRpc { - &self.vpn_portal_rpc_service - } - - fn get_proxy_service( - &self, - client_type: &str, - ) -> Option + Send + Sync>> - { - self.tcp_proxy_rpc_services - .get(client_type) - .map(|e| e.clone()) - } - - fn get_acl_manage_service(&self) -> &dyn AclManageRpc { - &self.acl_manage_rpc_service - } - - fn get_port_forward_manage_service( - &self, - ) -> &dyn PortForwardManageRpc { - &self.port_forward_manage_rpc_service - } - - fn get_stats_service(&self) -> &dyn StatsRpc { - &self.stats_rpc_service - } - - fn get_config_service(&self) -> &dyn ConfigRpc { - &self.config_rpc_service - } - - fn get_peer_center_service( - &self, - ) -> Arc + Send + Sync> { - self.peer_center_rpc_service.clone() - } - - fn get_credential_manage_service( - &self, - ) -> &dyn CredentialManageRpc { - &self.credential_manage_rpc_service - } - } - - ApiRpcServiceImpl { - peer_mgr_rpc_service: PeerManagerRpcService::new(self.peer_manager.clone()), - connector_mgr_rpc_service: ConnectorManagerRpcService(Arc::downgrade( - &self.conn_manager, - )), - mapped_listener_mgr_rpc_service: self.get_mapped_listener_manager_rpc_service(), - vpn_portal_rpc_service: self.get_vpn_portal_rpc_service(), - tcp_proxy_rpc_services: { - let tcp_proxy_rpc_services: dashmap::DashMap< - String, - Arc + Send + Sync>, - > = dashmap::DashMap::new(); - - if let Some(ip_proxy) = self.ip_proxy.as_ref() { - tcp_proxy_rpc_services.insert( - "tcp".to_string(), - Arc::new(TcpProxyRpcService::new(ip_proxy.tcp_proxy.clone())), - ); - } - #[cfg(feature = "kcp")] - if let Some(kcp_proxy) = self.kcp_proxy_src.as_ref() { - tcp_proxy_rpc_services.insert( - "kcp_src".to_string(), - Arc::new(TcpProxyRpcService::new(kcp_proxy.get_tcp_proxy())), - ); - } - - #[cfg(feature = "kcp")] - if let Some(kcp_proxy) = self.kcp_proxy_dst.as_ref() { - tcp_proxy_rpc_services.insert( - "kcp_dst".to_string(), - Arc::new(KcpProxyDstRpcService::new(kcp_proxy)), - ); - } - - #[cfg(feature = "quic")] - if let Some(quic_proxy) = self.quic_proxy.as_ref() { - if let Some(quic_src) = quic_proxy.src() { - tcp_proxy_rpc_services.insert( - "quic_src".to_string(), - Arc::new(TcpProxyRpcService::new(quic_src.get_tcp_proxy())), - ); - } - - if let Some(quic_dst) = quic_proxy.dst() { - tcp_proxy_rpc_services.insert( - "quic_dst".to_string(), - Arc::new(QuicProxyDstRpcService::new(quic_dst)), - ); - } - } - - tcp_proxy_rpc_services - }, - acl_manage_rpc_service: PeerManagerRpcService::new(self.peer_manager.clone()), - port_forward_manage_rpc_service: self.get_port_forward_manager_rpc_service(), - stats_rpc_service: self.get_stats_rpc_service(), - config_rpc_service: self.get_config_service(), - peer_center_rpc_service: Arc::new(self.peer_center.get_rpc_service()), - credential_manage_rpc_service: PeerManagerRpcService::new(self.peer_manager.clone()), - } - } - - pub fn get_global_ctx(&self) -> ArcGlobalCtx { - self.global_ctx.clone() - } - - pub fn get_vpn_portal_inst(&self) -> Arc>> { - self.vpn_portal.clone() - } - - #[cfg(feature = "tun")] - pub fn get_nic_ctx(&self) -> ArcNicCtx { - self.nic_ctx.clone() - } - - pub fn get_peer_packet_receiver(&self) -> Arc> { - self.peer_packet_receiver.clone() - } - - #[cfg(mobile)] - pub async fn setup_nic_ctx_for_mobile( - nic_ctx: ArcNicCtx, - global_ctx: ArcGlobalCtx, - peer_manager: Arc, - peer_packet_receiver: Arc>, - fd: i32, - ) -> Result<(), anyhow::Error> { - tracing::info!("setup_nic_ctx_for_mobile, fd: {}", fd); - Self::clear_nic_ctx(nic_ctx.clone(), peer_packet_receiver.clone()).await; - if fd <= 0 { - return Ok(()); - } - let close_notifier = Arc::new(Notify::new()); - let mut new_nic_ctx = NicCtx::new( - global_ctx.clone(), - &peer_manager, - peer_packet_receiver.clone(), - close_notifier.clone(), - ); - new_nic_ctx - .run_for_mobile(fd) - .await - .with_context(|| "add ip failed")?; - - let magic_dns_runner = if let Some(ipv4) = global_ctx.get_ipv4() { - Self::create_magic_dns_runner(peer_manager.clone(), None, ipv4) - } else { - None - }; - Self::use_new_nic_ctx(nic_ctx.clone(), new_nic_ctx, magic_dns_runner).await; - Ok(()) - } - - pub async fn clear_resources(&mut self) { - self.public_ipv6_provider_task.shutdown().await; - self.peer_manager.clear_resources().await; - #[cfg(feature = "tun")] - let _ = self.nic_ctx.lock().await.take(); - } -} - -impl Drop for Instance { - fn drop(&mut self) { - let my_peer_id = self.peer_manager.my_peer_id(); - let pm = Arc::downgrade(&self.peer_manager); - #[cfg(feature = "tun")] - let nic_ctx = self.nic_ctx.clone(); - tokio::spawn(async move { - #[cfg(feature = "tun")] - nic_ctx.lock().await.take(); - if let Some(pm) = pm.upgrade() { - pm.clear_resources().await; - }; - - let now = std::time::Instant::now(); - while now.elapsed().as_secs() < 10 { - tokio::time::sleep(std::time::Duration::from_millis(50)).await; - if pm.strong_count() == 0 { - tracing::info!( - "Instance for peer {} dropped, all resources cleared.", - my_peer_id - ); - return; - } - } - - debug_assert!( - false, - "Instance for peer {} dropped, but resources not cleared in 1 seconds.", - my_peer_id - ); - }); - } -} - -#[cfg(test)] -mod tests { - use crate::{ - common::global_ctx::tests::get_mock_global_ctx, - instance::instance::{InstanceConfigPatcher, InstanceRpcServerHook}, - proto::{api::config::InstanceConfigPatch, rpc_impl::standalone::RpcServerHook}, - }; - - #[tokio::test] - async fn test_rpc_portal_whitelist() { - use cidr::IpCidr; - - struct TestCase { - remote_url: String, - whitelist: Option>, - expected_result: bool, - } - - let test_cases: Vec = vec![ - // Test default whitelist (127.0.0.0/8, ::1/128) - TestCase { - remote_url: "tcp://127.0.0.1:15888".to_string(), - whitelist: None, - expected_result: true, - }, - TestCase { - remote_url: "tcp://127.1.2.3:15888".to_string(), - whitelist: None, - expected_result: true, - }, - TestCase { - remote_url: "tcp://192.168.1.1:15888".to_string(), - whitelist: None, - expected_result: false, - }, - // Test custom whitelist - TestCase { - remote_url: "tcp://192.168.1.10:15888".to_string(), - whitelist: Some(vec![ - "192.168.1.0/24".parse().unwrap(), - "10.0.0.0/8".parse().unwrap(), - ]), - expected_result: true, - }, - TestCase { - remote_url: "tcp://10.1.2.3:15888".to_string(), - whitelist: Some(vec![ - "192.168.1.0/24".parse().unwrap(), - "10.0.0.0/8".parse().unwrap(), - ]), - expected_result: true, - }, - TestCase { - remote_url: "tcp://172.16.0.1:15888".to_string(), - whitelist: Some(vec![ - "192.168.1.0/24".parse().unwrap(), - "10.0.0.0/8".parse().unwrap(), - ]), - expected_result: false, - }, - // Test empty whitelist (should reject all connections) - TestCase { - remote_url: "tcp://127.0.0.1:15888".to_string(), - whitelist: Some(vec![]), - expected_result: false, - }, - // Test broad whitelist (0.0.0.0/0 and ::/0 accept all IP addresses) - TestCase { - remote_url: "tcp://8.8.8.8:15888".to_string(), - whitelist: Some(vec!["0.0.0.0/0".parse().unwrap()]), - expected_result: true, - }, - // Test edge case: specific IP whitelist - TestCase { - remote_url: "tcp://192.168.1.5:15888".to_string(), - whitelist: Some(vec!["192.168.1.5/32".parse().unwrap()]), - expected_result: true, - }, - TestCase { - remote_url: "tcp://192.168.1.6:15888".to_string(), - whitelist: Some(vec!["192.168.1.5/32".parse().unwrap()]), - expected_result: false, - }, - // Test invalid URL (this case will fail during URL parsing) - TestCase { - remote_url: "invalid-url".to_string(), - whitelist: None, - expected_result: false, - }, - // Test URL without IP address (this case will fail during IP parsing) - TestCase { - remote_url: "tcp://localhost:15888".to_string(), - whitelist: None, - expected_result: false, - }, - ]; - - for case in test_cases { - let hook = InstanceRpcServerHook::new(case.whitelist.clone()); - let tunnel_info = Some(crate::proto::common::TunnelInfo { - remote_addr: Some(crate::proto::common::Url { - url: case.remote_url.clone(), - }), - ..Default::default() - }); - - let result = hook.on_new_client(tunnel_info).await; - if case.expected_result { - assert!( - result.is_ok(), - "Expected success for remote_url:{},whitelist:{:?},but got: {:?}", - case.remote_url, - case.whitelist, - result - ); - } else { - assert!( - result.is_err(), - "Expected failure for remote_url:{},whitelist:{:?},but got: {:?}", - case.remote_url, - case.whitelist, - result - ); - } - } - } - - #[tokio::test] - async fn validate_public_ipv6_patch_rejects_non_global_prefix() { - let global_ctx = get_mock_global_ctx(); - let patch = InstanceConfigPatch { - ipv6_public_addr_provider: Some(true), - ipv6_public_addr_prefix: Some("fd00::/64".to_string()), - ..Default::default() - }; - - let err = - InstanceConfigPatcher::validate_public_ipv6_patch(&global_ctx, &patch).unwrap_err(); - - assert!( - err.to_string() - .contains("not a valid global unicast IPv6 prefix") - ); - } - - #[tokio::test] - async fn public_ipv6_provider_task_slot_does_not_restart_after_shutdown() { - let global_ctx = get_mock_global_ctx(); - let slot = std::sync::Arc::new(super::PublicIpv6ProviderTaskSlot::new()); - global_ctx.config.set_ipv6_public_addr_provider(true); - global_ctx - .config - .set_ipv6_public_addr_prefix(Some("2001:db8::/48".parse().unwrap())); - - slot.shutdown().await; - super::ensure_public_ipv6_provider_reconcile_task(&global_ctx, &slot).await; - - assert!(slot.task.lock().await.is_none()); - } - - #[tokio::test] - async fn validate_public_ipv6_patch_allows_enabling_auto_with_manual_ipv6() { - let global_ctx = get_mock_global_ctx(); - global_ctx.set_ipv6(Some("fd00::1/64".parse().unwrap())); - - let patch = InstanceConfigPatch { - ipv6_public_addr_auto: Some(true), - ..Default::default() - }; - - assert!(InstanceConfigPatcher::validate_public_ipv6_patch(&global_ctx, &patch).is_ok()); - } - - #[tokio::test] - async fn validate_public_ipv6_patch_ignores_runtime_auto_ipv6_cache() { - let global_ctx = get_mock_global_ctx(); - global_ctx.config.set_ipv6_public_addr_auto(true); - global_ctx.set_ipv6(Some("2001:db8::10/64".parse().unwrap())); - - let patch = InstanceConfigPatch { - ipv6_public_addr_provider: Some(true), - ipv6_public_addr_prefix: Some("2001:db8:100::/64".to_string()), - ..Default::default() - }; - - assert!(InstanceConfigPatcher::validate_public_ipv6_patch(&global_ctx, &patch).is_ok()); - } -} diff --git a/easytier/src/instance/listeners.rs b/easytier/src/instance/listeners.rs index c0ecc60c..1192e838 100644 --- a/easytier/src/instance/listeners.rs +++ b/easytier/src/instance/listeners.rs @@ -1,417 +1,224 @@ -use std::{ - fmt::Debug, - net::IpAddr, - str::FromStr, - sync::{Arc, Weak}, -}; +use std::fmt::Debug; -use anyhow::Context; use async_trait::async_trait; -use tokio::task::JoinSet; - -use crate::{ - common::{ - error::Error, - global_ctx::{ArcGlobalCtx, GlobalCtxEvent}, - netns::NetNS, - }, - peers::peer_manager::PeerManager, - tunnel::{ - self, IpScheme, Tunnel, TunnelListener, TunnelScheme, ring::RingTunnelListener, - tcp::TcpTunnelListener, udp::UdpTunnelListener, - }, - utils::BoxExt, +use easytier_core::listener::{ + ExternalListenerFactory, ExternalListenerRequest, transport::AcceptedTransport, }; +use easytier_core::socket::SocketListener; -pub fn create_listener_by_url( - l: &url::Url, - global_ctx: ArcGlobalCtx, -) -> Result, Error> { - use crate::common::config::ConfigLoader; - let socket_mark = global_ctx.config.get_flags().socket_mark; - Ok(match l.try_into()? { - TunnelScheme::Ip(scheme) => match scheme { - IpScheme::Tcp => { - let mut l = TcpTunnelListener::new(l.clone()); - l.set_socket_mark(socket_mark); - l.boxed() - } - IpScheme::Udp => { - let mut l = UdpTunnelListener::new(l.clone()); - l.set_socket_mark(socket_mark); - l.boxed() - } - #[cfg(feature = "wireguard")] - IpScheme::Wg => { - use crate::tunnel::wireguard::{WgConfig, WgTunnelListener}; - let nid = global_ctx.get_network_identity(); - let wg_config = WgConfig::new_from_network_identity( - &nid.network_name, - &nid.network_secret.unwrap_or_default(), - ); - let mut l = WgTunnelListener::new(l.clone(), wg_config); - l.set_socket_mark(socket_mark); - l.boxed() - } - #[cfg(feature = "quic")] - IpScheme::Quic => { - // QUIC reads socket_mark from global_ctx in QuicEndpointManager - tunnel::quic::QuicTunnelListener::new(l.clone(), global_ctx.clone()).boxed() - } - #[cfg(feature = "websocket")] - IpScheme::Ws | IpScheme::Wss => { - let mut l = tunnel::websocket::WsTunnelListener::new(l.clone()); - l.set_socket_mark(socket_mark); - l.boxed() - } +#[cfg(feature = "faketcp")] +use crate::common::netns::NetNS; +use crate::socket::tcp::RuntimeTcpSocket; + +pub(crate) struct RuntimeExternalListenerFactory; + +impl ExternalListenerFactory> + for RuntimeExternalListenerFactory +{ + #[allow(clippy::match_like_matches_macro)] + fn supports_scheme(&self, scheme: &str) -> bool { + match scheme { + "faketcp" => cfg!(feature = "faketcp"), + "unix" => cfg!(unix), + _ => false, + } + } + + fn create( + &self, + request: ExternalListenerRequest, + ) -> Box>> { + match request.url.scheme() { #[cfg(feature = "faketcp")] - IpScheme::FakeTcp => tunnel::fake_tcp::FakeTcpTunnelListener::new(l.clone()).boxed(), - }, - #[cfg(unix)] - TunnelScheme::Unix => tunnel::unix::UnixSocketTunnelListener::new(l.clone()).boxed(), - _ => return Err(Error::InvalidUrl(l.to_string())), + "faketcp" => Box::new(RuntimeFakeTcpSocketListener::new( + request.url, + NetNS::from_socket_context(&request.socket_context), + )), + #[cfg(unix)] + "unix" => Box::new(RuntimeUnixStreamListener::new(request.url)), + scheme => unreachable!("core requested unsupported external listener: {scheme}"), + } + } +} + +#[cfg(unix)] +struct RuntimeUnixStreamListener { + url: url::Url, + inner: Option, +} + +#[cfg(unix)] +impl RuntimeUnixStreamListener { + fn new(url: url::Url) -> Self { + Self { url, inner: None } + } +} + +#[cfg(unix)] +fn unix_stream_remote_url(remote_addr: tokio::net::unix::SocketAddr) -> url::Url { + crate::socket::tcp::url_from_unix_socket_addr(remote_addr).unwrap_or_else(|| { + format!("unix://anonymous/{}", uuid::Uuid::new_v4()) + .parse() + .expect("synthetic Unix stream URL should be valid") }) } -pub fn is_url_host_ipv6(l: &url::Url) -> bool { - l.host_str().is_some_and(|h| h.contains(':')) -} - -pub fn is_url_host_unspecified(l: &url::Url) -> bool { - if let Ok(ip) = IpAddr::from_str(l.host_str().unwrap_or_default()) { - ip.is_unspecified() - } else { - false +#[cfg(unix)] +impl Debug for RuntimeUnixStreamListener { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("RuntimeUnixStreamListener") + .field("url", &self.url) + .field("listening", &self.inner.is_some()) + .finish() } } +#[cfg(unix)] #[async_trait] -pub trait TunnelHandlerForListener { - async fn handle_tunnel(&self, tunnel: Box) -> Result<(), Error>; -} +impl SocketListener for RuntimeUnixStreamListener { + type Accepted = AcceptedTransport; -#[async_trait] -impl TunnelHandlerForListener for PeerManager { - #[tracing::instrument] - async fn handle_tunnel(&self, tunnel: Box) -> Result<(), Error> { - self.add_tunnel_as_server(tunnel, true).await - } -} - -pub trait ListenerCreatorTrait: Fn() -> Box + Send + Sync {} -impl ListenerCreatorTrait for T where T: Fn() -> Box + Send {} -pub type ListenerCreator = Box; - -#[derive(Clone)] -struct ListenerFactory { - creator_fn: Arc, - must_succ: bool, -} - -pub struct ListenerManager { - global_ctx: ArcGlobalCtx, - net_ns: NetNS, - listeners: Vec, - peer_manager: Weak, - - tasks: JoinSet<()>, -} - -impl ListenerManager { - pub fn new(global_ctx: ArcGlobalCtx, peer_manager: Arc) -> Self { - Self { - global_ctx: global_ctx.clone(), - net_ns: global_ctx.net_ns.clone(), - listeners: Vec::new(), - peer_manager: Arc::downgrade(&peer_manager), - tasks: JoinSet::new(), + async fn listen(&mut self) -> anyhow::Result<()> { + if self.inner.is_none() { + self.inner = Some(tokio::net::UnixListener::bind(self.url.path())?); } + Ok(()) } - pub async fn prepare_listeners(&mut self) -> Result<(), Error> { - let self_id = self.global_ctx.get_id(); - self.add_listener( - move || { - Box::new(RingTunnelListener::new( - format!("ring://{}", self_id).parse().unwrap(), - )) - }, - true, - ) - .await?; - - for l in self.global_ctx.config.get_listener_uris().iter() { - let l = l.clone(); - let Ok(_) = create_listener_by_url(&l, self.global_ctx.clone()) else { - let msg = format!("failed to get listener by url: {}, maybe not supported", l); - self.global_ctx - .issue_event(GlobalCtxEvent::ListenerAddFailed(l.clone(), msg)); - continue; - }; - let ctx = self.global_ctx.clone(); - - let listener = l.clone(); - self.add_listener( - move || create_listener_by_url(&listener, ctx.clone()).unwrap(), - true, - ) + async fn accept(&mut self) -> anyhow::Result { + let (stream, remote_addr) = self + .inner + .as_ref() + .ok_or_else(|| anyhow::anyhow!("Unix stream listener is not started"))? + .accept() .await?; + Ok(AcceptedTransport::ByteStream { + socket: RuntimeTcpSocket::from_unix(stream), + local_url: self.url.clone(), + remote_url: Some(unix_stream_remote_url(remote_addr)), + }) + } - if self.global_ctx.config.get_flags().enable_ipv6 - && !is_url_host_ipv6(&l) - && is_url_host_unspecified(&l) - // quic enables dual-stack by default, may conflict with v4 listener - && l.scheme() != "quic" && l.scheme() != "faketcp" - { - let mut ipv6_listener = l.clone(); - ipv6_listener - .set_host(Some("[::]".to_string().as_str())) - .with_context(|| format!("failed to set ipv6 host for listener: {}", l))?; - let ctx = self.global_ctx.clone(); - self.add_listener( - move || create_listener_by_url(&ipv6_listener, ctx.clone()).unwrap(), - false, - ) - .await?; - } + fn local_url(&self) -> url::Url { + self.url.clone() + } +} + +#[cfg(unix)] +impl Drop for RuntimeUnixStreamListener { + fn drop(&mut self) { + let _ = std::fs::remove_file(self.url.path()); + } +} + +#[cfg(feature = "faketcp")] +struct RuntimeFakeTcpSocketListener { + net_ns: NetNS, + inner: crate::socket::fake_tcp::FakeTcpSocketListener, +} + +#[cfg(feature = "faketcp")] +impl RuntimeFakeTcpSocketListener { + fn new(url: url::Url, net_ns: NetNS) -> Self { + Self { + net_ns, + inner: crate::socket::fake_tcp::FakeTcpSocketListener::new(url), } + } +} +#[cfg(feature = "faketcp")] +impl Debug for RuntimeFakeTcpSocketListener { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("RuntimeFakeTcpSocketListener") + .field("url", &SocketListener::local_url(&self.inner)) + .finish() + } +} + +#[cfg(feature = "faketcp")] +#[async_trait] +impl SocketListener for RuntimeFakeTcpSocketListener { + type Accepted = AcceptedTransport; + + async fn listen(&mut self) -> anyhow::Result<()> { + let _guard = self.net_ns.guard(); + SocketListener::listen(&mut self.inner).await?; Ok(()) } - pub async fn add_listener( - &mut self, - creator: C, - must_succ: bool, - ) -> Result<(), Error> { - self.listeners.push(ListenerFactory { - creator_fn: Arc::new(Box::new(creator)), - must_succ, - }); - Ok(()) + async fn accept(&mut self) -> anyhow::Result { + let local_url = SocketListener::local_url(&self.inner); + let socket = self.inner.accept_socket().await?; + Ok(AcceptedTransport::Tcp { + socket: RuntimeTcpSocket::from_fake_tcp(socket), + local_url, + upgrade_permit: None, + }) } - #[tracing::instrument(skip(creator))] - async fn run_listener( - creator: Arc, - peer_manager: Weak, - global_ctx: ArcGlobalCtx, - ) { - let mut err_count = 0; - loop { - let mut l = (creator)(); - let _g = global_ctx.net_ns.guard(); - match l.listen().await { - Ok(_) => { - err_count = 0; - global_ctx.add_running_listener(l.local_url()); - global_ctx.issue_event(GlobalCtxEvent::ListenerAdded(l.local_url())); - } - Err(e) => { - tracing::error!(?e, ?l, "listener listen error"); - global_ctx.issue_event(GlobalCtxEvent::ListenerAddFailed( - l.local_url(), - format!("error: {:?}, retry listen later...", e), - )); - err_count += 1; - if err_count > 5 { - return; - } - tokio::time::sleep(std::time::Duration::from_secs(1)).await; - continue; - } - } - loop { - let ret = match l.accept().await { - Ok(ret) => ret, - Err(e) => { - global_ctx.issue_event(GlobalCtxEvent::ListenerAcceptFailed( - l.local_url(), - format!("error: {:?}, retry listen later...", e), - )); - tracing::error!(?e, ?l, "listener accept error"); - tokio::time::sleep(std::time::Duration::from_secs(1)).await; - break; - } - }; - - let tunnel_info = ret.info().unwrap(); - global_ctx.issue_event(GlobalCtxEvent::ConnectionAccepted( - tunnel_info - .local_addr - .clone() - .unwrap_or_default() - .to_string(), - tunnel_info - .remote_addr - .clone() - .unwrap_or_default() - .to_string(), - )); - tracing::info!(ret = ?ret, "conn accepted"); - let peer_manager = peer_manager.clone(); - let global_ctx = global_ctx.clone(); - tokio::spawn(async move { - let Some(peer_manager) = peer_manager.upgrade() else { - tracing::error!("peer manager is gone, cannot handle tunnel"); - return; - }; - let server_ret = peer_manager.handle_tunnel(ret).await; - if let Err(e) = &server_ret { - global_ctx.issue_event(GlobalCtxEvent::ConnectionError( - tunnel_info.local_addr.unwrap_or_default().to_string(), - tunnel_info.remote_addr.unwrap_or_default().to_string(), - e.to_string(), - )); - tracing::error!(error = ?e, "handle conn error"); - } - }); - } - } - } - - pub async fn run(&mut self) -> Result<(), Error> { - for listener in &self.listeners { - if listener.must_succ { - // try listen once - let mut l = (listener.creator_fn)(); - let _g = self.net_ns.guard(); - l.listen() - .await - .with_context(|| format!("failed to listen on {}", l.local_url()))?; - } - - self.tasks.spawn(Self::run_listener( - listener.creator_fn.clone(), - self.peer_manager.clone(), - self.global_ctx.clone(), - )); - } - - Ok(()) + fn local_url(&self) -> url::Url { + SocketListener::local_url(&self.inner) } } #[cfg(test)] mod tests { - use std::sync::atomic::{AtomicI32, Ordering}; - - use futures::{SinkExt, StreamExt}; - use tokio::time::timeout; - - use crate::{ - common::global_ctx::tests::get_mock_global_ctx, - tunnel::{TunnelConnector, TunnelError, packet_def::ZCPacket, ring::RingTunnelConnector}, - }; - use super::*; - #[derive(Debug)] - struct MockListenerHandler {} + #[test] + fn external_listener_capabilities_follow_native_build() { + let factory = RuntimeExternalListenerFactory; - #[async_trait] - impl TunnelHandlerForListener for MockListenerHandler { - async fn handle_tunnel(&self, tunnel: Box) -> Result<(), Error> { - let data = "abc"; - let (_recv, mut send) = tunnel.split(); - - let zc_packet = ZCPacket::new_with_payload(data.as_bytes()); - send.send(zc_packet).await.unwrap(); - Err(Error::Unknown) - } + assert_eq!( + factory.supports_scheme("faketcp"), + cfg!(feature = "faketcp") + ); + assert_eq!(factory.supports_scheme("unix"), cfg!(unix)); + assert!(!factory.supports_scheme("tcp")); } + #[cfg(unix)] #[tokio::test] - async fn handle_error_in_accept() { - let handler = Arc::new(MockListenerHandler {}); - let mut listener_mgr = ListenerManager::new(get_mock_global_ctx(), handler.clone()); + async fn unix_adapters_exchange_bytes_and_unlink_listener() { + use easytier_core::connectivity::composite::ConnectorRuntime; + use tokio::io::{AsyncReadExt, AsyncWriteExt}; - let ring_id = format!("ring://{}", uuid::Uuid::new_v4()); + let directory = tempfile::tempdir().unwrap(); + let path = directory.path().join("easytier.sock"); + let url: url::Url = format!("unix://{}", path.display()).parse().unwrap(); + let mut listener = RuntimeUnixStreamListener::new(url.clone()); + SocketListener::listen(&mut listener).await.unwrap(); - let ring_id_clone = ring_id.clone(); - listener_mgr - .add_listener( - move || Box::new(RingTunnelListener::new(ring_id_clone.parse().unwrap())), - true, - ) - .await - .unwrap(); - listener_mgr.run().await.unwrap(); - - tokio::time::sleep(std::time::Duration::from_secs(1)).await; - - let connect_once = |ring_id| async move { - let tunnel = RingTunnelConnector::new(ring_id).connect().await.unwrap(); - let (mut recv, _send) = tunnel.split(); - assert_eq!( - recv.next().await.unwrap().unwrap().payload(), - "abc".as_bytes() - ); - tunnel + let runtime = crate::host_runtime::native_host_runtime(); + let (accepted, connected) = tokio::join!( + SocketListener::accept(&mut listener), + runtime.connect_byte_stream(&url), + ); + let AcceptedTransport::ByteStream { + socket: mut server, + local_url, + .. + } = accepted.unwrap() + else { + panic!("Unix listener returned a non-byte-stream transport"); }; + let (mut client, _, _, _) = connected.unwrap().into_parts(); + assert_eq!(local_url, url); - timeout(std::time::Duration::from_secs(1), async move { - connect_once(ring_id.parse().unwrap()).await; - // handle tunnel fail should not impact the second connect - connect_once(ring_id.parse().unwrap()).await; - }) - .await - .unwrap(); - } + client.write_all(b"ping").await.unwrap(); + let mut request = [0; 4]; + server.read_exact(&mut request).await.unwrap(); + assert_eq!(&request, b"ping"); - #[tokio::test] - async fn retry_listen() { - let counter = Arc::new(AtomicI32::new(0)); - let drop_counter = Arc::new(AtomicI32::new(0)); - struct MockListener { - counter: Arc, - drop_counter: Arc, - } + server.write_all(b"pong").await.unwrap(); + let mut response = [0; 4]; + client.read_exact(&mut response).await.unwrap(); + assert_eq!(&response, b"pong"); - #[async_trait::async_trait] - impl TunnelListener for MockListener { - async fn listen(&mut self) -> Result<(), TunnelError> { - self.counter.fetch_add(1, Ordering::Relaxed); - Ok(()) - } - - async fn accept(&mut self) -> Result, TunnelError> { - tokio::time::sleep(std::time::Duration::from_secs(1)).await; - Err(TunnelError::BufferFull) - } - - fn local_url(&self) -> url::Url { - "mock://".parse().unwrap() - } - } - - impl Drop for MockListener { - fn drop(&mut self) { - self.drop_counter.fetch_add(1, Ordering::Relaxed); - } - } - - let handler = Arc::new(MockListenerHandler {}); - let mut listener_mgr = ListenerManager::new(get_mock_global_ctx(), handler.clone()); - let counter_clone = counter.clone(); - let drop_counter_clone = drop_counter.clone(); - listener_mgr - .add_listener( - move || { - Box::new(MockListener { - counter: counter_clone.clone(), - drop_counter: drop_counter_clone.clone(), - }) - }, - true, - ) - .await - .unwrap(); - listener_mgr.run().await.unwrap(); - - tokio::time::sleep(std::time::Duration::from_secs(3)).await; - - assert!(counter.load(Ordering::Relaxed) >= 2); - assert!(drop_counter.load(Ordering::Relaxed) >= 1); + drop(listener); + assert!(!path.exists()); } } diff --git a/easytier/src/instance/mod.rs b/easytier/src/instance/mod.rs index 2535fd1b..bc0fe9fa 100644 --- a/easytier/src/instance/mod.rs +++ b/easytier/src/instance/mod.rs @@ -1,15 +1,25 @@ +#[cfg(feature = "management")] +pub(crate) mod cli_event_logger; +pub(crate) mod composition; +pub(crate) mod config; +#[cfg(feature = "management")] +pub(crate) mod config_storage; pub mod dns_server; -#[allow(clippy::module_inception)] -pub mod instance; +pub mod factory; +pub mod host; +pub(crate) mod runtime_host; +#[cfg(test)] +pub(crate) mod test_instance; +#[cfg(feature = "upnp")] +pub(crate) mod udp_hole_punch; -pub mod listeners; +pub(crate) mod listeners; -mod public_ipv6_provider; - -pub mod proxy_cidrs_monitor; +#[cfg(feature = "public-ipv6-provider")] +pub(crate) mod public_ipv6_provider; #[cfg(feature = "tun")] pub mod virtual_nic; -#[cfg(any(windows, test))] +#[cfg(any(all(windows, feature = "tun"), test))] pub(crate) mod windows_udp_broadcast; diff --git a/easytier/src/instance/proxy_cidrs_monitor.rs b/easytier/src/instance/proxy_cidrs_monitor.rs deleted file mode 100644 index a44680c8..00000000 --- a/easytier/src/instance/proxy_cidrs_monitor.rs +++ /dev/null @@ -1,95 +0,0 @@ -use std::collections::BTreeSet; -use std::sync::{Arc, Weak}; - -use crate::common::global_ctx::{ArcGlobalCtx, GlobalCtxEvent}; -use crate::peers::peer_manager::PeerManager; -use quanta::Instant; -use tokio_util::task::AbortOnDropHandle; - -/// ProxyCidrsMonitor monitors changes in proxy CIDRs from peer routes -/// and emits GlobalCtxEvent::ProxyCidrsUpdated with added/removed diffs. -pub struct ProxyCidrsMonitor { - peer_mgr: Weak, - global_ctx: ArcGlobalCtx, -} - -impl ProxyCidrsMonitor { - pub fn new(peer_mgr: Arc, global_ctx: ArcGlobalCtx) -> Self { - Self { - peer_mgr: Arc::downgrade(&peer_mgr), - global_ctx, - } - } - - /// Collects current proxy_cidrs from peer routes, VPN portal config, and manual routes. - /// This is a static function that can be used for initial sync or recovery after Lagged errors. - pub async fn diff_proxy_cidrs( - peer_mgr: &PeerManager, - global_ctx: &ArcGlobalCtx, - cur_proxy_cidrs: &BTreeSet, - ) -> ( - BTreeSet, - Vec, - Vec, - ) { - let proxy_cidrs = if let Some(routes) = global_ctx.config.get_routes() { - // If manual routes exist, override entire proxy_cidrs - routes.into_iter().collect() - } else { - // Collect proxy_cidrs from routes - let mut proxy_cidrs = peer_mgr.list_proxy_cidrs().await; - - // Add VPN portal cidr to proxy_cidrs - if let Some(vpn_cfg) = global_ctx.config.get_vpn_portal_config() { - proxy_cidrs.insert(vpn_cfg.client_cidr); - } - - proxy_cidrs - }; - - // Calculate diff - if cur_proxy_cidrs == &proxy_cidrs { - return (proxy_cidrs, Vec::new(), Vec::new()); - } - let added = proxy_cidrs.difference(cur_proxy_cidrs).cloned().collect(); - let removed = cur_proxy_cidrs.difference(&proxy_cidrs).cloned().collect(); - - (proxy_cidrs, added, removed) - } - - /// Starts monitoring proxy_cidrs changes and emits events with diffs - pub fn start(self) -> AbortOnDropHandle<()> { - AbortOnDropHandle::new(tokio::spawn(async move { - let mut cur_proxy_cidrs = BTreeSet::new(); - let mut last_update = None::; - - loop { - tokio::time::sleep(std::time::Duration::from_secs(1)).await; - - let Some(peer_mgr) = self.peer_mgr.upgrade() else { - tracing::warn!("peer manager dropped, stopping ProxyCidrsMonitor"); - break; - }; - - // Check if route info has been updated - let last_update_time = peer_mgr.get_route_peer_info_last_update_time().await; - if last_update == Some(last_update_time) { - continue; - } - last_update = Some(last_update_time); - - let (new_proxy_cidrs, added, removed) = - Self::diff_proxy_cidrs(peer_mgr.as_ref(), &self.global_ctx, &cur_proxy_cidrs) - .await; - - cur_proxy_cidrs = new_proxy_cidrs; - - if added.is_empty() && removed.is_empty() { - continue; - } - self.global_ctx - .issue_event(GlobalCtxEvent::ProxyCidrsUpdated(added, removed)); - } - })) - } -} diff --git a/easytier/src/instance/public_ipv6_provider.rs b/easytier/src/instance/public_ipv6_provider.rs index 1be72a02..839a687d 100644 --- a/easytier/src/instance/public_ipv6_provider.rs +++ b/easytier/src/instance/public_ipv6_provider.rs @@ -1,1862 +1,49 @@ -#[cfg(target_os = "linux")] -use std::path::Path; -use std::sync::Arc; +use crate::common::global_ctx::GlobalCtxEvent; #[cfg(target_os = "linux")] -use anyhow::Context; -use cidr::{Ipv6Cidr, Ipv6Inet}; -#[cfg(target_os = "linux")] -use netlink_packet_route::route::{RouteAddress, RouteAttribute, RouteMessage, RouteType}; -use tokio_util::sync::CancellationToken; - -#[cfg(target_os = "linux")] -use crate::common::ifcfg::{ - add_ipv6_ndp_proxy, get_interface_index, list_ipv6_ndp_proxy, list_ipv6_route_messages, - remove_ipv6_ndp_proxy, -}; -use crate::common::{ - error::Error, - global_ctx::{ArcGlobalCtx, GlobalCtxEvent}, - netns::NetNS, -}; - -const PUBLIC_IPV6_PROVIDER_RECONCILE_INTERVAL: std::time::Duration = - std::time::Duration::from_secs(5); -const PUBLIC_IPV6_PROVIDER_RECONCILE_MAX_RETRIES: usize = 3; - -#[cfg(target_os = "linux")] -#[derive(Debug, Clone, PartialEq, Eq)] -struct NdpProxyTarget { - wan_iface: String, -} - -#[derive(Debug, Clone, PartialEq, Eq)] -struct PublicIpv6ProviderActiveState { - prefix: Ipv6Cidr, - #[cfg(target_os = "linux")] - ndp_proxy: Option, -} - -#[derive(Debug, Clone, PartialEq, Eq)] -enum PublicIpv6ProviderRuntimeState { - Disabled, - Pending(String), - Active(PublicIpv6ProviderActiveState), -} - -#[derive(Debug, Clone, Copy, PartialEq, Eq)] -struct PublicIpv6ProviderConfigSnapshot { - provider_enabled: bool, - configured_prefix: Option, -} - -fn read_public_ipv6_provider_config_snapshot( - global_ctx: &ArcGlobalCtx, -) -> PublicIpv6ProviderConfigSnapshot { - PublicIpv6ProviderConfigSnapshot { - provider_enabled: global_ctx.config.get_ipv6_public_addr_provider(), - configured_prefix: global_ctx.config.get_ipv6_public_addr_prefix(), - } -} - -fn should_run_public_ipv6_provider_reconcile_task( - config: PublicIpv6ProviderConfigSnapshot, -) -> bool { - config.provider_enabled -} - -pub(super) fn should_run_public_ipv6_provider_reconcile(global_ctx: &ArcGlobalCtx) -> bool { - should_run_public_ipv6_provider_reconcile_task(read_public_ipv6_provider_config_snapshot( - global_ctx, - )) -} - -fn is_global_routable_public_ipv6_prefix(prefix: Ipv6Cidr) -> bool { - let addr = prefix.first_address(); - !addr.is_loopback() - && !addr.is_multicast() - && !addr.is_unicast_link_local() - && !addr.is_unique_local() - && !addr.is_unspecified() -} - -pub(super) fn validate_public_ipv6_config_values( - _ipv6: Option, - provider_enabled: bool, - _auto_enabled: bool, - prefix: Option, -) -> Result<(), Error> { - if !provider_enabled { - return Ok(()); - } - - ensure_public_ipv6_provider_supported()?; - - if let Some(prefix) = prefix - && !is_global_routable_public_ipv6_prefix(prefix) - { - return Err(anyhow::anyhow!( - "the prefix {} is not a valid global unicast IPv6 prefix; it must be a routable address range, not a private, link-local, or multicast address", - prefix - ) - .into()); - } - - Ok(()) -} - -pub(super) fn validate_public_ipv6_config(global_ctx: &ArcGlobalCtx) -> Result<(), Error> { - validate_public_ipv6_config_values( - global_ctx.get_ipv6(), - global_ctx.config.get_ipv6_public_addr_provider(), - global_ctx.config.get_ipv6_public_addr_auto(), - global_ctx.config.get_ipv6_public_addr_prefix(), - ) -} - -fn ensure_public_ipv6_provider_supported() -> Result<(), Error> { - if cfg!(target_os = "linux") { - return Ok(()); - } - - Err(anyhow::anyhow!( - "the provider feature requires Linux; run without --ipv6-public-addr-provider on this node, or move the provider role to a Linux node. client mode (--ipv6-public-addr-auto) works on all platforms" - ) - .into()) -} - -fn public_ipv6_provider_auto_detect_error() -> Error { - anyhow::anyhow!( - "no public IPv6 prefix found on this system; set --ipv6-public-addr-prefix manually, or check that your ISP has delegated an IPv6 prefix and a default-from route exists in the kernel routing table" - ) - .into() -} - -#[cfg(target_os = "linux")] -fn read_linux_proc_bool(path: &Path) -> Result { - let value = std::fs::read_to_string(path) - .with_context(|| format!("failed to read {}", path.display()))?; - match value.trim() { - "0" => Ok(false), - "1" => Ok(true), - other => Err(anyhow::anyhow!("unexpected value '{}' in {}", other, path.display()).into()), - } -} - -#[cfg(target_os = "linux")] -fn write_linux_proc_bool(path: &Path, enabled: bool) -> Result<(), Error> { - let value = if enabled { "1\n" } else { "0\n" }; - std::fs::write(path, value).with_context(|| format!("failed to write {}", path.display()))?; - Ok(()) -} - -#[cfg(target_os = "linux")] -fn ensure_linux_ipv6_forwarding_at_paths( - all_path: &Path, - default_path: &Path, -) -> Result { - let all_enabled = read_linux_proc_bool(all_path)?; - let default_enabled = read_linux_proc_bool(default_path)?; - let mut changed = false; - - if !all_enabled { - write_linux_proc_bool(all_path, true)?; - changed = true; - } - - if !default_enabled { - write_linux_proc_bool(default_path, true)?; - changed = true; - } - - if !read_linux_proc_bool(all_path)? || !read_linux_proc_bool(default_path)? { - return Err(anyhow::anyhow!( - "failed to enable Linux IPv6 forwarding in {} and {}", - all_path.display(), - default_path.display() - ) - .into()); - } - - Ok(changed) -} - -#[cfg(target_os = "linux")] -fn ensure_linux_ipv6_forwarding() -> Result { - let all_path = Path::new("/proc/sys/net/ipv6/conf/all/forwarding"); - let default_path = Path::new("/proc/sys/net/ipv6/conf/default/forwarding"); - - ensure_linux_ipv6_forwarding_at_paths(all_path, default_path).map_err(|err| { - anyhow::anyhow!( - "public IPv6 provider requires Linux IPv6 forwarding; failed to enable net.ipv6.conf.all.forwarding=1 and net.ipv6.conf.default.forwarding=1 automatically: {}. run with sufficient privileges or set them manually", - err - ) - .into() - }) -} - -#[cfg(target_os = "linux")] -#[derive(Clone, Debug, PartialEq, Eq)] -struct DetectedIpv6Route { - dst: Option, - src: Option, - ifindex: Option, - kind: RouteType, -} - -#[cfg(target_os = "linux")] -#[derive(Clone, Debug, PartialEq, Eq)] -struct DetectedPublicIpv6Prefix { - prefix: Ipv6Cidr, - ndp_proxy: Option, -} - -#[cfg(target_os = "linux")] -fn ipv6_cidr_from_route_addr(addr: RouteAddress, prefix_len: u8) -> Option { - match addr { - RouteAddress::Inet6(addr) => Ipv6Cidr::new(addr, prefix_len).ok(), - _ => None, - } -} - -#[cfg(target_os = "linux")] -impl TryFrom for DetectedIpv6Route { - type Error = Error; - - fn try_from(message: RouteMessage) -> Result { - let dst = message.attributes.iter().find_map(|attr| match attr { - RouteAttribute::Destination(addr) => { - ipv6_cidr_from_route_addr(addr.clone(), message.header.destination_prefix_length) - } - _ => None, - }); - let src = message.attributes.iter().find_map(|attr| match attr { - RouteAttribute::Source(addr) => { - ipv6_cidr_from_route_addr(addr.clone(), message.header.source_prefix_length) - } - _ => None, - }); - let ifindex = message.attributes.iter().find_map(|attr| match attr { - RouteAttribute::Oif(index) => Some(*index), - _ => None, - }); - - Ok(Self { - dst, - src, - ifindex, - kind: message.header.kind, - }) - } -} - -#[cfg(target_os = "linux")] -fn is_ipv6_default_route(dst: Option) -> bool { - dst.is_none() || dst == Some(Ipv6Cidr::new(std::net::Ipv6Addr::UNSPECIFIED, 0).unwrap()) -} - -#[cfg(target_os = "linux")] -fn detect_public_ipv6_prefix_from_routes( - routes: &[DetectedIpv6Route], - loopback_ifindex: u32, -) -> Option { - routes - .iter() - .filter_map(|route| { - if !is_ipv6_default_route(route.dst) || route.kind != RouteType::Unicast { - return None; - } - - let prefix = route.src?; - let wan_ifindex = route.ifindex?; - if !is_global_routable_public_ipv6_prefix(prefix) { - return None; - } - - let delegated = routes.iter().any(|candidate| { - candidate.dst == Some(prefix) - && candidate.ifindex.is_some() - && candidate.ifindex != Some(wan_ifindex) - && candidate.ifindex != Some(loopback_ifindex) - && candidate.kind == RouteType::Unicast - }); - - delegated.then_some(DetectedPublicIpv6Prefix { - prefix, - ndp_proxy: None, - }) - }) - .min_by_key(|detected| detected.prefix.network_length()) -} - -#[cfg(target_os = "linux")] -#[derive(Clone, Debug, PartialEq, Eq)] -struct DetectedDefaultRouteIpv6Interface { - interface_name: String, - ifindex: u32, - address: std::net::Ipv6Addr, - prefix: Ipv6Cidr, -} - -#[cfg(target_os = "linux")] -#[derive(Clone, Debug, PartialEq, Eq)] -struct DefaultRouteIpv6InterfaceCandidate { - interface_name: String, - ifindex: u32, - address: std::net::Ipv6Addr, - prefix_len: u8, -} - -#[cfg(target_os = "linux")] -fn default_route_ifindices(routes: &[DetectedIpv6Route]) -> std::collections::BTreeSet { - routes - .iter() - .filter(|route| is_ipv6_default_route(route.dst) && route.kind == RouteType::Unicast) - .filter_map(|route| route.ifindex) - .collect() -} - -#[cfg(target_os = "linux")] -fn select_default_route_ipv6_interfaces( - candidates: impl IntoIterator, - wan_ifindices: &std::collections::BTreeSet, - max_prefix_len: u8, -) -> Vec { - candidates - .into_iter() - .filter_map(|candidate| { - if !wan_ifindices.contains(&candidate.ifindex) { - return None; - } - - if candidate.address.is_loopback() - || candidate.address.is_multicast() - || candidate.address.is_unicast_link_local() - || candidate.address.is_unique_local() - || candidate.address.is_unspecified() - { - return None; - } - - if candidate.prefix_len == 0 || candidate.prefix_len > max_prefix_len { - return None; - } - - let prefix = Ipv6Inet::new(candidate.address, candidate.prefix_len) - .ok() - .map(|inet| inet.network())?; - - Some(DetectedDefaultRouteIpv6Interface { - interface_name: candidate.interface_name, - ifindex: candidate.ifindex, - address: candidate.address, - prefix, - }) - }) - .collect() -} - -#[cfg(target_os = "linux")] -fn detect_default_route_ipv6_interfaces( - routes: &[DetectedIpv6Route], - max_prefix_len: u8, -) -> Vec { - use nix::ifaddrs::getifaddrs; - use nix::sys::socket::SockaddrLike; - use pnet::ipnetwork::ip_mask_to_prefix; - - let wan_ifindices = default_route_ifindices(routes); - if wan_ifindices.is_empty() { - return Vec::new(); - } - - let Ok(interfaces) = getifaddrs() else { - return Vec::new(); - }; - - let candidates = interfaces - .filter_map(|iface| { - let address = iface.address?; - let netmask = iface.netmask?; - let ifindex = get_interface_index(&iface.interface_name).ok()?; - - if address.family()? != nix::sys::socket::AddressFamily::Inet6 { - return None; - } - - let ipv6_addr = address.as_sockaddr_in6()?.ip(); - let netmask_ip = netmask.as_sockaddr_in6()?.ip(); - let prefix_len = ip_mask_to_prefix(std::net::IpAddr::V6(netmask_ip)).ok()?; - - Some(DefaultRouteIpv6InterfaceCandidate { - interface_name: iface.interface_name, - ifindex, - address: ipv6_addr, - prefix_len, - }) - }) - .collect::>(); - - select_default_route_ipv6_interfaces(candidates, &wan_ifindices, max_prefix_len) -} - -#[cfg(target_os = "linux")] -fn select_public_ipv6_prefix_from_default_route_interfaces( - candidates: impl IntoIterator, -) -> Option { - let iface = candidates - .into_iter() - .min_by_key(|iface| (iface.prefix.network_length(), iface.ifindex))?; - Some(DetectedPublicIpv6Prefix { - prefix: iface.prefix, - ndp_proxy: Some(NdpProxyTarget { - wan_iface: iface.interface_name, - }), - }) -} - -#[cfg(target_os = "linux")] -fn detect_public_ipv6_prefix_from_interfaces( - routes: &[DetectedIpv6Route], -) -> Option { - select_public_ipv6_prefix_from_default_route_interfaces(detect_default_route_ipv6_interfaces( - routes, 64, - )) -} - -#[cfg(target_os = "linux")] -fn ipv6_cidr_contains_cidr(outer: Ipv6Cidr, inner: Ipv6Cidr) -> bool { - outer.contains(&inner.first_address()) && outer.contains(&inner.last_address()) -} - -#[cfg(target_os = "linux")] -fn detect_configured_prefix_ndp_proxy_target( - routes: &[DetectedIpv6Route], - prefix: Ipv6Cidr, -) -> Option { - let wan_ifindices = default_route_ifindices(routes); - if wan_ifindices.is_empty() { - return None; - } - - let loopback_ifindex = get_interface_index("lo").ok(); - let routed = routes.iter().any(|route| { - route.dst == Some(prefix) - && route.kind == RouteType::Unicast - && route.ifindex.is_some_and(|ifindex| { - !wan_ifindices.contains(&ifindex) && Some(ifindex) != loopback_ifindex - }) - }); - if routed { - return None; - } - - detect_default_route_ipv6_interfaces(routes, 128) - .into_iter() - .filter(|iface| { - ipv6_cidr_contains_cidr(iface.prefix, prefix) - || (iface.prefix.network_length() == 128 && prefix.contains(&iface.address)) - }) - .min_by_key(|iface| (iface.prefix.network_length(), iface.ifindex)) - .map(|iface| NdpProxyTarget { - wan_iface: iface.interface_name, - }) -} - -#[cfg(target_os = "linux")] -fn list_detected_ipv6_routes() -> Result, Error> { - let routes = list_ipv6_route_messages().with_context(|| "failed to query linux ipv6 routes")?; - routes - .iter() - .cloned() - .map(DetectedIpv6Route::try_from) - .collect::, _>>() -} - -#[cfg(target_os = "linux")] -async fn detect_public_ipv6_prefix_linux() -> Result, Error> { - let routes = list_detected_ipv6_routes()?; - let loopback_ifindex = - get_interface_index("lo").with_context(|| "failed to resolve linux loopback ifindex")?; - - if let Some(prefix) = detect_public_ipv6_prefix_from_routes(&routes, loopback_ifindex) { - return Ok(Some(prefix)); - } - - // Fallback for DHCPv6 IA_NA / SLAAC — see https://github.com/EasyTier/EasyTier/issues/2333 - Ok(detect_public_ipv6_prefix_from_interfaces(&routes)) -} - +#[path = "public_ipv6_provider/linux.rs"] +mod platform; #[cfg(not(target_os = "linux"))] -async fn detect_public_ipv6_prefix_linux() -> Result, Error> { - Ok(None) -} +#[path = "public_ipv6_provider/unsupported.rs"] +mod platform; -fn invalid_public_ipv6_prefix_state( - prefix: Ipv6Cidr, - source: &str, -) -> PublicIpv6ProviderRuntimeState { - PublicIpv6ProviderRuntimeState::Pending(format!( - "the {} prefix {} is not a valid global unicast IPv6 prefix", - source, prefix - )) -} +pub(crate) use platform::runtime_public_ipv6_provider_platform; -#[cfg(target_os = "linux")] -fn active_public_ipv6_provider_state( - prefix: Ipv6Cidr, - ndp_proxy: Option, -) -> PublicIpv6ProviderRuntimeState { - PublicIpv6ProviderRuntimeState::Active(PublicIpv6ProviderActiveState { prefix, ndp_proxy }) -} - -#[cfg(not(target_os = "linux"))] -fn active_public_ipv6_provider_state(prefix: Ipv6Cidr) -> PublicIpv6ProviderRuntimeState { - PublicIpv6ProviderRuntimeState::Active(PublicIpv6ProviderActiveState { prefix }) -} - -#[cfg(target_os = "linux")] -async fn resolve_public_ipv6_provider_runtime_state_linux( - global_ctx: &ArcGlobalCtx, - configured_prefix: Option, -) -> PublicIpv6ProviderRuntimeState { - let _g = global_ctx.net_ns.guard(); - - if let Err(err) = ensure_linux_ipv6_forwarding() { - return PublicIpv6ProviderRuntimeState::Pending(err.to_string()); - } - - if let Some(prefix) = configured_prefix { - if !is_global_routable_public_ipv6_prefix(prefix) { - return invalid_public_ipv6_prefix_state(prefix, "configured"); - } - let ndp_proxy = match list_detected_ipv6_routes() { - Ok(routes) => detect_configured_prefix_ndp_proxy_target(&routes, prefix), - Err(err) => { - tracing::warn!( - prefix = %prefix, - ?err, - "failed to detect NDP proxy target for configured public IPv6 prefix" - ); - None - } - }; - return active_public_ipv6_provider_state(prefix, ndp_proxy); - } - - match detect_public_ipv6_prefix_linux().await { - Ok(Some(detected)) if is_global_routable_public_ipv6_prefix(detected.prefix) => { - active_public_ipv6_provider_state(detected.prefix, detected.ndp_proxy) - } - Ok(Some(detected)) => invalid_public_ipv6_prefix_state(detected.prefix, "detected"), - Ok(None) => PublicIpv6ProviderRuntimeState::Pending( - public_ipv6_provider_auto_detect_error().to_string(), - ), - Err(err) => PublicIpv6ProviderRuntimeState::Pending(err.to_string()), - } -} - -async fn resolve_public_ipv6_provider_runtime_state( - _global_ctx: &ArcGlobalCtx, - config: PublicIpv6ProviderConfigSnapshot, -) -> PublicIpv6ProviderRuntimeState { - if !config.provider_enabled { - return PublicIpv6ProviderRuntimeState::Disabled; - } - - #[cfg(target_os = "linux")] - { - return resolve_public_ipv6_provider_runtime_state_linux( - _global_ctx, - config.configured_prefix, - ) - .await; - } - - #[cfg(not(target_os = "linux"))] - { - let _ = config.configured_prefix; - PublicIpv6ProviderRuntimeState::Pending( - ensure_public_ipv6_provider_supported() - .unwrap_err() - .to_string(), - ) - } -} - -fn apply_public_ipv6_provider_runtime_state( - global_ctx: &ArcGlobalCtx, - state: &PublicIpv6ProviderRuntimeState, -) -> bool { - let next_prefix = match state { - PublicIpv6ProviderRuntimeState::Active(active) => Some(active.prefix), - PublicIpv6ProviderRuntimeState::Disabled | PublicIpv6ProviderRuntimeState::Pending(_) => { - None - } - }; - let prefix_changed = global_ctx.set_advertised_ipv6_public_addr_prefix(next_prefix); - - let next_provider_enabled = matches!(state, PublicIpv6ProviderRuntimeState::Active(_)); - let feature_changed = - global_ctx.set_ipv6_public_addr_provider_feature_flag(next_provider_enabled); - - prefix_changed || feature_changed -} - -fn try_apply_public_ipv6_provider_runtime_state( - global_ctx: &ArcGlobalCtx, - config: PublicIpv6ProviderConfigSnapshot, - state: &PublicIpv6ProviderRuntimeState, -) -> Option { - (read_public_ipv6_provider_config_snapshot(global_ctx) == config) - .then(|| apply_public_ipv6_provider_runtime_state(global_ctx, state)) -} - -fn current_public_ipv6_provider_runtime_state( - global_ctx: &ArcGlobalCtx, -) -> PublicIpv6ProviderRuntimeState { - match ( - global_ctx.get_feature_flags().ipv6_public_addr_provider, - global_ctx.get_advertised_ipv6_public_addr_prefix(), - ) { - (false, _) => PublicIpv6ProviderRuntimeState::Disabled, - #[cfg(target_os = "linux")] - (true, Some(prefix)) => active_public_ipv6_provider_state(prefix, None), - #[cfg(not(target_os = "linux"))] - (true, Some(prefix)) => active_public_ipv6_provider_state(prefix), - (true, None) => PublicIpv6ProviderRuntimeState::Pending( - "public IPv6 provider runtime is missing an advertised prefix".to_string(), - ), - } -} - -async fn reconcile_public_ipv6_provider_runtime_with_state( - global_ctx: &ArcGlobalCtx, -) -> (PublicIpv6ProviderRuntimeState, bool) { - for attempt in 0..PUBLIC_IPV6_PROVIDER_RECONCILE_MAX_RETRIES { - let config = read_public_ipv6_provider_config_snapshot(global_ctx); - let next_state = resolve_public_ipv6_provider_runtime_state(global_ctx, config).await; - - if let Some(changed) = - try_apply_public_ipv6_provider_runtime_state(global_ctx, config, &next_state) - { - return (next_state, changed); - } - - tracing::debug!( - attempt = attempt + 1, - max_retries = PUBLIC_IPV6_PROVIDER_RECONCILE_MAX_RETRIES, - "public IPv6 provider config changed during reconcile, retrying" - ); - } - - tracing::warn!( - max_retries = PUBLIC_IPV6_PROVIDER_RECONCILE_MAX_RETRIES, - "skipping public IPv6 provider reconcile because config kept changing" - ); - ( - current_public_ipv6_provider_runtime_state(global_ctx), - false, - ) -} - -pub(super) async fn reconcile_public_ipv6_provider_runtime(global_ctx: &ArcGlobalCtx) -> bool { - reconcile_public_ipv6_provider_runtime_with_state(global_ctx) - .await - .1 -} - -#[cfg(target_os = "linux")] -#[derive(Default)] -struct NdpProxyRuntime { - wan_iface: Option, - applied: std::collections::BTreeSet, -} - -#[cfg(target_os = "linux")] -impl NdpProxyRuntime { - fn reconcile( - &mut self, - global_ctx: &ArcGlobalCtx, - state: &PublicIpv6ProviderRuntimeState, - ) -> bool { - let Some((prefix, target)) = ndp_proxy_target(state) else { - return !self.clear_current(global_ctx); - }; - - let Some(tun_iface) = global_ctx.get_tun_device_name() else { - self.clear_current(global_ctx); - tracing::debug!("waiting for tun device before syncing NDP proxy entries"); - return self.cleanup_pending(); - }; - - let _g = global_ctx.net_ns.guard(); - - if self.wan_iface.as_deref() != Some(target.wan_iface.as_str()) { - if !self.clear_current_locked() { - tracing::warn!( - old_wan_iface = ?self.wan_iface, - new_wan_iface = %target.wan_iface, - remaining_entries = self.applied.len(), - "waiting to remove old NDP proxy entries before switching WAN interface" - ); - return true; - } - self.wan_iface = Some(target.wan_iface.clone()); - } - - if let Err(err) = sync_ndp_proxy_entries( - target.wan_iface.as_str(), - tun_iface.as_str(), - prefix, - &mut self.applied, - ) { - tracing::warn!( - wan_iface = %target.wan_iface, - tun_iface = %tun_iface, - ?err, - "failed to sync NDP proxy entries" - ); - } - self.cleanup_pending() - } - - fn clear_current(&mut self, global_ctx: &ArcGlobalCtx) -> bool { - self.clear_current_in_netns(&global_ctx.net_ns) - } - - fn clear_current_in_netns(&mut self, net_ns: &NetNS) -> bool { - let _g = net_ns.guard(); - self.clear_current_locked() - } - - fn clear_current_locked(&mut self) -> bool { - let Some(wan_iface) = self.wan_iface.clone() else { - return self.applied.is_empty(); - }; - - match list_ipv6_ndp_proxy(wan_iface.as_str()) { - Ok(current) => { - let candidates = self.applied.iter().copied().collect::>(); - clear_owned_ndp_proxy_entries( - wan_iface.as_str(), - ¤t, - &mut self.applied, - candidates, - ); - } - Err(err) if is_linux_missing_netlink_object_error(&err) => { - tracing::trace!( - wan_iface = %wan_iface, - ?err, - "forgetting NDP proxy ownership because WAN interface is gone" - ); - self.applied.clear(); - } - Err(err) => { - tracing::trace!( - wan_iface = %wan_iface, - ?err, - "failed to list NDP proxy entries before cleanup" - ); - } - } - - if self.applied.is_empty() { - self.wan_iface = None; - true - } else { - false - } - } - - fn cleanup_pending(&self) -> bool { - self.wan_iface.is_some() && !self.applied.is_empty() - } -} - -#[cfg(target_os = "linux")] -fn is_linux_missing_netlink_object_error(err: &Error) -> bool { - match err { - Error::IOError(err) => { - err.kind() == std::io::ErrorKind::NotFound - || matches!( - err.raw_os_error(), - Some(nix::libc::ESRCH | nix::libc::ENODEV | nix::libc::ENXIO) - ) - } +fn should_reconcile_immediately(event: &GlobalCtxEvent) -> bool { + match event { + #[cfg(feature = "management")] + GlobalCtxEvent::ConfigPatched(_) => true, + GlobalCtxEvent::TunDeviceReady(_) + | GlobalCtxEvent::TunDeviceError(_) + | GlobalCtxEvent::PublicIpv6RoutesUpdated(_, _) => true, _ => false, } } -#[cfg(target_os = "linux")] -fn clear_owned_ndp_proxy_entries( - wan_iface: &str, - current: &std::collections::BTreeSet, - applied: &mut std::collections::BTreeSet, - candidates: Vec, -) -> Option { - let mut first_err = None; - for addr in candidates { - if !current.contains(&addr) { - applied.remove(&addr); - continue; - } - - if let Err(err) = remove_ipv6_ndp_proxy(wan_iface, addr) { - if is_linux_missing_netlink_object_error(&err) { - applied.remove(&addr); - } else { - tracing::trace!( - wan_iface = %wan_iface, - addr = %addr, - ?err, - "failed to remove NDP proxy entry" - ); - first_err.get_or_insert(err); - } - } else { - applied.remove(&addr); - } - } - first_err -} - -#[cfg(target_os = "linux")] -fn ndp_proxy_target(state: &PublicIpv6ProviderRuntimeState) -> Option<(Ipv6Cidr, &NdpProxyTarget)> { - match state { - PublicIpv6ProviderRuntimeState::Active(active) => active - .ndp_proxy - .as_ref() - .map(|target| (active.prefix, target)), - PublicIpv6ProviderRuntimeState::Disabled | PublicIpv6ProviderRuntimeState::Pending(_) => { - None - } - } -} - -#[cfg(target_os = "linux")] -fn ensure_linux_ndp_proxy_enabled(wan_iface: &str) -> Result<(), Error> { - let path = Path::new("/proc/sys/net/ipv6/conf") - .join(wan_iface) - .join("proxy_ndp"); - if !read_linux_proc_bool(&path)? { - write_linux_proc_bool(&path, true)?; - tracing::info!(wan_iface = %wan_iface, "enabled Linux NDP proxy"); - } - Ok(()) -} - -#[cfg(target_os = "linux")] -fn collect_public_ipv6_tun_routes( - tun_iface: &str, - prefix: Ipv6Cidr, -) -> Result, Error> { - let tun_ifindex = match get_interface_index(tun_iface) { - Ok(ifindex) => ifindex, - Err(err) if is_linux_missing_netlink_object_error(&err) => { - tracing::debug!( - tun_iface = %tun_iface, - ?err, - "treating missing tun interface as empty public IPv6 route set" - ); - return Ok(Default::default()); - } - Err(err) => return Err(err), - }; - Ok(list_ipv6_route_messages()? - .into_iter() - .filter(|route| { - route.header.destination_prefix_length == 128 && route.header.kind == RouteType::Unicast - }) - .filter(|route| { - route - .attributes - .iter() - .any(|attr| matches!(attr, RouteAttribute::Oif(idx) if *idx == tun_ifindex)) - }) - .filter_map(|route| { - route.attributes.into_iter().find_map(|attr| match attr { - RouteAttribute::Destination(RouteAddress::Inet6(addr)) => Some(addr), - _ => None, - }) - }) - .filter(|addr| !addr.is_unicast_link_local() && prefix.contains(addr)) - .collect()) -} - -#[cfg(target_os = "linux")] -fn sync_ndp_proxy_entries( - wan_iface: &str, - tun_iface: &str, - prefix: Ipv6Cidr, - applied: &mut std::collections::BTreeSet, -) -> Result<(), Error> { - ensure_linux_ndp_proxy_enabled(wan_iface)?; - - let wanted = collect_public_ipv6_tun_routes(tun_iface, prefix)?; - let current = list_ipv6_ndp_proxy(wan_iface)?; - - let mut first_err = None; - for addr in wanted.difference(¤t) { - if let Err(err) = add_ipv6_ndp_proxy(wan_iface, *addr) { - first_err.get_or_insert(err); - } else { - applied.insert(*addr); - tracing::debug!(wan_iface = %wan_iface, addr = %addr, "added NDP proxy entry"); - } - } - - let stale = applied.difference(&wanted).copied().collect::>(); - let stale_cleanup_err = - clear_owned_ndp_proxy_entries(wan_iface, ¤t, applied, stale.clone()); - if !stale.is_empty() { - tracing::debug!( - wan_iface = %wan_iface, - stale_count = stale.len(), - remaining_count = stale.iter().filter(|addr| applied.contains(addr)).count(), - "synced stale NDP proxy entries" - ); - } - if let Some(err) = first_err.or(stale_cleanup_err) { - return Err(err); - } - - Ok(()) -} - -#[cfg(target_os = "linux")] -fn reconcile_ndp_proxy_runtime( - runtime: &mut NdpProxyRuntime, - global_ctx: &ArcGlobalCtx, - state: &PublicIpv6ProviderRuntimeState, -) -> bool { - runtime.reconcile(global_ctx, state) -} - -#[cfg(target_os = "linux")] -fn cleanup_ndp_proxy_runtime(runtime: &mut NdpProxyRuntime, net_ns: &NetNS) { - if !runtime.clear_current_in_netns(net_ns) { - tracing::warn!( - remaining_entries = runtime.applied.len(), - wan_iface = ?runtime.wan_iface, - "failed to clean all NDP proxy entries before stopping public IPv6 provider task" - ); - } -} - -#[cfg(not(target_os = "linux"))] -fn reconcile_ndp_proxy_runtime( - _runtime: &mut (), - _global_ctx: &ArcGlobalCtx, - _state: &PublicIpv6ProviderRuntimeState, -) -> bool { - false -} - -#[cfg(not(target_os = "linux"))] -fn cleanup_ndp_proxy_runtime(_runtime: &mut (), _net_ns: &NetNS) {} - -#[cfg(target_os = "linux")] -fn new_ndp_proxy_runtime() -> NdpProxyRuntime { - NdpProxyRuntime::default() -} - -#[cfg(not(target_os = "linux"))] -fn new_ndp_proxy_runtime() {} - -fn should_reconcile_immediately(event: &GlobalCtxEvent) -> bool { - matches!( - event, - GlobalCtxEvent::ConfigPatched(_) - | GlobalCtxEvent::TunDeviceReady(_) - | GlobalCtxEvent::TunDeviceError(_) - | GlobalCtxEvent::PublicIpv6RoutesUpdated(_, _) - ) -} - -async fn wait_for_public_ipv6_provider_reconcile_event( +pub(super) async fn wait_for_public_ipv6_provider_reconcile_event( event_receiver: &mut tokio::sync::broadcast::Receiver, - cancel_token: &CancellationToken, - reconcile_interval: std::time::Duration, ) -> bool { - let timer = tokio::time::sleep(reconcile_interval); - tokio::pin!(timer); loop { - tokio::select! { - _ = cancel_token.cancelled() => return false, - _ = &mut timer => return true, - recv = event_receiver.recv() => match recv { - Ok(event) if should_reconcile_immediately(&event) => return true, - Ok(_) => {} - Err(tokio::sync::broadcast::error::RecvError::Closed) => return false, - Err(tokio::sync::broadcast::error::RecvError::Lagged(_)) => { - *event_receiver = event_receiver.resubscribe(); - return true; - } + match event_receiver.recv().await { + Ok(event) if should_reconcile_immediately(&event) => return true, + Ok(_) => {} + Err(tokio::sync::broadcast::error::RecvError::Closed) => return false, + Err(tokio::sync::broadcast::error::RecvError::Lagged(_)) => { + *event_receiver = event_receiver.resubscribe(); + return true; } } } } -fn log_public_ipv6_provider_state_change( - last_state: Option<&PublicIpv6ProviderRuntimeState>, - next_state: &PublicIpv6ProviderRuntimeState, - changed: bool, -) { - if last_state != Some(next_state) { - match next_state { - PublicIpv6ProviderRuntimeState::Disabled if last_state.is_some() => { - tracing::info!("public IPv6 provider disabled"); - } - PublicIpv6ProviderRuntimeState::Disabled => {} - PublicIpv6ProviderRuntimeState::Pending(reason) => { - tracing::warn!(reason = %reason, "public IPv6 provider not ready"); - } - PublicIpv6ProviderRuntimeState::Active(active) => { - #[cfg(target_os = "linux")] - { - if let Some(target) = active.ndp_proxy.as_ref() { - tracing::info!( - prefix = %active.prefix, - wan_iface = %target.wan_iface, - "public IPv6 provider is active with NDP proxy" - ); - } else { - tracing::info!( - prefix = %active.prefix, - "public IPv6 provider is active" - ); - } - } - #[cfg(not(target_os = "linux"))] - tracing::info!(prefix = %active.prefix, "public IPv6 provider is active"); - } - } - } else if changed { - tracing::info!("public IPv6 provider runtime state changed"); - } -} - -pub(super) struct PublicIpv6ProviderReconcileTask { - cancel_token: CancellationToken, - handle: tokio::task::JoinHandle<()>, -} - -impl PublicIpv6ProviderReconcileTask { - pub(super) async fn shutdown(self) { - self.cancel_token.cancel(); - if let Err(err) = self.handle.await { - tracing::warn!( - ?err, - "public IPv6 provider reconcile task failed during shutdown" - ); - } - } -} - -pub(super) fn run_public_ipv6_provider_reconcile_task( - global_ctx: &ArcGlobalCtx, -) -> Option { - if !should_run_public_ipv6_provider_reconcile_task(read_public_ipv6_provider_config_snapshot( - global_ctx, - )) { - return None; - } - - let global_ctx = Arc::downgrade(global_ctx); - let cancel_token = CancellationToken::new(); - let task_cancel_token = cancel_token.clone(); - let handle = tokio::spawn(async move { - let Some(initial_ctx) = global_ctx.upgrade() else { - return; - }; - let net_ns = initial_ctx.net_ns.clone(); - let mut event_receiver = initial_ctx.subscribe(); - drop(initial_ctx); - let mut last_state: Option = None; - let mut ndp_proxy_runtime = new_ndp_proxy_runtime(); - - loop { - let Some(global_ctx) = global_ctx.upgrade() else { - tracing::debug!("global ctx dropped, stopping public ipv6 provider reconcile"); - break; - }; - - let (next_state, changed) = - reconcile_public_ipv6_provider_runtime_with_state(&global_ctx).await; - log_public_ipv6_provider_state_change(last_state.as_ref(), &next_state, changed); - let _ = reconcile_ndp_proxy_runtime(&mut ndp_proxy_runtime, &global_ctx, &next_state); - last_state = Some(next_state); - - if !wait_for_public_ipv6_provider_reconcile_event( - &mut event_receiver, - &task_cancel_token, - PUBLIC_IPV6_PROVIDER_RECONCILE_INTERVAL, - ) - .await - { - break; - } - } - - cleanup_ndp_proxy_runtime(&mut ndp_proxy_runtime, &net_ns); - }); - Some(PublicIpv6ProviderReconcileTask { - cancel_token, - handle, - }) -} - #[cfg(test)] mod tests { - #[cfg(target_os = "linux")] - use std::fs; - #[cfg(target_os = "linux")] - use std::path::PathBuf; - #[cfg(target_os = "linux")] - use std::process::Command; - use std::sync::Arc; - - #[cfg(target_os = "linux")] - use netlink_packet_route::route::RouteType; - - #[cfg(target_os = "linux")] - use super::{ - DefaultRouteIpv6InterfaceCandidate, DetectedIpv6Route, - detect_public_ipv6_prefix_from_interfaces, detect_public_ipv6_prefix_from_routes, - detect_public_ipv6_prefix_linux, ensure_linux_ipv6_forwarding_at_paths, - ensure_public_ipv6_provider_supported, public_ipv6_provider_auto_detect_error, - select_default_route_ipv6_interfaces, - select_public_ipv6_prefix_from_default_route_interfaces, sync_ndp_proxy_entries, - }; - - use super::{ - PublicIpv6ProviderConfigSnapshot, PublicIpv6ProviderRuntimeState, - active_public_ipv6_provider_state, read_public_ipv6_provider_config_snapshot, - should_run_public_ipv6_provider_reconcile_task, - try_apply_public_ipv6_provider_runtime_state, - }; - #[cfg(not(target_os = "linux"))] - use super::{ensure_public_ipv6_provider_supported, public_ipv6_provider_auto_detect_error}; - use crate::common::{ - config::{ConfigLoader, TomlConfigLoader}, - error::Error, - global_ctx::{GlobalCtx, GlobalCtxEvent}, - }; - - #[cfg(target_os = "linux")] - fn run_ip(args: &[&str]) { - let output = Command::new("ip") - .args(args) - .output() - .expect("failed to execute ip process"); - assert!( - output.status.success(), - "ip command failed: {:?}\nstdout: {}\nstderr: {}", - args, - String::from_utf8_lossy(&output.stdout), - String::from_utf8_lossy(&output.stderr), - ); - } - - #[cfg(target_os = "linux")] - fn test_iface_name(tag: &str) -> String { - format!("et{}{:x}", tag, std::process::id() & 0xffff) - } - - #[cfg(target_os = "linux")] - struct ScopedDummyLink { - name: String, - } - - #[cfg(target_os = "linux")] - impl ScopedDummyLink { - fn new(name: &str) -> Self { - let _ = Command::new("ip").args(["link", "del", name]).output(); - run_ip(&["link", "add", name, "type", "dummy"]); - run_ip(&["link", "set", name, "up"]); - Self { - name: name.to_string(), - } - } - } - - #[cfg(target_os = "linux")] - impl Drop for ScopedDummyLink { - fn drop(&mut self) { - let _ = Command::new("ip") - .args(["link", "del", &self.name]) - .output(); - } - } - - #[cfg(target_os = "linux")] - fn temp_forwarding_paths( - all_value: &str, - default_value: &str, - ) -> (tempfile::TempDir, PathBuf, PathBuf) { - let dir = tempfile::tempdir().unwrap(); - let all_path = dir.path().join("all_forwarding"); - let default_path = dir.path().join("default_forwarding"); - fs::write(&all_path, all_value).unwrap(); - fs::write(&default_path, default_value).unwrap(); - (dir, all_path, default_path) - } - - #[cfg(target_os = "linux")] - fn route( - dst: Option<&str>, - src: Option<&str>, - ifindex: Option, - kind: RouteType, - ) -> DetectedIpv6Route { - DetectedIpv6Route { - dst: dst.map(|cidr| cidr.parse().unwrap()), - src: src.map(|cidr| cidr.parse().unwrap()), - ifindex, - kind, - } - } - - fn active_state(prefix: cidr::Ipv6Cidr) -> PublicIpv6ProviderRuntimeState { - #[cfg(target_os = "linux")] - { - active_public_ipv6_provider_state(prefix, None) - } - #[cfg(not(target_os = "linux"))] - { - active_public_ipv6_provider_state(prefix) - } - } - - #[cfg(target_os = "linux")] - fn detected_prefix( - detected: Option, - ) -> Option { - detected.map(|detected| detected.prefix) - } - - #[cfg(target_os = "linux")] - fn iface_candidate( - interface_name: &str, - ifindex: u32, - address: &str, - prefix_len: u8, - ) -> DefaultRouteIpv6InterfaceCandidate { - DefaultRouteIpv6InterfaceCandidate { - interface_name: interface_name.to_string(), - ifindex, - address: address.parse().unwrap(), - prefix_len, - } - } - - #[cfg(target_os = "linux")] - #[test] - fn test_detect_public_ipv6_prefix_from_routes_selects_delegated_prefix() { - let routes = vec![ - route(None, Some("2001:db8:1::/56"), Some(2), RouteType::Unicast), - route(Some("2001:db8:1::/56"), None, Some(3), RouteType::Unicast), - ]; - - assert_eq!( - detected_prefix(detect_public_ipv6_prefix_from_routes(&routes, 1)), - Some("2001:db8:1::/56".parse().unwrap()) - ); - } - - #[cfg(target_os = "linux")] - #[test] - fn test_detect_public_ipv6_prefix_from_routes_rejects_non_public_prefixes() { - let routes = vec![ - route(Some("::/0"), Some("fd00::/48"), Some(2), RouteType::Unicast), - route(Some("fd00::/48"), None, Some(3), RouteType::Unicast), - route(None, Some("fe80::/64"), Some(4), RouteType::Unicast), - route(Some("fe80::/64"), None, Some(5), RouteType::Unicast), - route(None, Some("ff00::/8"), Some(6), RouteType::Unicast), - route(Some("ff00::/8"), None, Some(7), RouteType::Unicast), - route(None, Some("::/0"), Some(8), RouteType::Unicast), - route(Some("::/0"), None, Some(9), RouteType::Unicast), - ]; - - assert_eq!( - detected_prefix(detect_public_ipv6_prefix_from_routes(&routes, 1)), - None - ); - } - - #[cfg(target_os = "linux")] - #[test] - fn test_detect_public_ipv6_prefix_from_routes_requires_delegated_route() { - let routes = vec![route( - None, - Some("2001:db8:1::/56"), - Some(2), - RouteType::Unicast, - )]; - - assert_eq!( - detected_prefix(detect_public_ipv6_prefix_from_routes(&routes, 1)), - None - ); - } - - #[cfg(target_os = "linux")] - #[test] - fn test_detect_public_ipv6_prefix_from_routes_rejects_non_unicast_default_route() { - let routes = vec![ - route(None, Some("2001:db8:1::/56"), Some(2), RouteType::BlackHole), - route(Some("2001:db8:1::/56"), None, Some(3), RouteType::Unicast), - ]; - - assert_eq!( - detected_prefix(detect_public_ipv6_prefix_from_routes(&routes, 1)), - None - ); - } - - #[cfg(target_os = "linux")] - #[test] - fn test_detect_public_ipv6_prefix_from_routes_rejects_loopback_delegation() { - let routes = vec![ - route(None, Some("2001:db8:1::/56"), Some(2), RouteType::Unicast), - route(Some("2001:db8:1::/56"), None, Some(1), RouteType::Unicast), - ]; - - assert_eq!( - detected_prefix(detect_public_ipv6_prefix_from_routes(&routes, 1)), - None - ); - } - - #[cfg(target_os = "linux")] - #[test] - fn test_detect_public_ipv6_prefix_from_routes_prefers_shortest_prefix() { - let routes = vec![ - route(None, Some("2001:db8:1::/56"), Some(2), RouteType::Unicast), - route(Some("2001:db8:1::/56"), None, Some(3), RouteType::Unicast), - route(None, Some("2001:db8::/48"), Some(4), RouteType::Unicast), - route(Some("2001:db8::/48"), None, Some(5), RouteType::Unicast), - ]; - - assert_eq!( - detected_prefix(detect_public_ipv6_prefix_from_routes(&routes, 1)), - Some("2001:db8::/48".parse().unwrap()) - ); - } - - #[cfg(target_os = "linux")] - #[test] - fn test_detect_public_ipv6_prefix_from_routes_rejects_non_unicast_delegation() { - let routes = vec![ - route(None, Some("2001:db8:1::/56"), Some(2), RouteType::Unicast), - route(Some("2001:db8:1::/56"), None, Some(3), RouteType::BlackHole), - ]; - - assert_eq!( - detected_prefix(detect_public_ipv6_prefix_from_routes(&routes, 1)), - None - ); - } - - #[test] - fn test_public_ipv6_provider_auto_detect_error_mentions_manual_prefix() { - let err = public_ipv6_provider_auto_detect_error(); - let msg = err.to_string(); - - assert!(msg.contains("IPv6 prefix"), "{}", msg); - assert!(msg.contains("ipv6-public-addr-prefix"), "{}", msg); - } - - fn test_global_ctx() -> Arc { - Arc::new(GlobalCtx::new(TomlConfigLoader::default())) - } + use super::*; #[tokio::test] - async fn test_read_public_ipv6_provider_config_snapshot_reads_provider_fields() { - let global_ctx = test_global_ctx(); - let prefix = "2001:db8::/48".parse().unwrap(); - global_ctx.config.set_ipv6_public_addr_provider(true); - global_ctx.config.set_ipv6_public_addr_prefix(Some(prefix)); - - assert_eq!( - read_public_ipv6_provider_config_snapshot(&global_ctx), - PublicIpv6ProviderConfigSnapshot { - provider_enabled: true, - configured_prefix: Some(prefix), - } - ); - } - - #[test] - fn test_reconcile_task_runs_when_provider_enabled() { - assert!(!should_run_public_ipv6_provider_reconcile_task( - PublicIpv6ProviderConfigSnapshot { - provider_enabled: false, - configured_prefix: None, - } - )); - assert!(should_run_public_ipv6_provider_reconcile_task( - PublicIpv6ProviderConfigSnapshot { - provider_enabled: true, - configured_prefix: Some("2001:db8::/48".parse().unwrap()), - } - )); - assert!(should_run_public_ipv6_provider_reconcile_task( - PublicIpv6ProviderConfigSnapshot { - provider_enabled: true, - configured_prefix: None, - } - )); - } - - #[tokio::test] - async fn test_try_apply_public_ipv6_provider_runtime_state_rejects_stale_config() { - let global_ctx = test_global_ctx(); - let prefix = "2001:db8::/48".parse().unwrap(); - let config = PublicIpv6ProviderConfigSnapshot { - provider_enabled: true, - configured_prefix: Some(prefix), - }; - - global_ctx.config.set_ipv6_public_addr_provider(false); - global_ctx.config.set_ipv6_public_addr_prefix(None); - - let changed = try_apply_public_ipv6_provider_runtime_state( - &global_ctx, - config, - &active_state(prefix), - ); - - assert_eq!(changed, None); - assert_eq!(global_ctx.get_advertised_ipv6_public_addr_prefix(), None); - assert!(!global_ctx.get_feature_flags().ipv6_public_addr_provider); - } - - #[tokio::test] - async fn test_try_apply_public_ipv6_provider_runtime_state_applies_matching_config() { - let global_ctx = test_global_ctx(); - let prefix = "2001:db8::/48".parse().unwrap(); - global_ctx.config.set_ipv6_public_addr_provider(true); - global_ctx.config.set_ipv6_public_addr_prefix(Some(prefix)); - let config = read_public_ipv6_provider_config_snapshot(&global_ctx); - - let changed = try_apply_public_ipv6_provider_runtime_state( - &global_ctx, - config, - &active_state(prefix), - ); - - assert_eq!(changed, Some(true)); - assert_eq!( - global_ctx.get_advertised_ipv6_public_addr_prefix(), - Some(prefix) - ); - assert!(global_ctx.get_feature_flags().ipv6_public_addr_provider); - } - - #[cfg(target_os = "linux")] - #[test] - fn test_public_ipv6_provider_platform_check_accepts_linux() { - assert!(ensure_public_ipv6_provider_supported().is_ok()); - } - - #[cfg(target_os = "linux")] - #[test] - fn test_ensure_linux_ipv6_forwarding_enables_all_and_default() { - let (_dir, all_path, default_path) = temp_forwarding_paths("0\n", "0\n"); - - let changed = ensure_linux_ipv6_forwarding_at_paths(&all_path, &default_path).unwrap(); - - assert!(changed); - assert_eq!(fs::read_to_string(&all_path).unwrap(), "1\n"); - assert_eq!(fs::read_to_string(&default_path).unwrap(), "1\n"); - } - - #[cfg(target_os = "linux")] - #[test] - fn test_ensure_linux_ipv6_forwarding_is_noop_when_already_enabled() { - let (_dir, all_path, default_path) = temp_forwarding_paths("1\n", "1\n"); - - let changed = ensure_linux_ipv6_forwarding_at_paths(&all_path, &default_path).unwrap(); - - assert!(!changed); - assert_eq!(fs::read_to_string(&all_path).unwrap(), "1\n"); - assert_eq!(fs::read_to_string(&default_path).unwrap(), "1\n"); - } - - #[cfg(not(target_os = "linux"))] - #[test] - fn test_public_ipv6_provider_platform_check_reports_linux_only() { - let err = ensure_public_ipv6_provider_supported().unwrap_err(); - let msg = err.to_string(); - - assert!(msg.contains("Linux"), "{}", msg); - assert!(msg.contains("ipv6-public-addr-auto"), "{}", msg); - } - - #[cfg(target_os = "linux")] - #[serial_test::serial] - #[tokio::test] - async fn test_detect_public_ipv6_prefix_linux_reads_netlink_routes_from_kernel() { - let wan_if = test_iface_name("dw"); - let lan_if = test_iface_name("dl"); - let _wan = ScopedDummyLink::new(&wan_if); - let _lan = ScopedDummyLink::new(&lan_if); - - run_ip(&[ - "-6", - "addr", - "add", - "2001:db8:100:ffff::1/64", - "dev", - &wan_if, - ]); - run_ip(&[ - "-6", - "route", - "add", - "default", - "from", - "2001:db8:100::/56", - "dev", - &wan_if, - ]); - run_ip(&["-6", "route", "add", "2001:db8:100::/56", "dev", &lan_if]); - - assert_eq!( - detected_prefix(detect_public_ipv6_prefix_linux().await.unwrap()), - Some("2001:db8:100::/56".parse().unwrap()) - ); - } - - #[cfg(target_os = "linux")] - #[serial_test::serial] - #[tokio::test] - async fn test_detect_public_ipv6_prefix_linux_dhcpv6_ia_na_fallback() { - // DHCPv6 IA_NA scenario: prefix is directly on the WAN interface, - // with no delegated route on a LAN interface. - // The route-based detection should fail, and the interface-scanning - // fallback should pick up the prefix from the WAN address. - let wan_if = test_iface_name("ia"); - let _wan = ScopedDummyLink::new(&wan_if); - - run_ip(&[ - "-6", - "addr", - "add", - "2001:db8:aaaa:ffff::1/64", - "dev", - &wan_if, - ]); - run_ip(&[ - "-6", - "route", - "add", - "default", - "from", - "2001:db8:aaaa::/64", - "dev", - &wan_if, - ]); - // Also add a /48 address+route pair to verify shortest-prefix preference - run_ip(&["-6", "addr", "add", "2001:db8:bbbb::1/48", "dev", &wan_if]); - run_ip(&[ - "-6", - "route", - "add", - "default", - "from", - "2001:db8::/48", - "dev", - &wan_if, - ]); - - // NO delegated route on a LAN interface — this is the IA_NA case - // The fallback should find both prefixes via interface scanning and - // prefer the shorter /48. - let detected = detect_public_ipv6_prefix_linux().await.unwrap().unwrap(); - assert_eq!(detected.prefix, "2001:db8:bbbb::/48".parse().unwrap()); - assert_eq!(detected.ndp_proxy.unwrap().wan_iface, wan_if); - } - - #[cfg(target_os = "linux")] - #[serial_test::serial] - #[tokio::test] - async fn test_detect_public_ipv6_prefix_from_interfaces_uses_default_route_iface() { - let wan_if = test_iface_name("dw"); - let other_if = test_iface_name("do"); - let _wan = ScopedDummyLink::new(&wan_if); - let _other = ScopedDummyLink::new(&other_if); - - run_ip(&["-6", "addr", "add", "2001:db8:dddd::1/64", "dev", &wan_if]); - run_ip(&["-6", "addr", "add", "2001:db8::1/48", "dev", &other_if]); - - let wan_ifindex = crate::common::ifcfg::get_interface_index(&wan_if).unwrap(); - let other_ifindex = crate::common::ifcfg::get_interface_index(&other_if).unwrap(); - let routes = vec![ - route(None, None, Some(wan_ifindex), RouteType::Unicast), - route( - Some("2001:db8::/48"), - None, - Some(other_ifindex), - RouteType::Unicast, - ), - ]; - - let detected = detect_public_ipv6_prefix_from_interfaces(&routes) - .expect("fallback should select the default-route interface"); - assert_eq!(detected.prefix, "2001:db8:dddd::/64".parse().unwrap()); - assert_eq!(detected.ndp_proxy.unwrap().wan_iface, wan_if); - } - - #[cfg(target_os = "linux")] - #[test] - fn test_select_default_route_ipv6_interfaces_filters_candidates() { - let wan_ifindices = [2u32].into_iter().collect(); - let candidates = vec![ - iface_candidate("wan0", 2, "2001:db8:100::1", 64), - iface_candidate("nonwan0", 9, "2001:db8:200::1", 64), - iface_candidate("loopback0", 2, "::1", 128), - iface_candidate("linklocal0", 2, "fe80::1", 64), - iface_candidate("ula0", 2, "fd00::1", 64), - iface_candidate("multicast0", 2, "ff02::1", 64), - iface_candidate("unspecified0", 2, "::", 64), - iface_candidate("empty0", 2, "2001:db8:300::1", 0), - iface_candidate("host0", 2, "2001:db8:400::1", 128), - ]; - - let selected = select_default_route_ipv6_interfaces(candidates, &wan_ifindices, 64); - - assert_eq!(selected.len(), 1); - assert_eq!(selected[0].interface_name, "wan0"); - assert_eq!(selected[0].prefix, "2001:db8:100::/64".parse().unwrap()); - } - - #[cfg(target_os = "linux")] - #[test] - fn test_select_default_route_ipv6_interfaces_strips_host_bits() { - let wan_ifindices = [2u32].into_iter().collect(); - let candidates = vec![iface_candidate("wan0", 2, "2001:db8:aaaa::abcd", 64)]; - - let selected = select_default_route_ipv6_interfaces(candidates, &wan_ifindices, 64); - - assert_eq!(selected.len(), 1); - assert_eq!(selected[0].prefix, "2001:db8:aaaa::/64".parse().unwrap()); - } - - #[cfg(target_os = "linux")] - #[test] - fn test_select_public_ipv6_prefix_tie_breaks_by_lowest_ifindex() { - let wan_ifindices = [2u32, 3, 5].into_iter().collect(); - let candidates = vec![ - iface_candidate("wan5", 5, "2001:db8:5555::1", 48), - iface_candidate("wan64", 3, "2001:db8:3333::1", 64), - iface_candidate("wan2", 2, "2001:db9:2222::1", 48), - ]; - let interfaces = select_default_route_ipv6_interfaces(candidates, &wan_ifindices, 64); - - let detected = select_public_ipv6_prefix_from_default_route_interfaces(interfaces) - .expect("default-route public IPv6 prefix should be selected"); - - assert_eq!(detected.prefix, "2001:db9:2222::/48".parse().unwrap()); - assert_eq!(detected.ndp_proxy.unwrap().wan_iface, "wan2"); - } - - #[cfg(target_os = "linux")] - #[serial_test::serial] - #[tokio::test] - async fn test_configured_prefix_on_default_iface_gets_ndp_proxy_target() { - let wan_if = test_iface_name("cp"); - let _wan = ScopedDummyLink::new(&wan_if); - let configured_prefix = "2001:db8:feed::/64".parse().unwrap(); - - run_ip(&["-6", "addr", "add", "2001:db8:feed::1/128", "dev", &wan_if]); - - let ifindex = crate::common::ifcfg::get_interface_index(&wan_if).unwrap(); - let routes = vec![route(None, None, Some(ifindex), RouteType::Unicast)]; - - let target = super::detect_configured_prefix_ndp_proxy_target(&routes, configured_prefix) - .expect("configured on-link prefix should require NDP proxy"); - assert_eq!(target.wan_iface, wan_if); - } - - #[cfg(target_os = "linux")] - #[serial_test::serial] - #[tokio::test] - async fn test_configured_prefix_broader_than_default_iface_does_not_get_ndp_proxy_target() { - let wan_if = test_iface_name("cb"); - let _wan = ScopedDummyLink::new(&wan_if); - let configured_prefix = "2001:db8:beef::/48".parse().unwrap(); - - run_ip(&["-6", "addr", "add", "2001:db8:beef:1::1/64", "dev", &wan_if]); - - let ifindex = crate::common::ifcfg::get_interface_index(&wan_if).unwrap(); - let routes = vec![route(None, None, Some(ifindex), RouteType::Unicast)]; - - assert_eq!( - super::detect_configured_prefix_ndp_proxy_target(&routes, configured_prefix), - None - ); - } - - #[cfg(target_os = "linux")] - #[serial_test::serial] - #[tokio::test] - async fn test_detect_public_ipv6_prefix_linux_dhcpv6_ia_na_single_prefix() { - // DHCPv6 IA_NA: the WAN interface has a global prefix with a - // default route. Use a /48 dummy prefix so the fallback prefers it - // over any real /64 on the test machine. - let wan_if = test_iface_name("ib"); - let _wan = ScopedDummyLink::new(&wan_if); - - run_ip(&["-6", "addr", "add", "2001:db8:cccc::1/48", "dev", &wan_if]); - run_ip(&["-6", "route", "add", "default", "dev", &wan_if]); - - let detected = detect_public_ipv6_prefix_linux().await.unwrap().unwrap(); - assert_eq!(detected.prefix, "2001:db8:cccc::/48".parse().unwrap()); - assert_eq!(detected.ndp_proxy.unwrap().wan_iface, wan_if); - } - - #[cfg(target_os = "linux")] - #[serial_test::serial] - #[tokio::test] - async fn test_detect_public_ipv6_prefix_from_interfaces_skips_non_global() { - // Create a dummy interface with only a link-local address. - // The interface fallback should return None because there is no - // global unicast address. - let iface = test_iface_name("ng"); - let _link = ScopedDummyLink::new(&iface); - - // Bring up the interface so it auto-configures a link-local address - run_ip(&["link", "set", &iface, "up"]); - let ifindex = crate::common::ifcfg::get_interface_index(&iface).unwrap(); - let routes = vec![route(None, None, Some(ifindex), RouteType::Unicast)]; - - // No global address added — only link-local should be present - let result = detect_public_ipv6_prefix_from_interfaces(&routes); - assert_eq!(detected_prefix(result), None); - } - - #[cfg(target_os = "linux")] - #[serial_test::serial] - #[tokio::test] - async fn test_ndp_proxy_sync_uses_configured_tun_iface_without_shell_neigh() { - let wan_if = test_iface_name("nw"); - let tun_if = test_iface_name("nt"); - let _wan = ScopedDummyLink::new(&wan_if); - let _tun = ScopedDummyLink::new(&tun_if); - let addr = "2001:db8:abcd::123".parse::().unwrap(); - let prefix = "2001:db8:abcd::/64".parse().unwrap(); - let mut applied = std::collections::BTreeSet::new(); - - run_ip(&["-6", "route", "add", &format!("{addr}/128"), "dev", &tun_if]); - - sync_ndp_proxy_entries(&wan_if, &tun_if, prefix, &mut applied).unwrap(); - assert!( - crate::common::ifcfg::list_ipv6_ndp_proxy(&wan_if) - .unwrap() - .contains(&addr) - ); - assert!(applied.contains(&addr)); - - run_ip(&["-6", "route", "del", &format!("{addr}/128"), "dev", &tun_if]); - sync_ndp_proxy_entries(&wan_if, &tun_if, prefix, &mut applied).unwrap(); - assert!( - !crate::common::ifcfg::list_ipv6_ndp_proxy(&wan_if) - .unwrap() - .contains(&addr) - ); - assert!(!applied.contains(&addr)); - } - - #[cfg(target_os = "linux")] - #[serial_test::serial] - #[tokio::test] - async fn test_ndp_proxy_sync_does_not_delete_preexisting_proxy_entry() { - let wan_if = test_iface_name("pw"); - let tun_if = test_iface_name("pt"); - let _wan = ScopedDummyLink::new(&wan_if); - let _tun = ScopedDummyLink::new(&tun_if); - let addr = "2001:db8:beef::123".parse::().unwrap(); - let prefix = "2001:db8:beef::/64".parse().unwrap(); - let mut applied = std::collections::BTreeSet::new(); - - super::ensure_linux_ndp_proxy_enabled(&wan_if).unwrap(); - crate::common::ifcfg::add_ipv6_ndp_proxy(&wan_if, addr).unwrap(); - run_ip(&["-6", "route", "add", &format!("{addr}/128"), "dev", &tun_if]); - - sync_ndp_proxy_entries(&wan_if, &tun_if, prefix, &mut applied).unwrap(); - assert!( - crate::common::ifcfg::list_ipv6_ndp_proxy(&wan_if) - .unwrap() - .contains(&addr) - ); - assert!(!applied.contains(&addr)); - - run_ip(&["-6", "route", "del", &format!("{addr}/128"), "dev", &tun_if]); - sync_ndp_proxy_entries(&wan_if, &tun_if, prefix, &mut applied).unwrap(); - assert!( - crate::common::ifcfg::list_ipv6_ndp_proxy(&wan_if) - .unwrap() - .contains(&addr) - ); - assert!(!applied.contains(&addr)); - - crate::common::ifcfg::remove_ipv6_ndp_proxy(&wan_if, addr).unwrap(); - } - - #[cfg(target_os = "linux")] - #[serial_test::serial] - #[tokio::test] - async fn test_ndp_proxy_sync_removes_owned_entry_when_tun_iface_is_gone() { - let wan_if = test_iface_name("gw"); - let _wan = ScopedDummyLink::new(&wan_if); - let addr = "2001:db8:face::123".parse::().unwrap(); - let prefix = "2001:db8:face::/64".parse().unwrap(); - let mut applied = std::collections::BTreeSet::from([addr]); - - super::ensure_linux_ndp_proxy_enabled(&wan_if).unwrap(); - crate::common::ifcfg::add_ipv6_ndp_proxy(&wan_if, addr).unwrap(); - - sync_ndp_proxy_entries(&wan_if, "missing-easytier-tun", prefix, &mut applied).unwrap(); - assert!( - !crate::common::ifcfg::list_ipv6_ndp_proxy(&wan_if) - .unwrap() - .contains(&addr) - ); - assert!(applied.is_empty()); - } - - #[cfg(target_os = "linux")] - #[serial_test::serial] - #[tokio::test] - async fn test_cleanup_ndp_proxy_runtime_removes_owned_entry_on_task_exit() { - let wan_if = test_iface_name("cw"); - let _wan = ScopedDummyLink::new(&wan_if); - let addr = "2001:db8:cafe::123".parse::().unwrap(); - let mut runtime = super::NdpProxyRuntime { - wan_iface: Some(wan_if.clone()), - applied: std::collections::BTreeSet::from([addr]), - }; - - super::ensure_linux_ndp_proxy_enabled(&wan_if).unwrap(); - crate::common::ifcfg::add_ipv6_ndp_proxy(&wan_if, addr).unwrap(); - - super::cleanup_ndp_proxy_runtime(&mut runtime, &crate::common::netns::NetNS::new(None)); - assert!( - !crate::common::ifcfg::list_ipv6_ndp_proxy(&wan_if) - .unwrap() - .contains(&addr) - ); - assert!(!runtime.cleanup_pending()); - } - - #[tokio::test] - async fn test_wait_for_reconcile_ignores_unrelated_events_without_resetting_timer() { + async fn wait_for_reconcile_ignores_unrelated_events() { let (tx, mut rx) = tokio::sync::broadcast::channel(16); - let cancel_token = tokio_util::sync::CancellationToken::new(); + let trigger_tx = tx.clone(); let spam_task = tokio::spawn(async move { loop { if tx.send(GlobalCtxEvent::PeerAdded(1)).is_err() { @@ -1865,184 +52,20 @@ mod tests { tokio::time::sleep(std::time::Duration::from_millis(5)).await; } }); + let trigger_task = tokio::spawn(async move { + tokio::time::sleep(std::time::Duration::from_millis(50)).await; + let _ = trigger_tx.send(GlobalCtxEvent::TunDeviceReady("et-test".to_owned())); + }); let reconciled = tokio::time::timeout( std::time::Duration::from_millis(250), - super::wait_for_public_ipv6_provider_reconcile_event( - &mut rx, - &cancel_token, - std::time::Duration::from_millis(50), - ), + wait_for_public_ipv6_provider_reconcile_event(&mut rx), ) .await - .expect("unrelated events should not keep resetting the reconcile timer"); + .expect("a relevant event should wake the reconcile loop"); spam_task.abort(); + trigger_task.await.unwrap(); assert!(reconciled); } - - #[cfg(target_os = "linux")] - async fn wait_for_ndp_proxy_entry(wan_if: &str, addr: std::net::Ipv6Addr, present: bool) { - for _ in 0..50 { - let current = crate::common::ifcfg::list_ipv6_ndp_proxy(wan_if).unwrap(); - if current.contains(&addr) == present { - return; - } - tokio::time::sleep(std::time::Duration::from_millis(20)).await; - } - - let current = crate::common::ifcfg::list_ipv6_ndp_proxy(wan_if).unwrap(); - assert_eq!(current.contains(&addr), present); - } - - #[cfg(target_os = "linux")] - #[serial_test::serial] - #[tokio::test] - async fn test_reconcile_task_shutdown_removes_owned_ndp_proxy_entry() { - let wan_if = test_iface_name("tw"); - let tun_if = test_iface_name("tt"); - let _wan = ScopedDummyLink::new(&wan_if); - let _tun = ScopedDummyLink::new(&tun_if); - let prefix = "2001:db8:fade::/64".parse().unwrap(); - let wan_addr = "2001:db8:fade::1"; - let leased_addr = "2001:db8:fade::123".parse::().unwrap(); - let global_ctx = test_global_ctx(); - - run_ip(&[ - "-6", - "addr", - "add", - &format!("{wan_addr}/128"), - "dev", - &wan_if, - ]); - run_ip(&["-6", "route", "add", "default", "dev", &wan_if]); - run_ip(&[ - "-6", - "route", - "add", - &format!("{leased_addr}/128"), - "dev", - &tun_if, - ]); - - global_ctx.config.set_ipv6_public_addr_provider(true); - global_ctx.config.set_ipv6_public_addr_prefix(Some(prefix)); - global_ctx.set_tun_device_ready(tun_if); - - let task = super::run_public_ipv6_provider_reconcile_task(&global_ctx) - .expect("provider task should start when provider is enabled"); - wait_for_ndp_proxy_entry(&wan_if, leased_addr, true).await; - - task.shutdown().await; - wait_for_ndp_proxy_entry(&wan_if, leased_addr, false).await; - } - - #[cfg(target_os = "linux")] - #[test] - fn test_missing_netlink_object_errors_release_ndp_ownership() { - for errno in [ - nix::libc::ENOENT, - nix::libc::ESRCH, - nix::libc::ENODEV, - nix::libc::ENXIO, - ] { - let err = Error::IOError(std::io::Error::from_raw_os_error(errno)); - assert!(super::is_linux_missing_netlink_object_error(&err)); - } - } - - #[cfg(target_os = "linux")] - #[test] - fn test_clear_owned_ndp_proxy_entries_forgets_already_absent_entry() { - let addr = "2001:db8:dead::111".parse::().unwrap(); - let mut applied = std::collections::BTreeSet::from([addr]); - let current = std::collections::BTreeSet::new(); - - assert!( - super::clear_owned_ndp_proxy_entries("missing", ¤t, &mut applied, vec![addr],) - .is_none() - ); - assert!(!applied.contains(&addr)); - } - - #[cfg(target_os = "linux")] - #[test] - fn test_ndp_proxy_runtime_forgets_entries_when_wan_interface_is_gone() { - let addr = "2001:db8:dead::123".parse::().unwrap(); - let mut runtime = super::NdpProxyRuntime { - wan_iface: Some(test_iface_name("missing")), - applied: std::collections::BTreeSet::from([addr]), - }; - - assert!(runtime.clear_current_locked()); - assert!(runtime.wan_iface.is_none()); - assert!(runtime.applied.is_empty()); - } - - #[cfg(target_os = "linux")] - #[tokio::test] - async fn test_ndp_proxy_runtime_finishes_cleanup_when_wan_interface_is_gone_after_disable() { - let addr = "2001:db8:dead::456".parse::().unwrap(); - let global_ctx = test_global_ctx(); - let mut runtime = super::NdpProxyRuntime { - wan_iface: Some(test_iface_name("missing")), - applied: std::collections::BTreeSet::from([addr]), - }; - - assert!(!runtime.reconcile(&global_ctx, &PublicIpv6ProviderRuntimeState::Disabled)); - assert!(!runtime.cleanup_pending()); - } - - #[cfg(target_os = "linux")] - #[serial_test::serial] - #[tokio::test] - async fn test_detect_public_ipv6_prefix_linux_prefers_shortest_prefix_from_kernel() { - let wan_if_1 = test_iface_name("sw1"); - let lan_if_1 = test_iface_name("sl1"); - let wan_if_2 = test_iface_name("sw2"); - let lan_if_2 = test_iface_name("sl2"); - let _wan_1 = ScopedDummyLink::new(&wan_if_1); - let _lan_1 = ScopedDummyLink::new(&lan_if_1); - let _wan_2 = ScopedDummyLink::new(&wan_if_2); - let _lan_2 = ScopedDummyLink::new(&lan_if_2); - - run_ip(&[ - "-6", - "addr", - "add", - "2001:db8:3000:ffff::1/64", - "dev", - &wan_if_1, - ]); - run_ip(&[ - "-6", - "route", - "add", - "default", - "from", - "2001:db8:3000::/56", - "dev", - &wan_if_1, - ]); - run_ip(&["-6", "route", "add", "2001:db8:3000::/56", "dev", &lan_if_1]); - - run_ip(&["-6", "addr", "add", "2001:db9:ffff::1/64", "dev", &wan_if_2]); - run_ip(&[ - "-6", - "route", - "add", - "default", - "from", - "2001:db9::/48", - "dev", - &wan_if_2, - ]); - run_ip(&["-6", "route", "add", "2001:db9::/48", "dev", &lan_if_2]); - - assert_eq!( - detected_prefix(detect_public_ipv6_prefix_linux().await.unwrap()), - Some("2001:db9::/48".parse().unwrap()) - ); - } } diff --git a/easytier/src/instance/public_ipv6_provider/linux.rs b/easytier/src/instance/public_ipv6_provider/linux.rs new file mode 100644 index 00000000..9d38cceb --- /dev/null +++ b/easytier/src/instance/public_ipv6_provider/linux.rs @@ -0,0 +1,1487 @@ +use std::path::Path; +use std::sync::Arc; + +use super::wait_for_public_ipv6_provider_reconcile_event; + +use crate::common::ifcfg::{ + RouteMessage, RouteType, add_ipv6_ndp_proxy, get_interface_index, list_ipv6_ndp_proxy, + list_ipv6_route_messages, remove_ipv6_ndp_proxy, +}; +use crate::common::netns::NetNS; +use crate::common::{ + error::Error, + global_ctx::{ArcGlobalCtx, GlobalCtxEvent}, +}; +use anyhow::Context; +use cidr::{Ipv6Cidr, Ipv6Inet}; +use easytier_core::config::peers::PublicIpv6ProviderConfig; +use easytier_core::peers::public_ipv6::is_global_routable_public_ipv6_prefix; +use easytier_core::peers::public_ipv6::provider::PublicIpv6NdpTarget; +use easytier_core::peers::public_ipv6::provider::{ + PublicIpv6NdpDesired, PublicIpv6PlatformError, PublicIpv6PlatformObservation, + PublicIpv6ProviderPlatform, +}; + +#[cfg(test)] +fn ensure_public_ipv6_provider_supported() -> Result<(), Error> { + PublicIpv6ProviderConfig { + provider_enabled: true, + configured_prefix: None, + provider_supported: cfg!(target_os = "linux"), + } + .validate() + .map_err(|error| anyhow::Error::new(error).into()) +} + +fn read_linux_proc_bool(path: &Path) -> Result { + let value = std::fs::read_to_string(path) + .with_context(|| format!("failed to read {}", path.display()))?; + match value.trim() { + "0" => Ok(false), + "1" => Ok(true), + other => Err(anyhow::anyhow!("unexpected value '{}' in {}", other, path.display()).into()), + } +} + +fn write_linux_proc_bool(path: &Path, enabled: bool) -> Result<(), Error> { + let value = if enabled { "1\n" } else { "0\n" }; + std::fs::write(path, value).with_context(|| format!("failed to write {}", path.display()))?; + Ok(()) +} + +fn ensure_linux_ipv6_forwarding_at_paths( + all_path: &Path, + default_path: &Path, +) -> Result { + let all_enabled = read_linux_proc_bool(all_path)?; + let default_enabled = read_linux_proc_bool(default_path)?; + let mut changed = false; + + if !all_enabled { + write_linux_proc_bool(all_path, true)?; + changed = true; + } + + if !default_enabled { + write_linux_proc_bool(default_path, true)?; + changed = true; + } + + if !read_linux_proc_bool(all_path)? || !read_linux_proc_bool(default_path)? { + return Err(anyhow::anyhow!( + "failed to enable Linux IPv6 forwarding in {} and {}", + all_path.display(), + default_path.display() + ) + .into()); + } + + Ok(changed) +} + +fn ensure_linux_ipv6_forwarding() -> Result { + let all_path = Path::new("/proc/sys/net/ipv6/conf/all/forwarding"); + let default_path = Path::new("/proc/sys/net/ipv6/conf/default/forwarding"); + + ensure_linux_ipv6_forwarding_at_paths(all_path, default_path).map_err(|err| { + anyhow::anyhow!( + "public IPv6 provider requires Linux IPv6 forwarding; failed to enable net.ipv6.conf.all.forwarding=1 and net.ipv6.conf.default.forwarding=1 automatically: {}. run with sufficient privileges or set them manually", + err + ) + .into() + }) +} + +#[derive(Clone, Debug, PartialEq, Eq)] +struct DetectedIpv6Route { + dst: Option, + src: Option, + ifindex: Option, + kind: RouteType, +} + +#[derive(Clone, Debug, PartialEq, Eq)] +struct DetectedPublicIpv6Prefix { + prefix: Ipv6Cidr, + ndp_target: Option, +} + +fn ipv6_cidr_from_route_addr(addr: &std::net::IpAddr, prefix_len: u8) -> Option { + match addr { + std::net::IpAddr::V6(addr) => Ipv6Cidr::new(*addr, prefix_len).ok(), + _ => None, + } +} + +impl TryFrom for DetectedIpv6Route { + type Error = Error; + + fn try_from(message: RouteMessage) -> Result { + let dst = message + .destination() + .and_then(|addr| ipv6_cidr_from_route_addr(addr, message.dst_len())); + let src = message + .source() + .and_then(|addr| ipv6_cidr_from_route_addr(addr, message.src_len())); + + Ok(Self { + dst, + src, + ifindex: message.oif(), + kind: message.route_type(), + }) + } +} + +fn is_ipv6_default_route(dst: Option) -> bool { + dst.is_none() || dst == Some(Ipv6Cidr::new(std::net::Ipv6Addr::UNSPECIFIED, 0).unwrap()) +} + +fn detect_public_ipv6_prefix_from_routes( + routes: &[DetectedIpv6Route], + loopback_ifindex: u32, +) -> Option { + routes + .iter() + .filter_map(|route| { + if !is_ipv6_default_route(route.dst) || route.kind != RouteType::Unicast { + return None; + } + + let prefix = route.src?; + let wan_ifindex = route.ifindex?; + if !is_global_routable_public_ipv6_prefix(prefix) { + return None; + } + + let delegated = routes.iter().any(|candidate| { + candidate.dst == Some(prefix) + && candidate.ifindex.is_some() + && candidate.ifindex != Some(wan_ifindex) + && candidate.ifindex != Some(loopback_ifindex) + && candidate.kind == RouteType::Unicast + }); + + delegated.then_some(DetectedPublicIpv6Prefix { + prefix, + ndp_target: None, + }) + }) + .min_by_key(|detected| detected.prefix.network_length()) +} + +#[derive(Clone, Debug, PartialEq, Eq)] +struct DetectedDefaultRouteIpv6Interface { + interface_name: String, + ifindex: u32, + address: std::net::Ipv6Addr, + prefix: Ipv6Cidr, +} + +#[derive(Clone, Debug, PartialEq, Eq)] +struct DefaultRouteIpv6InterfaceCandidate { + interface_name: String, + ifindex: u32, + address: std::net::Ipv6Addr, + prefix_len: u8, +} + +fn default_route_ifindices(routes: &[DetectedIpv6Route]) -> std::collections::BTreeSet { + routes + .iter() + .filter(|route| is_ipv6_default_route(route.dst) && route.kind == RouteType::Unicast) + .filter_map(|route| route.ifindex) + .collect() +} + +fn select_default_route_ipv6_interfaces( + candidates: impl IntoIterator, + wan_ifindices: &std::collections::BTreeSet, + max_prefix_len: u8, +) -> Vec { + candidates + .into_iter() + .filter_map(|candidate| { + if !wan_ifindices.contains(&candidate.ifindex) { + return None; + } + + if candidate.address.is_loopback() + || candidate.address.is_multicast() + || candidate.address.is_unicast_link_local() + || candidate.address.is_unique_local() + || candidate.address.is_unspecified() + { + return None; + } + + if candidate.prefix_len == 0 || candidate.prefix_len > max_prefix_len { + return None; + } + + let prefix = Ipv6Inet::new(candidate.address, candidate.prefix_len) + .ok() + .map(|inet| inet.network())?; + + Some(DetectedDefaultRouteIpv6Interface { + interface_name: candidate.interface_name, + ifindex: candidate.ifindex, + address: candidate.address, + prefix, + }) + }) + .collect() +} + +fn detect_default_route_ipv6_interfaces( + routes: &[DetectedIpv6Route], + max_prefix_len: u8, +) -> Vec { + use nix::ifaddrs::getifaddrs; + use nix::sys::socket::SockaddrLike; + use pnet::ipnetwork::ip_mask_to_prefix; + + let wan_ifindices = default_route_ifindices(routes); + if wan_ifindices.is_empty() { + return Vec::new(); + } + + let Ok(interfaces) = getifaddrs() else { + return Vec::new(); + }; + + let candidates = interfaces + .filter_map(|iface| { + let address = iface.address?; + let netmask = iface.netmask?; + let ifindex = get_interface_index(&iface.interface_name).ok()?; + + if address.family()? != nix::sys::socket::AddressFamily::Inet6 { + return None; + } + + let ipv6_addr = address.as_sockaddr_in6()?.ip(); + let netmask_ip = netmask.as_sockaddr_in6()?.ip(); + let prefix_len = ip_mask_to_prefix(std::net::IpAddr::V6(netmask_ip)).ok()?; + + Some(DefaultRouteIpv6InterfaceCandidate { + interface_name: iface.interface_name, + ifindex, + address: ipv6_addr, + prefix_len, + }) + }) + .collect::>(); + + select_default_route_ipv6_interfaces(candidates, &wan_ifindices, max_prefix_len) +} + +fn select_public_ipv6_prefix_from_default_route_interfaces( + candidates: impl IntoIterator, +) -> Option { + let iface = candidates + .into_iter() + .min_by_key(|iface| (iface.prefix.network_length(), iface.ifindex))?; + Some(DetectedPublicIpv6Prefix { + prefix: iface.prefix, + ndp_target: Some(PublicIpv6NdpTarget { + wan_interface: iface.interface_name, + }), + }) +} + +fn detect_public_ipv6_prefix_from_interfaces( + routes: &[DetectedIpv6Route], +) -> Option { + select_public_ipv6_prefix_from_default_route_interfaces(detect_default_route_ipv6_interfaces( + routes, 64, + )) +} + +fn ipv6_cidr_contains_cidr(outer: Ipv6Cidr, inner: Ipv6Cidr) -> bool { + outer.contains(&inner.first_address()) && outer.contains(&inner.last_address()) +} + +fn detect_configured_prefix_ndp_proxy_target( + routes: &[DetectedIpv6Route], + prefix: Ipv6Cidr, +) -> Option { + let wan_ifindices = default_route_ifindices(routes); + if wan_ifindices.is_empty() { + return None; + } + + let loopback_ifindex = get_interface_index("lo").ok(); + let routed = routes.iter().any(|route| { + route.dst == Some(prefix) + && route.kind == RouteType::Unicast + && route.ifindex.is_some_and(|ifindex| { + !wan_ifindices.contains(&ifindex) && Some(ifindex) != loopback_ifindex + }) + }); + if routed { + return None; + } + + detect_default_route_ipv6_interfaces(routes, 128) + .into_iter() + .filter(|iface| { + ipv6_cidr_contains_cidr(iface.prefix, prefix) + || (iface.prefix.network_length() == 128 && prefix.contains(&iface.address)) + }) + .min_by_key(|iface| (iface.prefix.network_length(), iface.ifindex)) + .map(|iface| PublicIpv6NdpTarget { + wan_interface: iface.interface_name, + }) +} + +fn list_detected_ipv6_routes() -> Result, Error> { + let routes = list_ipv6_route_messages().with_context(|| "failed to query linux ipv6 routes")?; + routes + .iter() + .cloned() + .map(DetectedIpv6Route::try_from) + .collect::, _>>() +} + +fn detect_public_ipv6_prefix_linux() -> Result, Error> { + let routes = list_detected_ipv6_routes()?; + let loopback_ifindex = + get_interface_index("lo").with_context(|| "failed to resolve linux loopback ifindex")?; + + if let Some(prefix) = detect_public_ipv6_prefix_from_routes(&routes, loopback_ifindex) { + return Ok(Some(prefix)); + } + + // Fallback for DHCPv6 IA_NA / SLAAC — see https://github.com/EasyTier/EasyTier/issues/2333 + Ok(detect_public_ipv6_prefix_from_interfaces(&routes)) +} + +#[derive(Default)] +struct NdpProxyRuntime { + wan_iface: Option, + applied: std::collections::BTreeSet, +} + +impl NdpProxyRuntime { + fn reconcile( + &mut self, + global_ctx: &ArcGlobalCtx, + desired: Option<&PublicIpv6NdpDesired>, + ) -> bool { + let Some(desired) = desired else { + return !self.clear_current(global_ctx); + }; + + let Some(tun_iface) = global_ctx.get_tun_device_name() else { + self.clear_current(global_ctx); + tracing::debug!("waiting for tun device before syncing NDP proxy entries"); + return self.cleanup_pending(); + }; + + let _g = global_ctx.net_ns.guard(); + + if self.wan_iface.as_deref() != Some(desired.target.wan_interface.as_str()) { + if !self.clear_current_locked() { + tracing::warn!( + old_wan_iface = ?self.wan_iface, + new_wan_iface = %desired.target.wan_interface, + remaining_entries = self.applied.len(), + "waiting to remove old NDP proxy entries before switching WAN interface" + ); + return true; + } + self.wan_iface = Some(desired.target.wan_interface.clone()); + } + + if let Err(err) = sync_ndp_proxy_entries( + desired.target.wan_interface.as_str(), + tun_iface.as_str(), + desired.prefix, + &mut self.applied, + ) { + tracing::warn!( + wan_iface = %desired.target.wan_interface, + tun_iface = %tun_iface, + ?err, + "failed to sync NDP proxy entries" + ); + } + self.cleanup_pending() + } + + fn clear_current(&mut self, global_ctx: &ArcGlobalCtx) -> bool { + self.clear_current_in_netns(&global_ctx.net_ns) + } + + fn clear_current_in_netns(&mut self, net_ns: &NetNS) -> bool { + let _g = net_ns.guard(); + self.clear_current_locked() + } + + fn clear_current_locked(&mut self) -> bool { + let Some(wan_iface) = self.wan_iface.clone() else { + return self.applied.is_empty(); + }; + + match list_ipv6_ndp_proxy(wan_iface.as_str()) { + Ok(current) => { + let candidates = self.applied.iter().copied().collect::>(); + clear_owned_ndp_proxy_entries( + wan_iface.as_str(), + ¤t, + &mut self.applied, + candidates, + ); + } + Err(err) if is_linux_missing_netlink_object_error(&err) => { + tracing::trace!( + wan_iface = %wan_iface, + ?err, + "forgetting NDP proxy ownership because WAN interface is gone" + ); + self.applied.clear(); + } + Err(err) => { + tracing::trace!( + wan_iface = %wan_iface, + ?err, + "failed to list NDP proxy entries before cleanup" + ); + } + } + + if self.applied.is_empty() { + self.wan_iface = None; + true + } else { + false + } + } + + fn cleanup_pending(&self) -> bool { + self.wan_iface.is_some() && !self.applied.is_empty() + } +} + +fn is_linux_missing_netlink_object_error(err: &Error) -> bool { + match err { + Error::IOError(err) => { + err.kind() == std::io::ErrorKind::NotFound + || matches!( + err.raw_os_error(), + Some(nix::libc::ESRCH | nix::libc::ENODEV | nix::libc::ENXIO) + ) + } + _ => false, + } +} + +fn clear_owned_ndp_proxy_entries( + wan_iface: &str, + current: &std::collections::BTreeSet, + applied: &mut std::collections::BTreeSet, + candidates: Vec, +) -> Option { + let mut first_err = None; + for addr in candidates { + if !current.contains(&addr) { + applied.remove(&addr); + continue; + } + + if let Err(err) = remove_ipv6_ndp_proxy(wan_iface, addr) { + if is_linux_missing_netlink_object_error(&err) { + applied.remove(&addr); + } else { + tracing::trace!( + wan_iface = %wan_iface, + addr = %addr, + ?err, + "failed to remove NDP proxy entry" + ); + first_err.get_or_insert(err); + } + } else { + applied.remove(&addr); + } + } + first_err +} + +fn ensure_linux_ndp_proxy_enabled(wan_iface: &str) -> Result<(), Error> { + let path = Path::new("/proc/sys/net/ipv6/conf") + .join(wan_iface) + .join("proxy_ndp"); + if !read_linux_proc_bool(&path)? { + write_linux_proc_bool(&path, true)?; + tracing::info!(wan_iface = %wan_iface, "enabled Linux NDP proxy"); + } + Ok(()) +} + +fn collect_public_ipv6_tun_routes( + tun_iface: &str, + prefix: Ipv6Cidr, +) -> Result, Error> { + let tun_ifindex = match get_interface_index(tun_iface) { + Ok(ifindex) => ifindex, + Err(err) if is_linux_missing_netlink_object_error(&err) => { + tracing::debug!( + tun_iface = %tun_iface, + ?err, + "treating missing tun interface as empty public IPv6 route set" + ); + return Ok(Default::default()); + } + Err(err) => return Err(err), + }; + Ok(list_ipv6_route_messages()? + .into_iter() + .filter(|route| route.dst_len() == 128 && route.route_type() == RouteType::Unicast) + .filter(|route| route.oif() == Some(tun_ifindex)) + .filter_map(|route| match route.destination() { + Some(std::net::IpAddr::V6(addr)) => Some(*addr), + _ => None, + }) + .filter(|addr| !addr.is_unicast_link_local() && prefix.contains(addr)) + .collect()) +} + +fn sync_ndp_proxy_entries( + wan_iface: &str, + tun_iface: &str, + prefix: Ipv6Cidr, + applied: &mut std::collections::BTreeSet, +) -> Result<(), Error> { + ensure_linux_ndp_proxy_enabled(wan_iface)?; + + let wanted = collect_public_ipv6_tun_routes(tun_iface, prefix)?; + let current = list_ipv6_ndp_proxy(wan_iface)?; + + let mut first_err = None; + for addr in wanted.difference(¤t) { + if let Err(err) = add_ipv6_ndp_proxy(wan_iface, *addr) { + first_err.get_or_insert(err); + } else { + applied.insert(*addr); + tracing::debug!(wan_iface = %wan_iface, addr = %addr, "added NDP proxy entry"); + } + } + + let stale = applied.difference(&wanted).copied().collect::>(); + let stale_cleanup_err = + clear_owned_ndp_proxy_entries(wan_iface, ¤t, applied, stale.clone()); + if !stale.is_empty() { + tracing::debug!( + wan_iface = %wan_iface, + stale_count = stale.len(), + remaining_count = stale.iter().filter(|addr| applied.contains(addr)).count(), + "synced stale NDP proxy entries" + ); + } + if let Some(err) = first_err.or(stale_cleanup_err) { + return Err(err); + } + + Ok(()) +} + +fn cleanup_ndp_proxy_runtime(runtime: &mut NdpProxyRuntime, net_ns: &NetNS) -> bool { + let complete = runtime.clear_current_in_netns(net_ns); + if !complete { + tracing::warn!( + remaining_entries = runtime.applied.len(), + wan_iface = ?runtime.wan_iface, + "failed to clean all NDP proxy entries before stopping public IPv6 provider task" + ); + } + complete +} + +pub(super) struct RuntimePublicIpv6ProviderPlatform { + global_ctx: std::sync::Weak, + net_ns: NetNS, + event_receiver: tokio::sync::Mutex>, + ndp_proxy: std::sync::Mutex, +} + +impl RuntimePublicIpv6ProviderPlatform { + pub(super) fn new(global_ctx: &ArcGlobalCtx) -> Arc { + Arc::new(Self { + global_ctx: Arc::downgrade(global_ctx), + net_ns: global_ctx.net_ns.clone(), + event_receiver: tokio::sync::Mutex::new(global_ctx.subscribe()), + ndp_proxy: std::sync::Mutex::new(NdpProxyRuntime::default()), + }) + } +} + +#[async_trait::async_trait] +impl PublicIpv6ProviderPlatform for RuntimePublicIpv6ProviderPlatform { + fn inspect( + &self, + config: PublicIpv6ProviderConfig, + ) -> Result { + let Some(global_ctx) = self.global_ctx.upgrade() else { + return Err(PublicIpv6PlatformError::Unavailable); + }; + + let _guard = global_ctx.net_ns.guard(); + ensure_linux_ipv6_forwarding() + .map_err(|error| PublicIpv6PlatformError::Failed(error.to_string()))?; + + if let Some(prefix) = config.configured_prefix { + let ndp_target = match list_detected_ipv6_routes() { + Ok(routes) => detect_configured_prefix_ndp_proxy_target(&routes, prefix), + Err(error) => { + tracing::warn!( + %prefix, + ?error, + "failed to detect NDP proxy target for configured public IPv6 prefix" + ); + None + } + }; + return Ok(PublicIpv6PlatformObservation { + detected_prefix: None, + ndp_target, + }); + } + + detect_public_ipv6_prefix_linux() + .map(|detected| PublicIpv6PlatformObservation { + detected_prefix: detected.as_ref().map(|detected| detected.prefix), + ndp_target: detected.and_then(|detected| detected.ndp_target), + }) + .map_err(|error| PublicIpv6PlatformError::Failed(error.to_string())) + } + + fn sync_ndp( + &self, + desired: Option, + ) -> Result<(), PublicIpv6PlatformError> { + let mut runtime = self.ndp_proxy.lock().unwrap(); + if let Some(global_ctx) = self.global_ctx.upgrade() { + let cleanup_pending = runtime.reconcile(&global_ctx, desired.as_ref()); + if desired.is_none() && cleanup_pending { + return Err(PublicIpv6PlatformError::Failed(format!( + "failed to clean all public IPv6 NDP proxy entries (remaining: {})", + runtime.applied.len() + ))); + } + return Ok(()); + } + if desired.is_none() { + return cleanup_ndp_proxy_runtime(&mut runtime, &self.net_ns) + .then_some(()) + .ok_or_else(|| { + PublicIpv6PlatformError::Failed(format!( + "failed to clean all public IPv6 NDP proxy entries (remaining: {})", + runtime.applied.len() + )) + }); + } + Err(PublicIpv6PlatformError::Unavailable) + } + + async fn wait_for_change(&self) -> bool { + let mut event_receiver = self.event_receiver.lock().await; + wait_for_public_ipv6_provider_reconcile_event(&mut event_receiver).await + } +} + +pub(crate) fn runtime_public_ipv6_provider_platform( + global_ctx: &ArcGlobalCtx, +) -> Arc { + RuntimePublicIpv6ProviderPlatform::new(global_ctx) +} + +#[cfg(test)] +mod tests { + use std::fs; + use std::path::PathBuf; + use std::process::Command; + use std::sync::Arc; + + use crate::common::ifcfg::RouteType; + + use super::{ + DefaultRouteIpv6InterfaceCandidate, DetectedIpv6Route, + detect_public_ipv6_prefix_from_interfaces, detect_public_ipv6_prefix_from_routes, + detect_public_ipv6_prefix_linux, ensure_linux_ipv6_forwarding_at_paths, + ensure_public_ipv6_provider_supported, select_default_route_ipv6_interfaces, + select_public_ipv6_prefix_from_default_route_interfaces, sync_ndp_proxy_entries, + }; + + use super::PublicIpv6ProviderConfig; + use crate::common::{ + config::{ConfigLoader, TomlConfigLoader}, + error::Error, + global_ctx::GlobalCtx, + }; + use crate::instance::config::test_core_instance_config; + + fn run_ip(args: &[&str]) { + let output = Command::new("ip") + .args(args) + .output() + .expect("failed to execute ip process"); + assert!( + output.status.success(), + "ip command failed: {:?}\nstdout: {}\nstderr: {}", + args, + String::from_utf8_lossy(&output.stdout), + String::from_utf8_lossy(&output.stderr), + ); + } + + fn test_iface_name(tag: &str) -> String { + format!("et{}{:x}", tag, std::process::id() & 0xffff) + } + + struct ScopedDummyLink { + name: String, + } + + impl ScopedDummyLink { + fn new(name: &str) -> Self { + let _ = Command::new("ip").args(["link", "del", name]).output(); + run_ip(&["link", "add", name, "type", "dummy"]); + run_ip(&["link", "set", name, "up"]); + Self { + name: name.to_string(), + } + } + } + + impl Drop for ScopedDummyLink { + fn drop(&mut self) { + let _ = Command::new("ip") + .args(["link", "del", &self.name]) + .output(); + } + } + + fn temp_forwarding_paths( + all_value: &str, + default_value: &str, + ) -> (tempfile::TempDir, PathBuf, PathBuf) { + let dir = tempfile::tempdir().unwrap(); + let all_path = dir.path().join("all_forwarding"); + let default_path = dir.path().join("default_forwarding"); + fs::write(&all_path, all_value).unwrap(); + fs::write(&default_path, default_value).unwrap(); + (dir, all_path, default_path) + } + + fn route( + dst: Option<&str>, + src: Option<&str>, + ifindex: Option, + kind: RouteType, + ) -> DetectedIpv6Route { + DetectedIpv6Route { + dst: dst.map(|cidr| cidr.parse().unwrap()), + src: src.map(|cidr| cidr.parse().unwrap()), + ifindex, + kind, + } + } + + fn detected_prefix( + detected: Option, + ) -> Option { + detected.map(|detected| detected.prefix) + } + + fn iface_candidate( + interface_name: &str, + ifindex: u32, + address: &str, + prefix_len: u8, + ) -> DefaultRouteIpv6InterfaceCandidate { + DefaultRouteIpv6InterfaceCandidate { + interface_name: interface_name.to_string(), + ifindex, + address: address.parse().unwrap(), + prefix_len, + } + } + + #[test] + fn test_detect_public_ipv6_prefix_from_routes_selects_delegated_prefix() { + let routes = vec![ + route(None, Some("2001:db8:1::/56"), Some(2), RouteType::Unicast), + route(Some("2001:db8:1::/56"), None, Some(3), RouteType::Unicast), + ]; + + assert_eq!( + detected_prefix(detect_public_ipv6_prefix_from_routes(&routes, 1)), + Some("2001:db8:1::/56".parse().unwrap()) + ); + } + + #[test] + fn test_detect_public_ipv6_prefix_from_routes_rejects_non_public_prefixes() { + let routes = vec![ + route(Some("::/0"), Some("fd00::/48"), Some(2), RouteType::Unicast), + route(Some("fd00::/48"), None, Some(3), RouteType::Unicast), + route(None, Some("fe80::/64"), Some(4), RouteType::Unicast), + route(Some("fe80::/64"), None, Some(5), RouteType::Unicast), + route(None, Some("ff00::/8"), Some(6), RouteType::Unicast), + route(Some("ff00::/8"), None, Some(7), RouteType::Unicast), + route(None, Some("::/0"), Some(8), RouteType::Unicast), + route(Some("::/0"), None, Some(9), RouteType::Unicast), + ]; + + assert_eq!( + detected_prefix(detect_public_ipv6_prefix_from_routes(&routes, 1)), + None + ); + } + + #[test] + fn test_detect_public_ipv6_prefix_from_routes_requires_delegated_route() { + let routes = vec![route( + None, + Some("2001:db8:1::/56"), + Some(2), + RouteType::Unicast, + )]; + + assert_eq!( + detected_prefix(detect_public_ipv6_prefix_from_routes(&routes, 1)), + None + ); + } + + #[test] + fn test_detect_public_ipv6_prefix_from_routes_rejects_non_unicast_default_route() { + let routes = vec![ + route(None, Some("2001:db8:1::/56"), Some(2), RouteType::Blackhole), + route(Some("2001:db8:1::/56"), None, Some(3), RouteType::Unicast), + ]; + + assert_eq!( + detected_prefix(detect_public_ipv6_prefix_from_routes(&routes, 1)), + None + ); + } + + #[test] + fn test_detect_public_ipv6_prefix_from_routes_rejects_loopback_delegation() { + let routes = vec![ + route(None, Some("2001:db8:1::/56"), Some(2), RouteType::Unicast), + route(Some("2001:db8:1::/56"), None, Some(1), RouteType::Unicast), + ]; + + assert_eq!( + detected_prefix(detect_public_ipv6_prefix_from_routes(&routes, 1)), + None + ); + } + + #[test] + fn test_detect_public_ipv6_prefix_from_routes_prefers_shortest_prefix() { + let routes = vec![ + route(None, Some("2001:db8:1::/56"), Some(2), RouteType::Unicast), + route(Some("2001:db8:1::/56"), None, Some(3), RouteType::Unicast), + route(None, Some("2001:db8::/48"), Some(4), RouteType::Unicast), + route(Some("2001:db8::/48"), None, Some(5), RouteType::Unicast), + ]; + + assert_eq!( + detected_prefix(detect_public_ipv6_prefix_from_routes(&routes, 1)), + Some("2001:db8::/48".parse().unwrap()) + ); + } + + #[test] + fn test_detect_public_ipv6_prefix_from_routes_rejects_non_unicast_delegation() { + let routes = vec![ + route(None, Some("2001:db8:1::/56"), Some(2), RouteType::Unicast), + route(Some("2001:db8:1::/56"), None, Some(3), RouteType::Blackhole), + ]; + + assert_eq!( + detected_prefix(detect_public_ipv6_prefix_from_routes(&routes, 1)), + None + ); + } + + fn test_global_ctx() -> Arc { + Arc::new(GlobalCtx::new(TomlConfigLoader::default())) + } + + #[tokio::test] + async fn test_runtime_public_ipv6_provider_config_reads_provider_fields() { + let global_ctx = test_global_ctx(); + let prefix = "2001:db8::/48".parse().unwrap(); + global_ctx.config.set_ipv6_public_addr_provider(true); + global_ctx.config.set_ipv6_public_addr_prefix(Some(prefix)); + + assert_eq!( + test_core_instance_config(&global_ctx) + .connectivity + .runtime + .public_ipv6_provider, + PublicIpv6ProviderConfig { + provider_enabled: true, + configured_prefix: Some(prefix), + provider_supported: cfg!(target_os = "linux"), + } + ); + } + + #[test] + fn test_public_ipv6_provider_platform_check_accepts_linux() { + assert!(ensure_public_ipv6_provider_supported().is_ok()); + } + + #[test] + fn test_ensure_linux_ipv6_forwarding_enables_all_and_default() { + let (_dir, all_path, default_path) = temp_forwarding_paths("0\n", "0\n"); + + let changed = ensure_linux_ipv6_forwarding_at_paths(&all_path, &default_path).unwrap(); + + assert!(changed); + assert_eq!(fs::read_to_string(&all_path).unwrap(), "1\n"); + assert_eq!(fs::read_to_string(&default_path).unwrap(), "1\n"); + } + + #[test] + fn test_ensure_linux_ipv6_forwarding_is_noop_when_already_enabled() { + let (_dir, all_path, default_path) = temp_forwarding_paths("1\n", "1\n"); + + let changed = ensure_linux_ipv6_forwarding_at_paths(&all_path, &default_path).unwrap(); + + assert!(!changed); + assert_eq!(fs::read_to_string(&all_path).unwrap(), "1\n"); + assert_eq!(fs::read_to_string(&default_path).unwrap(), "1\n"); + } + + #[serial_test::serial] + #[tokio::test] + async fn test_detect_public_ipv6_prefix_linux_reads_netlink_routes_from_kernel() { + let wan_if = test_iface_name("dw"); + let lan_if = test_iface_name("dl"); + let _wan = ScopedDummyLink::new(&wan_if); + let _lan = ScopedDummyLink::new(&lan_if); + + run_ip(&[ + "-6", + "addr", + "add", + "2001:db8:100:ffff::1/64", + "dev", + &wan_if, + ]); + run_ip(&[ + "-6", + "route", + "add", + "default", + "from", + "2001:db8:100::/56", + "dev", + &wan_if, + ]); + run_ip(&["-6", "route", "add", "2001:db8:100::/56", "dev", &lan_if]); + + assert_eq!( + detected_prefix(detect_public_ipv6_prefix_linux().unwrap()), + Some("2001:db8:100::/56".parse().unwrap()) + ); + } + + #[serial_test::serial] + #[tokio::test] + async fn test_detect_public_ipv6_prefix_linux_dhcpv6_ia_na_fallback() { + // DHCPv6 IA_NA scenario: prefix is directly on the WAN interface, + // with no delegated route on a LAN interface. + // The route-based detection should fail, and the interface-scanning + // fallback should pick up the prefix from the WAN address. + let wan_if = test_iface_name("ia"); + let _wan = ScopedDummyLink::new(&wan_if); + + run_ip(&[ + "-6", + "addr", + "add", + "2001:db8:aaaa:ffff::1/64", + "dev", + &wan_if, + ]); + run_ip(&[ + "-6", + "route", + "add", + "default", + "from", + "2001:db8:aaaa::/64", + "dev", + &wan_if, + ]); + // Also add a /48 address+route pair to verify shortest-prefix preference + run_ip(&["-6", "addr", "add", "2001:db8:bbbb::1/48", "dev", &wan_if]); + run_ip(&[ + "-6", + "route", + "add", + "default", + "from", + "2001:db8::/48", + "dev", + &wan_if, + ]); + + // NO delegated route on a LAN interface — this is the IA_NA case + // The fallback should find both prefixes via interface scanning and + // prefer the shorter /48. + let detected = detect_public_ipv6_prefix_linux().unwrap().unwrap(); + assert_eq!(detected.prefix, "2001:db8:bbbb::/48".parse().unwrap()); + assert_eq!(detected.ndp_target.unwrap().wan_interface, wan_if); + } + + #[serial_test::serial] + #[tokio::test] + async fn test_detect_public_ipv6_prefix_from_interfaces_uses_default_route_iface() { + let wan_if = test_iface_name("dw"); + let other_if = test_iface_name("do"); + let _wan = ScopedDummyLink::new(&wan_if); + let _other = ScopedDummyLink::new(&other_if); + + run_ip(&["-6", "addr", "add", "2001:db8:dddd::1/64", "dev", &wan_if]); + run_ip(&["-6", "addr", "add", "2001:db8::1/48", "dev", &other_if]); + + let wan_ifindex = crate::common::ifcfg::get_interface_index(&wan_if).unwrap(); + let other_ifindex = crate::common::ifcfg::get_interface_index(&other_if).unwrap(); + let routes = vec![ + route(None, None, Some(wan_ifindex), RouteType::Unicast), + route( + Some("2001:db8::/48"), + None, + Some(other_ifindex), + RouteType::Unicast, + ), + ]; + + let detected = detect_public_ipv6_prefix_from_interfaces(&routes) + .expect("fallback should select the default-route interface"); + assert_eq!(detected.prefix, "2001:db8:dddd::/64".parse().unwrap()); + assert_eq!(detected.ndp_target.unwrap().wan_interface, wan_if); + } + + #[test] + fn test_select_default_route_ipv6_interfaces_filters_candidates() { + let wan_ifindices = [2u32].into_iter().collect(); + let candidates = vec![ + iface_candidate("wan0", 2, "2001:db8:100::1", 64), + iface_candidate("nonwan0", 9, "2001:db8:200::1", 64), + iface_candidate("loopback0", 2, "::1", 128), + iface_candidate("linklocal0", 2, "fe80::1", 64), + iface_candidate("ula0", 2, "fd00::1", 64), + iface_candidate("multicast0", 2, "ff02::1", 64), + iface_candidate("unspecified0", 2, "::", 64), + iface_candidate("empty0", 2, "2001:db8:300::1", 0), + iface_candidate("host0", 2, "2001:db8:400::1", 128), + ]; + + let selected = select_default_route_ipv6_interfaces(candidates, &wan_ifindices, 64); + + assert_eq!(selected.len(), 1); + assert_eq!(selected[0].interface_name, "wan0"); + assert_eq!(selected[0].prefix, "2001:db8:100::/64".parse().unwrap()); + } + + #[test] + fn test_select_default_route_ipv6_interfaces_strips_host_bits() { + let wan_ifindices = [2u32].into_iter().collect(); + let candidates = vec![iface_candidate("wan0", 2, "2001:db8:aaaa::abcd", 64)]; + + let selected = select_default_route_ipv6_interfaces(candidates, &wan_ifindices, 64); + + assert_eq!(selected.len(), 1); + assert_eq!(selected[0].prefix, "2001:db8:aaaa::/64".parse().unwrap()); + } + + #[test] + fn test_select_public_ipv6_prefix_tie_breaks_by_lowest_ifindex() { + let wan_ifindices = [2u32, 3, 5].into_iter().collect(); + let candidates = vec![ + iface_candidate("wan5", 5, "2001:db8:5555::1", 48), + iface_candidate("wan64", 3, "2001:db8:3333::1", 64), + iface_candidate("wan2", 2, "2001:db9:2222::1", 48), + ]; + let interfaces = select_default_route_ipv6_interfaces(candidates, &wan_ifindices, 64); + + let detected = select_public_ipv6_prefix_from_default_route_interfaces(interfaces) + .expect("default-route public IPv6 prefix should be selected"); + + assert_eq!(detected.prefix, "2001:db9:2222::/48".parse().unwrap()); + assert_eq!(detected.ndp_target.unwrap().wan_interface, "wan2"); + } + + #[serial_test::serial] + #[tokio::test] + async fn test_configured_prefix_on_default_iface_gets_ndp_proxy_target() { + let wan_if = test_iface_name("cp"); + let _wan = ScopedDummyLink::new(&wan_if); + let configured_prefix = "2001:db8:feed::/64".parse().unwrap(); + + run_ip(&["-6", "addr", "add", "2001:db8:feed::1/128", "dev", &wan_if]); + + let ifindex = crate::common::ifcfg::get_interface_index(&wan_if).unwrap(); + let routes = vec![route(None, None, Some(ifindex), RouteType::Unicast)]; + + let target = super::detect_configured_prefix_ndp_proxy_target(&routes, configured_prefix) + .expect("configured on-link prefix should require NDP proxy"); + assert_eq!(target.wan_interface, wan_if); + } + + #[serial_test::serial] + #[tokio::test] + async fn test_configured_prefix_broader_than_default_iface_does_not_get_ndp_proxy_target() { + let wan_if = test_iface_name("cb"); + let _wan = ScopedDummyLink::new(&wan_if); + let configured_prefix = "2001:db8:beef::/48".parse().unwrap(); + + run_ip(&["-6", "addr", "add", "2001:db8:beef:1::1/64", "dev", &wan_if]); + + let ifindex = crate::common::ifcfg::get_interface_index(&wan_if).unwrap(); + let routes = vec![route(None, None, Some(ifindex), RouteType::Unicast)]; + + assert_eq!( + super::detect_configured_prefix_ndp_proxy_target(&routes, configured_prefix), + None + ); + } + + #[serial_test::serial] + #[tokio::test] + async fn test_detect_public_ipv6_prefix_linux_dhcpv6_ia_na_single_prefix() { + // DHCPv6 IA_NA: the WAN interface has a global prefix with a + // default route. Use a /48 dummy prefix so the fallback prefers it + // over any real /64 on the test machine. + let wan_if = test_iface_name("ib"); + let _wan = ScopedDummyLink::new(&wan_if); + + run_ip(&["-6", "addr", "add", "2001:db8:cccc::1/48", "dev", &wan_if]); + run_ip(&["-6", "route", "add", "default", "dev", &wan_if]); + + let detected = detect_public_ipv6_prefix_linux().unwrap().unwrap(); + assert_eq!(detected.prefix, "2001:db8:cccc::/48".parse().unwrap()); + assert_eq!(detected.ndp_target.unwrap().wan_interface, wan_if); + } + + #[serial_test::serial] + #[tokio::test] + async fn test_detect_public_ipv6_prefix_from_interfaces_skips_non_global() { + // Create a dummy interface with only a link-local address. + // The interface fallback should return None because there is no + // global unicast address. + let iface = test_iface_name("ng"); + let _link = ScopedDummyLink::new(&iface); + + // Bring up the interface so it auto-configures a link-local address + run_ip(&["link", "set", &iface, "up"]); + let ifindex = crate::common::ifcfg::get_interface_index(&iface).unwrap(); + let routes = vec![route(None, None, Some(ifindex), RouteType::Unicast)]; + + // No global address added — only link-local should be present + let result = detect_public_ipv6_prefix_from_interfaces(&routes); + assert_eq!(detected_prefix(result), None); + } + + #[serial_test::serial] + #[tokio::test] + async fn test_ndp_proxy_sync_uses_configured_tun_iface_without_shell_neigh() { + let wan_if = test_iface_name("nw"); + let tun_if = test_iface_name("nt"); + let _wan = ScopedDummyLink::new(&wan_if); + let _tun = ScopedDummyLink::new(&tun_if); + let addr = "2001:db8:abcd::123".parse::().unwrap(); + let prefix = "2001:db8:abcd::/64".parse().unwrap(); + let mut applied = std::collections::BTreeSet::new(); + + run_ip(&["-6", "route", "add", &format!("{addr}/128"), "dev", &tun_if]); + + sync_ndp_proxy_entries(&wan_if, &tun_if, prefix, &mut applied).unwrap(); + assert!( + crate::common::ifcfg::list_ipv6_ndp_proxy(&wan_if) + .unwrap() + .contains(&addr) + ); + assert!(applied.contains(&addr)); + + run_ip(&["-6", "route", "del", &format!("{addr}/128"), "dev", &tun_if]); + sync_ndp_proxy_entries(&wan_if, &tun_if, prefix, &mut applied).unwrap(); + assert!( + !crate::common::ifcfg::list_ipv6_ndp_proxy(&wan_if) + .unwrap() + .contains(&addr) + ); + assert!(!applied.contains(&addr)); + } + + #[serial_test::serial] + #[tokio::test] + async fn test_ndp_proxy_sync_does_not_delete_preexisting_proxy_entry() { + let wan_if = test_iface_name("pw"); + let tun_if = test_iface_name("pt"); + let _wan = ScopedDummyLink::new(&wan_if); + let _tun = ScopedDummyLink::new(&tun_if); + let addr = "2001:db8:beef::123".parse::().unwrap(); + let prefix = "2001:db8:beef::/64".parse().unwrap(); + let mut applied = std::collections::BTreeSet::new(); + + super::ensure_linux_ndp_proxy_enabled(&wan_if).unwrap(); + crate::common::ifcfg::add_ipv6_ndp_proxy(&wan_if, addr).unwrap(); + run_ip(&["-6", "route", "add", &format!("{addr}/128"), "dev", &tun_if]); + + sync_ndp_proxy_entries(&wan_if, &tun_if, prefix, &mut applied).unwrap(); + assert!( + crate::common::ifcfg::list_ipv6_ndp_proxy(&wan_if) + .unwrap() + .contains(&addr) + ); + assert!(!applied.contains(&addr)); + + run_ip(&["-6", "route", "del", &format!("{addr}/128"), "dev", &tun_if]); + sync_ndp_proxy_entries(&wan_if, &tun_if, prefix, &mut applied).unwrap(); + assert!( + crate::common::ifcfg::list_ipv6_ndp_proxy(&wan_if) + .unwrap() + .contains(&addr) + ); + assert!(!applied.contains(&addr)); + + crate::common::ifcfg::remove_ipv6_ndp_proxy(&wan_if, addr).unwrap(); + } + + #[serial_test::serial] + #[tokio::test] + async fn test_ndp_proxy_sync_removes_owned_entry_when_tun_iface_is_gone() { + let wan_if = test_iface_name("gw"); + let _wan = ScopedDummyLink::new(&wan_if); + let addr = "2001:db8:face::123".parse::().unwrap(); + let prefix = "2001:db8:face::/64".parse().unwrap(); + let mut applied = std::collections::BTreeSet::from([addr]); + + super::ensure_linux_ndp_proxy_enabled(&wan_if).unwrap(); + crate::common::ifcfg::add_ipv6_ndp_proxy(&wan_if, addr).unwrap(); + + sync_ndp_proxy_entries(&wan_if, "missing-easytier-tun", prefix, &mut applied).unwrap(); + assert!( + !crate::common::ifcfg::list_ipv6_ndp_proxy(&wan_if) + .unwrap() + .contains(&addr) + ); + assert!(applied.is_empty()); + } + + #[serial_test::serial] + #[tokio::test] + async fn test_cleanup_ndp_proxy_runtime_removes_owned_entry_on_task_exit() { + let wan_if = test_iface_name("cw"); + let _wan = ScopedDummyLink::new(&wan_if); + let addr = "2001:db8:cafe::123".parse::().unwrap(); + let mut runtime = super::NdpProxyRuntime { + wan_iface: Some(wan_if.clone()), + applied: std::collections::BTreeSet::from([addr]), + }; + + super::ensure_linux_ndp_proxy_enabled(&wan_if).unwrap(); + crate::common::ifcfg::add_ipv6_ndp_proxy(&wan_if, addr).unwrap(); + + assert!(super::cleanup_ndp_proxy_runtime( + &mut runtime, + &crate::common::netns::NetNS::new(None) + )); + assert!( + !crate::common::ifcfg::list_ipv6_ndp_proxy(&wan_if) + .unwrap() + .contains(&addr) + ); + assert!(!runtime.cleanup_pending()); + } + + async fn wait_for_ndp_proxy_entry(wan_if: &str, addr: std::net::Ipv6Addr, present: bool) { + for _ in 0..50 { + let current = crate::common::ifcfg::list_ipv6_ndp_proxy(wan_if).unwrap(); + if current.contains(&addr) == present { + return; + } + tokio::time::sleep(std::time::Duration::from_millis(20)).await; + } + + let current = crate::common::ifcfg::list_ipv6_ndp_proxy(wan_if).unwrap(); + assert_eq!(current.contains(&addr), present); + } + + #[serial_test::serial] + #[tokio::test] + async fn test_reconcile_task_shutdown_removes_owned_ndp_proxy_entry() { + let wan_if = test_iface_name("tw"); + let tun_if = test_iface_name("tt"); + let _wan = ScopedDummyLink::new(&wan_if); + let _tun = ScopedDummyLink::new(&tun_if); + let prefix = "2001:db8:fade::/64".parse().unwrap(); + let wan_addr = "2001:db8:fade::1"; + let leased_addr = "2001:db8:fade::123".parse::().unwrap(); + let global_ctx = test_global_ctx(); + + run_ip(&[ + "-6", + "addr", + "add", + &format!("{wan_addr}/128"), + "dev", + &wan_if, + ]); + run_ip(&["-6", "route", "add", "default", "dev", &wan_if]); + run_ip(&[ + "-6", + "route", + "add", + &format!("{leased_addr}/128"), + "dev", + &tun_if, + ]); + + global_ctx.config.set_ipv6_public_addr_provider(true); + global_ctx.config.set_ipv6_public_addr_prefix(Some(prefix)); + global_ctx.set_tun_device_ready(tun_if); + + let platform = super::RuntimePublicIpv6ProviderPlatform::new(&global_ctx); + let runtime_config = easytier_core::config::runtime::CoreRuntimeConfigStore::new( + easytier_core::config::runtime::CoreRuntimeConfig { + public_ipv6_provider: test_core_instance_config(&global_ctx) + .connectivity + .runtime + .public_ipv6_provider, + ..Default::default() + }, + Arc::new(easytier_core::config::peers::PeerRuntimeSnapshot::default()), + ); + let runtime = easytier_core::peers::public_ipv6::CorePublicIpv6Runtime::new( + runtime_config.clone(), + global_ctx.clone(), + global_ctx.clone(), + ); + let service = easytier_core::peers::public_ipv6::provider::PublicIpv6ProviderService::new( + platform, + runtime_config, + runtime, + ); + service.start().await; + wait_for_ndp_proxy_entry(&wan_if, leased_addr, true).await; + + service.stop().await; + wait_for_ndp_proxy_entry(&wan_if, leased_addr, false).await; + } + + #[test] + fn test_missing_netlink_object_errors_release_ndp_ownership() { + for errno in [ + nix::libc::ENOENT, + nix::libc::ESRCH, + nix::libc::ENODEV, + nix::libc::ENXIO, + ] { + let err = Error::IOError(std::io::Error::from_raw_os_error(errno)); + assert!(super::is_linux_missing_netlink_object_error(&err)); + } + } + + #[test] + fn test_clear_owned_ndp_proxy_entries_forgets_already_absent_entry() { + let addr = "2001:db8:dead::111".parse::().unwrap(); + let mut applied = std::collections::BTreeSet::from([addr]); + let current = std::collections::BTreeSet::new(); + + assert!( + super::clear_owned_ndp_proxy_entries("missing", ¤t, &mut applied, vec![addr],) + .is_none() + ); + assert!(!applied.contains(&addr)); + } + + #[test] + fn test_ndp_proxy_runtime_forgets_entries_when_wan_interface_is_gone() { + let addr = "2001:db8:dead::123".parse::().unwrap(); + let mut runtime = super::NdpProxyRuntime { + wan_iface: Some(test_iface_name("missing")), + applied: std::collections::BTreeSet::from([addr]), + }; + + assert!(runtime.clear_current_locked()); + assert!(runtime.wan_iface.is_none()); + assert!(runtime.applied.is_empty()); + } + + #[tokio::test] + async fn test_ndp_proxy_runtime_finishes_cleanup_when_wan_interface_is_gone_after_disable() { + let addr = "2001:db8:dead::456".parse::().unwrap(); + let global_ctx = test_global_ctx(); + let mut runtime = super::NdpProxyRuntime { + wan_iface: Some(test_iface_name("missing")), + applied: std::collections::BTreeSet::from([addr]), + }; + + assert!(!runtime.reconcile(&global_ctx, None)); + assert!(!runtime.cleanup_pending()); + } + + #[serial_test::serial] + #[tokio::test] + async fn test_detect_public_ipv6_prefix_linux_prefers_shortest_prefix_from_kernel() { + let wan_if_1 = test_iface_name("sw1"); + let lan_if_1 = test_iface_name("sl1"); + let wan_if_2 = test_iface_name("sw2"); + let lan_if_2 = test_iface_name("sl2"); + let _wan_1 = ScopedDummyLink::new(&wan_if_1); + let _lan_1 = ScopedDummyLink::new(&lan_if_1); + let _wan_2 = ScopedDummyLink::new(&wan_if_2); + let _lan_2 = ScopedDummyLink::new(&lan_if_2); + + run_ip(&[ + "-6", + "addr", + "add", + "2001:db8:3000:ffff::1/64", + "dev", + &wan_if_1, + ]); + run_ip(&[ + "-6", + "route", + "add", + "default", + "from", + "2001:db8:3000::/56", + "dev", + &wan_if_1, + ]); + run_ip(&["-6", "route", "add", "2001:db8:3000::/56", "dev", &lan_if_1]); + + run_ip(&["-6", "addr", "add", "2001:db9:ffff::1/64", "dev", &wan_if_2]); + run_ip(&[ + "-6", + "route", + "add", + "default", + "from", + "2001:db9::/48", + "dev", + &wan_if_2, + ]); + run_ip(&["-6", "route", "add", "2001:db9::/48", "dev", &lan_if_2]); + + assert_eq!( + detected_prefix(detect_public_ipv6_prefix_linux().unwrap()), + Some("2001:db9::/48".parse().unwrap()) + ); + } +} diff --git a/easytier/src/instance/public_ipv6_provider/unsupported.rs b/easytier/src/instance/public_ipv6_provider/unsupported.rs new file mode 100644 index 00000000..c0513cb3 --- /dev/null +++ b/easytier/src/instance/public_ipv6_provider/unsupported.rs @@ -0,0 +1,79 @@ +use std::sync::Arc; + +use easytier_core::{ + config::peers::PublicIpv6ProviderConfig, + peers::public_ipv6::provider::{ + PublicIpv6NdpDesired, PublicIpv6PlatformError, PublicIpv6PlatformObservation, + PublicIpv6ProviderPlatform, + }, +}; + +use super::wait_for_public_ipv6_provider_reconcile_event; +use crate::common::global_ctx::{ArcGlobalCtx, GlobalCtxEvent}; + +pub(super) struct RuntimePublicIpv6ProviderPlatform { + global_ctx: std::sync::Weak, + event_receiver: tokio::sync::Mutex>, +} + +impl RuntimePublicIpv6ProviderPlatform { + pub(super) fn new(global_ctx: &ArcGlobalCtx) -> Arc { + Arc::new(Self { + global_ctx: Arc::downgrade(global_ctx), + event_receiver: tokio::sync::Mutex::new(global_ctx.subscribe()), + }) + } +} + +#[async_trait::async_trait] +impl PublicIpv6ProviderPlatform for RuntimePublicIpv6ProviderPlatform { + fn inspect( + &self, + config: PublicIpv6ProviderConfig, + ) -> Result { + let Some(global_ctx) = self.global_ctx.upgrade() else { + return Err(PublicIpv6PlatformError::Unavailable); + }; + let _ = (global_ctx, config); + Ok(PublicIpv6PlatformObservation::default()) + } + + fn sync_ndp( + &self, + desired: Option, + ) -> Result<(), PublicIpv6PlatformError> { + let _ = desired; + Ok(()) + } + + async fn wait_for_change(&self) -> bool { + let mut event_receiver = self.event_receiver.lock().await; + wait_for_public_ipv6_provider_reconcile_event(&mut event_receiver).await + } +} + +pub(crate) fn runtime_public_ipv6_provider_platform( + global_ctx: &ArcGlobalCtx, +) -> Arc { + RuntimePublicIpv6ProviderPlatform::new(global_ctx) +} + +#[cfg(test)] +mod tests { + use easytier_core::config::peers::PublicIpv6ProviderConfig; + + #[test] + fn public_ipv6_provider_platform_check_reports_linux_only() { + let err = PublicIpv6ProviderConfig { + provider_enabled: true, + configured_prefix: None, + provider_supported: false, + } + .validate() + .unwrap_err(); + let msg = err.to_string(); + + assert!(msg.contains("Linux"), "{msg}"); + assert!(msg.contains("ipv6-public-addr-auto"), "{msg}"); + } +} diff --git a/easytier/src/instance/runtime_host.rs b/easytier/src/instance/runtime_host.rs new file mode 100644 index 00000000..c1f38295 --- /dev/null +++ b/easytier/src/instance/runtime_host.rs @@ -0,0 +1,114 @@ +use std::sync::Arc; + +use easytier_core::{gateway::dhcp::DhcpIpv4Host, instance::CorePacketPlane}; +use tokio::sync::{Mutex, mpsc}; +use tokio_util::sync::CancellationToken; + +use crate::common::global_ctx::ArcGlobalCtx; + +mod event_journal; +mod implementation; +#[cfg(feature = "tun")] +mod magic_dns; +#[cfg(feature = "tun")] +mod tun_common; +#[cfg(not(feature = "tun"))] +#[path = "runtime_host/tun_disabled.rs"] +mod tun_runtime; +#[cfg(all(feature = "tun", not(mobile)))] +#[path = "runtime_host/tun_desktop.rs"] +mod tun_runtime; +#[cfg(all(feature = "tun", mobile))] +#[path = "runtime_host/tun_mobile.rs"] +mod tun_runtime; + +use event_journal::EventJournal; +#[cfg(feature = "tun")] +use magic_dns::MagicDnsRuntime; +use tun_runtime::NativeTunRuntime; + +pub(super) type HostPacketReceiver = mpsc::Receiver>; + +pub(crate) struct NativeInstanceRuntimeHost { + global_ctx: ArcGlobalCtx, + operation: Arc>, + cancel: CancellationToken, + event_journal: EventJournal, + tun: NativeTunRuntime, +} + +impl NativeInstanceRuntimeHost { + pub(crate) fn new( + global_ctx: ArcGlobalCtx, + peer_packet_receiver: HostPacketReceiver, + ) -> Arc { + let cancel = CancellationToken::new(); + let tun = NativeTunRuntime::new(global_ctx.clone(), cancel.clone(), peer_packet_receiver); + let event_journal = EventJournal::new(&global_ctx); + Arc::new(Self { + global_ctx, + event_journal, + operation: Arc::new(Mutex::new(())), + cancel, + tun, + }) + } + + async fn prepare_runtime( + &self, + packet_plane: Arc, + ) -> anyhow::Result>> { + self.event_journal.start(self.cancel.clone()).await; + self.tun.prepare(packet_plane.clone()).await?; + Ok(Some( + self.tun.dhcp_host(self.operation.clone(), packet_plane), + )) + } + + async fn shutdown_runtime(&self) { + self.cancel.cancel(); + let _operation = self.operation.lock().await; + self.event_journal.stop().await; + self.tun.shutdown().await; + } + + fn request_runtime_shutdown(&self) { + self.cancel.cancel(); + } + + fn management_events_snapshot(&self) -> Vec { + self.event_journal.events() + } + + pub(crate) fn subscribe_event(&self) -> crate::common::global_ctx::EventBusSubscriber { + self.global_ctx.subscribe() + } + + fn attach_runtime_tun_fd(&self, fd: i32) -> anyhow::Result<()> { + self.tun.attach_fd(fd) + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::common::{ + config::TomlConfig, + global_ctx::{GlobalCtx, GlobalCtxEvent}, + }; + + #[test] + fn runtime_host_owns_event_subscription_context() { + let global_ctx = Arc::new(GlobalCtx::new(TomlConfig::default())); + let (_packet_sender, packet_receiver) = mpsc::channel(1); + let runtime_host = NativeInstanceRuntimeHost::new(global_ctx.clone(), packet_receiver); + let mut events = runtime_host.subscribe_event(); + + global_ctx.issue_event(GlobalCtxEvent::CredentialChanged); + + assert_eq!( + events.try_recv().unwrap(), + GlobalCtxEvent::CredentialChanged + ); + } +} diff --git a/easytier/src/instance/runtime_host/event_journal.rs b/easytier/src/instance/runtime_host/event_journal.rs new file mode 100644 index 00000000..5071fec0 --- /dev/null +++ b/easytier/src/instance/runtime_host/event_journal.rs @@ -0,0 +1,132 @@ +#[cfg(feature = "management")] +use std::{ + collections::VecDeque, + sync::{Arc, RwLock}, +}; + +#[cfg(feature = "management")] +use tokio::sync::Mutex; +use tokio_util::sync::CancellationToken; +#[cfg(feature = "management")] +use tokio_util::task::AbortOnDropHandle; + +use crate::common::global_ctx::ArcGlobalCtx; +#[cfg(feature = "management")] +use crate::common::global_ctx::{EventBusSubscriber, GlobalCtxEvent}; + +#[cfg(feature = "management")] +#[derive(serde::Serialize)] +struct ManagementEvent { + time: chrono::DateTime, + event: GlobalCtxEvent, +} + +#[cfg(feature = "management")] +pub(super) struct EventJournal { + global_ctx: ArcGlobalCtx, + events: Arc>>, + receiver: Mutex>, + task: Mutex>>, +} + +#[cfg(not(feature = "management"))] +pub(super) struct EventJournal; + +#[cfg(feature = "management")] +impl EventJournal { + pub(super) fn new(global_ctx: &ArcGlobalCtx) -> Self { + Self { + global_ctx: global_ctx.clone(), + events: Arc::new(RwLock::new(VecDeque::new())), + receiver: Mutex::new(Some(global_ctx.subscribe())), + task: Mutex::new(None), + } + } + + pub(super) async fn start(&self, cancel: CancellationToken) { + let Some(mut receiver) = self.receiver.lock().await.take() else { + return; + }; + let events = self.events.clone(); + let task = tokio::spawn(async move { + loop { + let event = tokio::select! { + _ = cancel.cancelled() => return, + event = receiver.recv() => match event { + Ok(event) => event, + Err(tokio::sync::broadcast::error::RecvError::Closed) => return, + Err(tokio::sync::broadcast::error::RecvError::Lagged(_)) => continue, + }, + }; + let event = ManagementEvent { + time: chrono::Local::now(), + event, + }; + let Ok(event) = serde_json::to_string(&event) else { + continue; + }; + let mut events = events.write().unwrap(); + events.push_front(event); + if events.len() > 20 { + events.pop_back(); + } + } + }); + self.task.lock().await.replace(AbortOnDropHandle::new(task)); + } + + pub(super) async fn stop(&self) { + if let Some(task) = self.task.lock().await.take() { + let _ = task.await; + } + } + + pub(super) fn events(&self) -> Vec { + self.events.read().unwrap().iter().cloned().collect() + } + + pub(super) fn synchronize_config( + &self, + patch: &crate::proto::api::config::InstanceConfigPatch, + ) { + if let Some(hostname) = &patch.hostname { + self.global_ctx.set_hostname(hostname.clone()); + } + if let Some(ipv4) = patch.ipv4.as_ref() + && !self.global_ctx.config.get_dhcp() + { + self.global_ctx.set_ipv4(Some((*ipv4).into())); + } + if let Some(ipv6) = patch.ipv6.as_ref() { + self.global_ctx.set_ipv6(Some((*ipv6).into())); + } + if let Some(disable_relay_data) = patch.disable_relay_data { + let mut flags = self.global_ctx.get_flags(); + flags.disable_relay_data = disable_relay_data; + self.global_ctx.set_flags(flags); + } + } + + pub(super) fn publish_config_patch( + &self, + patch: crate::proto::api::config::InstanceConfigPatch, + ) { + self.global_ctx + .issue_event(GlobalCtxEvent::ConfigPatched(patch)); + } +} + +#[cfg(not(feature = "management"))] +impl EventJournal { + pub(super) fn new(_global_ctx: &ArcGlobalCtx) -> Self { + Self + } + + pub(super) async fn start(&self, _cancel: CancellationToken) {} + + pub(super) async fn stop(&self) {} + + pub(super) fn events(&self) -> Vec { + Vec::new() + } +} diff --git a/easytier/src/instance/runtime_host/implementation.rs b/easytier/src/instance/runtime_host/implementation.rs new file mode 100644 index 00000000..1e4e6e75 --- /dev/null +++ b/easytier/src/instance/runtime_host/implementation.rs @@ -0,0 +1,44 @@ +use std::sync::Arc; + +use easytier_core::{ + gateway::dhcp::DhcpIpv4Host, + instance::{CorePacketPlane, InstanceRuntimeHost}, +}; + +use super::NativeInstanceRuntimeHost; + +#[async_trait::async_trait] +impl InstanceRuntimeHost for NativeInstanceRuntimeHost { + async fn prepare( + &self, + packet_plane: Arc, + ) -> anyhow::Result>> { + self.prepare_runtime(packet_plane).await + } + + async fn shutdown(&self) { + self.shutdown_runtime().await; + } + + fn request_shutdown(&self) { + self.request_runtime_shutdown(); + } + + fn management_events(&self) -> Vec { + self.management_events_snapshot() + } + + #[cfg(feature = "management")] + fn synchronize_config(&self, patch: &crate::proto::api::config::InstanceConfigPatch) { + self.event_journal.synchronize_config(patch); + } + + #[cfg(feature = "management")] + fn publish_config_patch(&self, patch: crate::proto::api::config::InstanceConfigPatch) { + self.event_journal.publish_config_patch(patch); + } + + fn attach_tun_fd(&self, fd: i32) -> anyhow::Result<()> { + self.attach_runtime_tun_fd(fd) + } +} diff --git a/easytier/src/instance/runtime_host/magic_dns.rs b/easytier/src/instance/runtime_host/magic_dns.rs new file mode 100644 index 00000000..0b22fed7 --- /dev/null +++ b/easytier/src/instance/runtime_host/magic_dns.rs @@ -0,0 +1,74 @@ +use cidr::Ipv4Inet; +use easytier_core::instance::CorePacketPlane; +#[cfg(feature = "magic-dns")] +use tokio_util::{sync::CancellationToken, task::AbortOnDropHandle}; + +use crate::common::global_ctx::ArcGlobalCtx; +#[cfg(feature = "magic-dns")] +use crate::{ + common::config::ConfigLoader as _, + instance::dns_server::{MAGIC_DNS_FAKE_IP, runner::DnsRunner}, +}; + +#[derive(Default)] +pub(super) struct MagicDnsRuntime { + #[cfg(feature = "magic-dns")] + active: Option, +} + +#[cfg(feature = "magic-dns")] +struct MagicDnsTask { + task: AbortOnDropHandle<()>, + cancel: CancellationToken, +} + +impl MagicDnsRuntime { + #[cfg(feature = "magic-dns")] + pub(super) fn start( + global_ctx: ArcGlobalCtx, + packet_plane: std::sync::Arc, + tun_dev: Option, + tun_ip: Ipv4Inet, + ) -> Self { + let active = global_ctx.config.get_flags().accept_dns.then(|| { + let mut runner = DnsRunner::new( + packet_plane, + global_ctx, + tun_dev, + tun_ip, + MAGIC_DNS_FAKE_IP.parse().unwrap(), + ); + let cancel = CancellationToken::new(); + let task_cancel = cancel.clone(); + let task = tokio::spawn(async move { + let _ = runner.run(task_cancel).await; + }); + MagicDnsTask { + task: AbortOnDropHandle::new(task), + cancel, + } + }); + Self { active } + } + + #[cfg(not(feature = "magic-dns"))] + pub(super) fn start( + _global_ctx: ArcGlobalCtx, + _packet_plane: std::sync::Arc, + _tun_dev: Option, + _tun_ip: Ipv4Inet, + ) -> Self { + Self::default() + } + + #[cfg(feature = "magic-dns")] + pub(super) async fn stop(&mut self) { + if let Some(active) = self.active.take() { + active.cancel.cancel(); + let _ = active.task.await; + } + } + + #[cfg(not(feature = "magic-dns"))] + pub(super) async fn stop(&mut self) {} +} diff --git a/easytier/src/instance/runtime_host/tun_common.rs b/easytier/src/instance/runtime_host/tun_common.rs new file mode 100644 index 00000000..c4bfa91b --- /dev/null +++ b/easytier/src/instance/runtime_host/tun_common.rs @@ -0,0 +1,78 @@ +use std::{any::Any, sync::Arc}; + +use tokio::{sync::Mutex, task::JoinSet}; + +use super::{HostPacketReceiver, MagicDnsRuntime}; +use crate::instance::virtual_nic::NicCtx; + +struct NicCtxContainer { + _nic_ctx: Option>, + magic_dns: MagicDnsRuntime, +} + +impl NicCtxContainer { + fn new(nic_ctx: NicCtx, magic_dns: MagicDnsRuntime) -> Self { + Self { + _nic_ctx: Some(Box::new(nic_ctx)), + magic_dns, + } + } + + fn packet_drain(tasks: JoinSet<()>) -> Self { + Self { + _nic_ctx: Some(Box::new(tasks)), + magic_dns: MagicDnsRuntime::default(), + } + } +} + +#[derive(Clone)] +pub(super) struct TunNicState { + nic_ctx: Arc>>, + receiver: Arc>, +} + +impl TunNicState { + pub(super) fn new(receiver: HostPacketReceiver) -> Self { + Self { + nic_ctx: Arc::new(Mutex::new(None)), + receiver: Arc::new(Mutex::new(receiver)), + } + } + + pub(super) fn receiver(&self) -> Arc> { + self.receiver.clone() + } + + pub(super) async fn stop(&self) { + let mut old = self.nic_ctx.lock().await.take(); + if let Some(nic) = old.as_mut() { + nic.magic_dns.stop().await; + } + drop(old); + } + + pub(super) async fn drain(&self) { + self.stop().await; + let receiver = self.receiver.clone(); + let mut tasks = JoinSet::new(); + tasks.spawn(async move { + let mut receiver = receiver.lock().await; + while let Some(packet) = receiver.recv().await { + tracing::trace!(?packet, "discarded packet without a native interface"); + } + }); + self.nic_ctx + .lock() + .await + .replace(NicCtxContainer::packet_drain(tasks)); + } + + pub(super) async fn install(&self, nic: NicCtx, magic_dns: MagicDnsRuntime) { + self.stop().await; + self.nic_ctx + .lock() + .await + .replace(NicCtxContainer::new(nic, magic_dns)); + } +} diff --git a/easytier/src/instance/runtime_host/tun_desktop.rs b/easytier/src/instance/runtime_host/tun_desktop.rs new file mode 100644 index 00000000..96225c7a --- /dev/null +++ b/easytier/src/instance/runtime_host/tun_desktop.rs @@ -0,0 +1,253 @@ +use std::{sync::Arc, time::Duration}; + +use anyhow::Context as _; +use cidr::Ipv4Inet; +use easytier_core::{ + gateway::dhcp::{DhcpIpv4ApplyOutcome, DhcpIpv4ApplyPermit, DhcpIpv4Host}, + instance::CorePacketPlane, +}; +use futures::FutureExt as _; +use tokio::{ + sync::{Mutex, Notify, oneshot}, + task::JoinHandle, +}; +use tokio_util::sync::CancellationToken; + +use super::{HostPacketReceiver, MagicDnsRuntime, tun_common::TunNicState}; +use crate::{ + common::{ + config::ConfigLoader as _, + error::Error, + global_ctx::{ArcGlobalCtx, GlobalCtxEvent}, + }, + instance::virtual_nic::NicCtx, +}; + +pub(super) struct NativeTunRuntime { + global_ctx: ArcGlobalCtx, + cancel: CancellationToken, + nic: TunNicState, + static_ip_task: Mutex>>, +} + +impl NativeTunRuntime { + pub(super) fn new( + global_ctx: ArcGlobalCtx, + cancel: CancellationToken, + peer_packet_receiver: HostPacketReceiver, + ) -> Self { + Self { + global_ctx, + cancel, + nic: TunNicState::new(peer_packet_receiver), + static_ip_task: Mutex::new(None), + } + } + + fn report_static_ip_cancelled(output: &mut Option>>) { + if let Some(output) = output.take() { + let _ = output.send(Err(anyhow::anyhow!( + "instance is closing; static IP setup cancelled" + ) + .into())); + } + } + + async fn start_static_ip(&self, packet_plane: Arc) -> anyhow::Result<()> { + let ipv4 = self.global_ctx.get_ipv4(); + let ipv6 = self.global_ctx.get_ipv6(); + if ipv4.is_none() && ipv6.is_none() { + return Ok(()); + } + + let nic_state = self.nic.clone(); + let cancel = self.cancel.clone(); + let global_ctx = self.global_ctx.clone(); + let receiver = self.nic.receiver(); + let (output, first_round) = oneshot::channel(); + let task = tokio::spawn(async move { + let mut output = Some(output); + loop { + if cancel.is_cancelled() { + Self::report_static_ip_cancelled(&mut output); + return; + } + let closed = Arc::new(Notify::new()); + let mut nic = NicCtx::new( + global_ctx.clone(), + packet_plane.clone(), + receiver.clone(), + closed.clone(), + ); + let result = tokio::select! { + biased; + _ = cancel.cancelled() => { + Self::report_static_ip_cancelled(&mut output); + return; + } + result = nic.run(ipv4, ipv6) => result, + }; + if let Err(error) = result { + if let Some(output) = output.take() { + let _ = output.send(Err(error)); + return; + } + tracing::error!(?error, "failed to create native interface"); + tokio::select! { + _ = cancel.cancelled() => return, + _ = tokio::time::sleep(Duration::from_secs(1)) => {} + } + continue; + } + + let magic_dns = if let Some(ip) = ipv4 { + MagicDnsRuntime::start( + global_ctx.clone(), + packet_plane.clone(), + nic.ifname().await, + ip, + ) + } else { + MagicDnsRuntime::default() + }; + nic_state.install(nic, magic_dns).await; + if let Some(output) = output.take() { + let _ = output.send(Ok(())); + } + tokio::select! { + _ = cancel.cancelled() => return, + _ = closed.notified() => {} + } + } + }); + self.static_ip_task.lock().await.replace(task); + first_round + .await + .context("static IP setup task stopped")??; + Ok(()) + } + + pub(super) async fn prepare(&self, packet_plane: Arc) -> anyhow::Result<()> { + self.nic.drain().await; + if !self.global_ctx.config.get_flags().no_tun { + self.start_static_ip(packet_plane).await?; + } + Ok(()) + } + + pub(super) async fn shutdown(&self) { + if let Some(task) = self.static_ip_task.lock().await.take() { + let _ = task.await; + } + self.nic.stop().await; + } + + pub(super) fn attach_fd(&self, _fd: i32) -> anyhow::Result<()> { + anyhow::bail!("external TUN attachment is only supported on mobile Hosts") + } + + pub(super) fn dhcp_host( + &self, + operation: Arc>, + packet_plane: Arc, + ) -> Arc { + Arc::new(NativeDhcpIpv4Host { + global_ctx: self.global_ctx.clone(), + operation, + cancel: self.cancel.clone(), + nic: self.nic.clone(), + closed: Arc::new(Notify::new()), + packet_plane, + }) + } +} + +struct NativeDhcpIpv4Host { + global_ctx: ArcGlobalCtx, + operation: Arc>, + cancel: CancellationToken, + nic: TunNicState, + closed: Arc, + packet_plane: Arc, +} + +impl NativeDhcpIpv4Host { + fn ensure_open(&self) -> anyhow::Result<()> { + if self.cancel.is_cancelled() { + anyhow::bail!("instance is closing; DHCP update cancelled"); + } + Ok(()) + } + + async fn apply(&self, next: Option) -> anyhow::Result> { + self.ensure_open()?; + tokio::select! { + _ = self.cancel.cancelled() => anyhow::bail!("instance is closing; DHCP update cancelled"), + _ = self.nic.drain() => {} + } + self.ensure_open()?; + + let Some(ip) = next else { + self.global_ctx.set_ipv4(None); + return Ok(None); + }; + if self.global_ctx.no_tun() { + self.global_ctx.set_ipv4(Some(ip)); + return Ok(Some(ip)); + } + + let mut nic = NicCtx::new( + self.global_ctx.clone(), + self.packet_plane.clone(), + self.nic.receiver(), + self.closed.clone(), + ); + tokio::select! { + _ = self.cancel.cancelled() => anyhow::bail!("instance is closing; DHCP update cancelled"), + result = nic.run(Some(ip), self.global_ctx.get_ipv6()) => result?, + } + let magic_dns = MagicDnsRuntime::start( + self.global_ctx.clone(), + self.packet_plane.clone(), + nic.ifname().await, + ip, + ); + self.nic.install(nic, magic_dns).await; + self.global_ctx.set_ipv4(Some(ip)); + Ok(Some(ip)) + } +} + +#[async_trait::async_trait] +impl DhcpIpv4Host for NativeDhcpIpv4Host { + fn take_interface_closed(&self) -> bool { + self.closed.notified().now_or_never().is_some() + } + + async fn apply_dhcp_ipv4( + &self, + _previous: Option, + next: Option, + ) -> DhcpIpv4ApplyOutcome { + let permit = self.operation.clone().lock_owned().await; + let outcome = match self.apply(next).await { + Ok(actual) => DhcpIpv4ApplyOutcome::applied(actual), + Err(error) => DhcpIpv4ApplyOutcome::failed(self.global_ctx.get_ipv4(), error), + }; + outcome.with_permit(DhcpIpv4ApplyPermit::new(permit)) + } + + fn publish_dhcp_ipv4( + &self, + previous: Option, + requested: Option, + actual: Option, + ) { + let event = if requested.is_none() { + GlobalCtxEvent::DhcpIpv4Conflicted(previous) + } else { + GlobalCtxEvent::DhcpIpv4Changed(previous, actual) + }; + self.global_ctx.issue_event(event); + } +} diff --git a/easytier/src/instance/runtime_host/tun_disabled.rs b/easytier/src/instance/runtime_host/tun_disabled.rs new file mode 100644 index 00000000..6a34dd6a --- /dev/null +++ b/easytier/src/instance/runtime_host/tun_disabled.rs @@ -0,0 +1,109 @@ +use std::sync::Arc; + +use cidr::Ipv4Inet; +use easytier_core::{ + gateway::dhcp::{DhcpIpv4ApplyOutcome, DhcpIpv4ApplyPermit, DhcpIpv4Host}, + instance::CorePacketPlane, +}; +use tokio::sync::Mutex; +use tokio_util::sync::CancellationToken; + +use super::HostPacketReceiver; +use crate::common::global_ctx::{ArcGlobalCtx, GlobalCtxEvent}; + +pub(super) struct NativeTunRuntime { + global_ctx: ArcGlobalCtx, + cancel: CancellationToken, +} + +impl NativeTunRuntime { + pub(super) fn new( + global_ctx: ArcGlobalCtx, + cancel: CancellationToken, + peer_packet_receiver: HostPacketReceiver, + ) -> Self { + drop(peer_packet_receiver); + Self { global_ctx, cancel } + } + + pub(super) async fn prepare(&self, _packet_plane: Arc) -> anyhow::Result<()> { + Ok(()) + } + + pub(super) async fn shutdown(&self) {} + + pub(super) fn attach_fd(&self, _fd: i32) -> anyhow::Result<()> { + anyhow::bail!("external TUN attachment is only supported on mobile Hosts") + } + + pub(super) fn dhcp_host( + &self, + operation: Arc>, + _packet_plane: Arc, + ) -> Arc { + Arc::new(NativeDhcpIpv4Host { + global_ctx: self.global_ctx.clone(), + operation, + cancel: self.cancel.clone(), + }) + } +} + +struct NativeDhcpIpv4Host { + global_ctx: ArcGlobalCtx, + operation: Arc>, + cancel: CancellationToken, +} + +impl NativeDhcpIpv4Host { + fn ensure_open(&self) -> anyhow::Result<()> { + if self.cancel.is_cancelled() { + anyhow::bail!("instance is closing; DHCP update cancelled"); + } + Ok(()) + } + + async fn apply(&self, next: Option) -> anyhow::Result> { + self.ensure_open()?; + let Some(ip) = next else { + self.global_ctx.set_ipv4(None); + return Ok(None); + }; + self.global_ctx.set_ipv4(Some(ip)); + Ok(Some(ip)) + } +} + +#[async_trait::async_trait] +impl DhcpIpv4Host for NativeDhcpIpv4Host { + fn take_interface_closed(&self) -> bool { + false + } + + async fn apply_dhcp_ipv4( + &self, + _previous: Option, + next: Option, + ) -> DhcpIpv4ApplyOutcome { + let permit = self.operation.clone().lock_owned().await; + let outcome = match self.apply(next).await { + Ok(actual) => DhcpIpv4ApplyOutcome::applied(actual), + Err(error) => DhcpIpv4ApplyOutcome::failed(self.global_ctx.get_ipv4(), error), + }; + outcome.with_permit(DhcpIpv4ApplyPermit::new(permit)) + } + + fn publish_dhcp_ipv4( + &self, + previous: Option, + requested: Option, + actual: Option, + ) { + let event = if requested.is_none() { + GlobalCtxEvent::DhcpIpv4Conflicted(previous) + } else { + GlobalCtxEvent::DhcpIpv4Changed(previous, actual) + }; + self.global_ctx.issue_event(event); + } +} diff --git a/easytier/src/instance/runtime_host/tun_mobile.rs b/easytier/src/instance/runtime_host/tun_mobile.rs new file mode 100644 index 00000000..79216a55 --- /dev/null +++ b/easytier/src/instance/runtime_host/tun_mobile.rs @@ -0,0 +1,193 @@ +use std::sync::Arc; + +use anyhow::Context as _; +use cidr::Ipv4Inet; +use easytier_core::{ + gateway::dhcp::{DhcpIpv4ApplyOutcome, DhcpIpv4ApplyPermit, DhcpIpv4Host}, + instance::CorePacketPlane, +}; +use futures::FutureExt as _; +use tokio::sync::{Mutex, Notify, mpsc}; +use tokio_util::sync::CancellationToken; + +use super::{HostPacketReceiver, MagicDnsRuntime, tun_common::TunNicState}; +use crate::{ + common::global_ctx::{ArcGlobalCtx, GlobalCtxEvent}, + instance::virtual_nic::NicCtx, +}; + +pub(super) struct NativeTunRuntime { + global_ctx: ArcGlobalCtx, + cancel: CancellationToken, + nic: TunNicState, + tun_fd: mpsc::Sender, + tun_fd_receiver: Mutex>>, + task: Mutex>>, +} + +impl NativeTunRuntime { + pub(super) fn new( + global_ctx: ArcGlobalCtx, + cancel: CancellationToken, + peer_packet_receiver: HostPacketReceiver, + ) -> Self { + let (tun_fd, tun_fd_receiver) = mpsc::channel(16); + Self { + global_ctx, + cancel, + nic: TunNicState::new(peer_packet_receiver), + tun_fd, + tun_fd_receiver: Mutex::new(Some(tun_fd_receiver)), + task: Mutex::new(None), + } + } + + async fn install_mobile_tun( + nic_state: TunNicState, + global_ctx: ArcGlobalCtx, + packet_plane: Arc, + fd: i32, + ) -> anyhow::Result<()> { + nic_state.drain().await; + if fd <= 0 { + return Ok(()); + } + let closed = Arc::new(Notify::new()); + let mut nic = NicCtx::new( + global_ctx.clone(), + packet_plane.clone(), + nic_state.receiver(), + closed, + ); + nic.run_for_mobile(fd).await.context("add ip failed")?; + let magic_dns = global_ctx + .get_ipv4() + .map(|ip| MagicDnsRuntime::start(global_ctx, packet_plane, None, ip)) + .unwrap_or_default(); + nic_state.install(nic, magic_dns).await; + Ok(()) + } + + pub(super) async fn prepare(&self, packet_plane: Arc) -> anyhow::Result<()> { + self.nic.drain().await; + let Some(mut tun_fds) = self.tun_fd_receiver.lock().await.take() else { + return Ok(()); + }; + let nic_state = self.nic.clone(); + let global_ctx = self.global_ctx.clone(); + let cancel = self.cancel.clone(); + self.task.lock().await.replace(tokio::spawn(async move { + loop { + let fd = tokio::select! { + _ = cancel.cancelled() => return, + fd = tun_fds.recv() => match fd { Some(fd) => fd, None => return }, + }; + if let Err(error) = Self::install_mobile_tun( + nic_state.clone(), + global_ctx.clone(), + packet_plane.clone(), + fd, + ) + .await + { + tracing::error!(?error, "failed to attach mobile TUN fd"); + } + } + })); + Ok(()) + } + + pub(super) async fn shutdown(&self) { + if let Some(task) = self.task.lock().await.take() { + let _ = task.await; + } + self.nic.stop().await; + } + + pub(super) fn attach_fd(&self, fd: i32) -> anyhow::Result<()> { + self.tun_fd + .try_send(fd) + .map_err(|error| anyhow::anyhow!("failed to send TUN fd: {error}")) + } + + pub(super) fn dhcp_host( + &self, + operation: Arc>, + _packet_plane: Arc, + ) -> Arc { + Arc::new(NativeDhcpIpv4Host { + global_ctx: self.global_ctx.clone(), + operation, + cancel: self.cancel.clone(), + nic: self.nic.clone(), + closed: Arc::new(Notify::new()), + }) + } +} + +struct NativeDhcpIpv4Host { + global_ctx: ArcGlobalCtx, + operation: Arc>, + cancel: CancellationToken, + nic: TunNicState, + closed: Arc, +} + +impl NativeDhcpIpv4Host { + fn ensure_open(&self) -> anyhow::Result<()> { + if self.cancel.is_cancelled() { + anyhow::bail!("instance is closing; DHCP update cancelled"); + } + Ok(()) + } + + async fn apply(&self, next: Option) -> anyhow::Result> { + self.ensure_open()?; + tokio::select! { + _ = self.cancel.cancelled() => anyhow::bail!("instance is closing; DHCP update cancelled"), + _ = self.nic.drain() => {} + } + self.ensure_open()?; + + let Some(ip) = next else { + self.global_ctx.set_ipv4(None); + return Ok(None); + }; + self.global_ctx.set_ipv4(Some(ip)); + Ok(Some(ip)) + } +} + +#[async_trait::async_trait] +impl DhcpIpv4Host for NativeDhcpIpv4Host { + fn take_interface_closed(&self) -> bool { + self.closed.notified().now_or_never().is_some() + } + + async fn apply_dhcp_ipv4( + &self, + _previous: Option, + next: Option, + ) -> DhcpIpv4ApplyOutcome { + let permit = self.operation.clone().lock_owned().await; + let outcome = match self.apply(next).await { + Ok(actual) => DhcpIpv4ApplyOutcome::applied(actual), + Err(error) => DhcpIpv4ApplyOutcome::failed(self.global_ctx.get_ipv4(), error), + }; + outcome.with_permit(DhcpIpv4ApplyPermit::new(permit)) + } + + fn publish_dhcp_ipv4( + &self, + previous: Option, + requested: Option, + actual: Option, + ) { + let event = if requested.is_none() { + GlobalCtxEvent::DhcpIpv4Conflicted(previous) + } else { + GlobalCtxEvent::DhcpIpv4Changed(previous, actual) + }; + self.global_ctx.issue_event(event); + } +} diff --git a/easytier/src/instance/test_instance.rs b/easytier/src/instance/test_instance.rs new file mode 100644 index 00000000..6058cce7 --- /dev/null +++ b/easytier/src/instance/test_instance.rs @@ -0,0 +1,135 @@ +//! Test-only convenience around the production CoreInstance composition. + +use std::sync::Arc; + +use easytier_core::{ + config::toml::TomlConfig, connectivity::stun::StunSocketMapper, instance::CoreInstance, + process_runtime::CoreProcessRuntime, +}; + +use crate::{ + common::global_ctx::{ArcGlobalCtx, GlobalCtx}, + instance::{ + composition::{NativeCoreInstance, runtime_core_host_adapters}, + runtime_host::NativeInstanceRuntimeHost, + }, + socket::udp::RuntimeUdpSocket, +}; + +pub(crate) struct TestInstance { + core: Arc, + global_ctx: ArcGlobalCtx, +} + +impl TestInstance { + pub fn new_with_process_runtime( + config: TomlConfig, + process_runtime: Arc, + ) -> Self { + Self::compose(config, process_runtime, |_| {}) + } + + pub fn new_with_process_runtime_and_stun_provider( + config: TomlConfig, + process_runtime: Arc, + provider: Box>, + ) -> Self { + let provider: Arc> = Arc::from(provider); + Self::compose(config, process_runtime, move |adapters| { + adapters.replace_stun_provider(provider); + }) + } + + fn compose( + config: TomlConfig, + process_runtime: Arc, + customize: impl FnOnce( + &mut easytier_core::instance::CoreHostAdapters< + crate::instance::host::NativeInstanceHost, + >, + ), + ) -> Self { + let global_ctx = Arc::new(GlobalCtx::new(config.clone())); + let (packet_sender, packet_receiver) = tokio::sync::mpsc::channel(128); + let mut adapters = runtime_core_host_adapters( + global_ctx.clone(), + process_runtime, + Arc::new(packet_sender), + ); + customize(&mut adapters); + adapters.instance_runtime = + NativeInstanceRuntimeHost::new(global_ctx.clone(), packet_receiver); + let core = CoreInstance::from_toml(config, adapters) + .expect("test CoreInstance composition should be valid"); + Self { core, global_ctx } + } + + pub async fn run(&mut self) -> anyhow::Result<()> { + self.core.start().await + } + + pub async fn clear_resources(&mut self) { + self.core.stop().await; + } + + pub fn get_core_instance(&self) -> Arc { + self.core.clone() + } + + pub fn get_global_ctx(&self) -> ArcGlobalCtx { + self.global_ctx.clone() + } + + pub fn get_config_patcher(&self) -> TestConfigPatcher { + TestConfigPatcher { + core: self.core.clone(), + } + } +} + +pub(crate) struct TestConfigPatcher { + core: Arc, +} + +impl TestConfigPatcher { + pub async fn apply_patch( + &self, + patch: crate::proto::api::config::InstanceConfigPatch, + ) -> anyhow::Result<()> { + easytier_core::management::apply_config_patch(&self.core, patch).await + } +} + +#[cfg(test)] +mod tests { + use easytier_core::config::{ + normalize_secure_mode_config, + toml::{ConfigLoader as _, TomlConfig}, + }; + + use super::*; + + #[tokio::test] + async fn composition_preserves_secure_admin_identity() { + let config = TomlConfig::default(); + config.set_secure_mode(Some( + normalize_secure_mode_config(crate::proto::common::SecureModeConfig { + enabled: true, + ..Default::default() + }) + .unwrap(), + )); + + let instance = TestInstance::new_with_process_runtime(config, CoreProcessRuntime::new()); + + assert_eq!( + instance + .get_global_ctx() + .config + .get_network_identity() + .network_secret + .as_deref(), + Some("") + ); + } +} diff --git a/easytier/src/instance/udp_hole_punch.rs b/easytier/src/instance/udp_hole_punch.rs new file mode 100644 index 00000000..2e159722 --- /dev/null +++ b/easytier/src/instance/udp_hole_punch.rs @@ -0,0 +1,41 @@ +//! Native UDP hole-punch platform adapter. +//! +//! Peer selection, signaling, RPC registration, socket/session ownership and +//! lifecycle live in `easytier-core`. Native only supplies OS port mapping. + +use std::sync::Arc; + +use async_trait::async_trait; +use easytier_core::connectivity::hole_punch::port_mapping::{ + ActiveUdpPortMapping, UdpPortMappingAttemptError, UdpPortMappingBackend, + UdpPortMappingLifecycle, UdpPortMappingPlatform, +}; + +use crate::common::{netns::NetNS, upnp}; + +struct RuntimeUdpHolePunchPlatform { + net_ns: NetNS, +} + +#[async_trait] +impl UdpPortMappingPlatform for RuntimeUdpHolePunchPlatform { + async fn establish_udp_port_mapping( + &self, + backend: UdpPortMappingBackend, + local_listener: &url::Url, + ) -> Result, UdpPortMappingAttemptError> { + upnp::establish_udp_port_mapping(self.net_ns.clone(), backend, local_listener.clone()).await + } + + fn spawn_udp_port_mapping_lifecycle( + &self, + local_listener: url::Url, + lifecycle: UdpPortMappingLifecycle, + ) { + upnp::spawn_udp_port_mapping_lifecycle(self.net_ns.clone(), local_listener, lifecycle); + } +} + +pub(crate) fn runtime_udp_hole_punch_platform(net_ns: NetNS) -> Arc { + Arc::new(RuntimeUdpHolePunchPlatform { net_ns }) +} diff --git a/easytier/src/instance/virtual_nic.rs b/easytier/src/instance/virtual_nic.rs index 2ff52af4..49845388 100644 --- a/easytier/src/instance/virtual_nic.rs +++ b/easytier/src/instance/virtual_nic.rs @@ -1,25 +1,25 @@ use std::{ collections::BTreeSet, io, - net::{IpAddr, Ipv4Addr, Ipv6Addr}, + net::{Ipv4Addr, Ipv6Addr}, pin::Pin, - sync::{Arc, Weak}, + sync::Arc, task::{Context, Poll}, }; -use crate::{ - common::{ - error::Error, - global_ctx::{ArcGlobalCtx, GlobalCtxEvent}, - ifcfg::{IfConfiger, IfConfiguerTrait}, - log, - }, - instance::proxy_cidrs_monitor::ProxyCidrsMonitor, - peers::{PacketRecvChanReceiver, peer_manager::PeerManager, recv_packet_from_chan}, +use crate::common::{ + error::Error, + global_ctx::{ArcGlobalCtx, GlobalCtxEvent}, + ifcfg::{IfConfiger, IfConfiguerTrait}, +}; + +use easytier_core::{ + instance::CorePacketPlane, + packet::{TAIL_RESERVED_SIZE, ZCPacket, ZCPacketType}, tunnel::{ StreamItem, Tunnel, TunnelError, ZCPacketSink, ZCPacketStream, - common::{FramedWriter, TunnelWrapper, ZCPacketToBytes, reserve_buf}, - packet_def::{TAIL_RESERVED_SIZE, ZCPacket, ZCPacketType}, + framed::{FramedWriter, ZCPacketToBytes, reserve_buf}, + wrapper::TunnelWrapper, }, }; @@ -28,7 +28,6 @@ use bytes::{Buf, BufMut, BytesMut}; use cidr::{Ipv4Inet, Ipv6Inet}; use futures::{SinkExt, Stream, StreamExt, lock::BiLock, ready}; use pin_project_lite::pin_project; -use pnet::packet::{ipv4::Ipv4Packet, ipv6::Ipv6Packet}; use tokio::{ io::{AsyncRead, AsyncWrite, ReadBuf}, sync::{Mutex, Notify}, @@ -43,6 +42,8 @@ use zerocopy::{NativeEndian, NetworkEndian}; #[cfg(target_os = "windows")] use crate::common::ifcfg::RegistryManager; +type HostPacketReceiver = tokio::sync::mpsc::Receiver>; + pin_project! { pub struct TunStream { #[pin] @@ -98,7 +99,7 @@ impl Stream for TunStream { match ret { Ok(_) => Poll::Ready(Some(Ok(ZCPacket::new_from_buf(ret_buf, ZCPacketType::NIC)))), Err(err) => { - log::error!("tun stream error: {:?}", err); + tracing::error!("tun stream error: {:?}", err); Poll::Ready(None) } } @@ -110,7 +111,7 @@ enum PacketProtocol { #[default] IPv4, IPv6, - Other(u8), + Other, } // Note: the protocol in the packet information header is platform dependent. @@ -121,7 +122,7 @@ impl PacketProtocol { match self { PacketProtocol::IPv4 => Ok(libc::ETH_P_IP as u16), PacketProtocol::IPv6 => Ok(libc::ETH_P_IPV6 as u16), - PacketProtocol::Other(_) => Err(io::Error::other("neither an IPv4 nor IPv6 packet")), + PacketProtocol::Other => Err(io::Error::other("neither an IPv4 nor IPv6 packet")), } } @@ -131,7 +132,7 @@ impl PacketProtocol { match self { PacketProtocol::IPv4 => Ok(libc::PF_INET as u16), PacketProtocol::IPv6 => Ok(libc::PF_INET6 as u16), - PacketProtocol::Other(_) => Err(io::Error::other("neither an IPv4 nor IPv6 packet")), + PacketProtocol::Other => Err(io::Error::other("neither an IPv4 nor IPv6 packet")), } } @@ -146,7 +147,7 @@ fn infer_proto(buf: &[u8]) -> PacketProtocol { match buf[0] >> 4 { 4 => PacketProtocol::IPv4, 6 => PacketProtocol::IPv6, - p => PacketProtocol::Other(p), + _ => PacketProtocol::Other, } } @@ -254,7 +255,7 @@ impl Drop for VirtualNic { if let Some(ref ifname) = self.ifname { // Try to clean up firewall rules, but don't panic in destructor if let Err(error) = crate::arch::windows::remove_interface_firewall_rules(ifname) { - log::warn!( + tracing::warn!( %error, "failed to remove firewall rules for interface {}", ifname @@ -298,19 +299,19 @@ impl VirtualNic { .unwrap_or(false); if !tun_module_available { - log::warn!("TUN kernel module may not be available."); - log::warn!("\tYou may need to load it with: sudo modprobe tun."); + tracing::warn!("TUN kernel module may not be available."); + tracing::warn!("\tYou may need to load it with: sudo modprobe tun."); } // Try to create /dev/net directory if it doesn't exist if tokio::fs::metadata(TUN_DIR_PATH).await.is_err() { if let Err(error) = tokio::fs::create_dir_all(TUN_DIR_PATH).await { - log::warn!( + tracing::warn!( ?error, "Failed to create directory {}. TUN device creation may fail. Continuing anyway.", TUN_DIR_PATH ); - log::warn!( + tracing::warn!( "\tYou may need to run with root privileges or manually create the TUN device." ); Self::print_troubleshooting_info(); @@ -330,7 +331,7 @@ impl VirtualNic { dev_node, ) { Ok(_) => { - log::info!("Successfully created TUN device node {}", TUN_DEV_PATH); + tracing::info!("Successfully created TUN device node {}", TUN_DEV_PATH); } Err(error) => { tracing::warn!( @@ -346,7 +347,7 @@ impl VirtualNic { /// Print troubleshooting information for TUN device issues #[cfg(target_os = "linux")] fn print_troubleshooting_info() { - log::info!( + tracing::info!( "Possible solutions:\ \n\t1. Run with root privileges: sudo ./easytier-core [options]\ \n\t2. Manually create TUN device: sudo mkdir -p /dev/net && sudo mknod /dev/net/tun c 10 200\ @@ -357,12 +358,6 @@ impl VirtualNic { ); } - /// For non-Linux systems, this is a no-op - #[cfg(not(target_os = "linux"))] - async fn ensure_tun_device_node() -> Result<(), Error> { - Ok(()) - } - /// FreeBSD specific: Rename a TUN interface #[cfg(target_os = "freebsd")] async fn rename_tun_interface(old_name: &str, new_name: &str) -> Result<(), Error> { @@ -532,8 +527,8 @@ impl VirtualNic { match crate::arch::windows::add_self_to_firewall_allowlist() { Ok(_) => tracing::info!("add_self_to_firewall_allowlist successful!"), Err(error) => { - log::warn!(%error, "Failed to add Easytier to firewall allowlist, Subnet proxy and KCP proxy may not work properly."); - log::warn!( + tracing::warn!(%error, "Failed to add Easytier to firewall allowlist, Subnet proxy and KCP proxy may not work properly."); + tracing::warn!( "You can add firewall rules manually, or use --use-smoltcp to run with user-space TCP/IP stack." ); } @@ -585,7 +580,7 @@ impl VirtualNic { &mut self, tun_fd: std::os::fd::RawFd, ) -> Result, Error> { - log::debug!(%tun_fd); + tracing::debug!(%tun_fd); let mut config = Configuration::default(); config.layer(Layer::L3); @@ -708,7 +703,7 @@ impl VirtualNic { ); } Err(error) => { - log::warn!(%error, "Failed to configure Windows Firewall for interface {}\ + tracing::warn!(%error, "Failed to configure Windows Firewall for interface {}\ \n\tThis may cause connectivity issues with ping and other network functions.\ \n\tPlease run as Administrator or manually configure Windows Firewall.\ \n\tAlternatively, you can disable Windows Firewall for testing purposes.", ifname); @@ -797,8 +792,8 @@ impl VirtualNic { pub struct NicCtx { global_ctx: ArcGlobalCtx, - peer_mgr: Weak, - peer_packet_receiver: Arc>, + packet_plane: Arc, + peer_packet_receiver: Arc>, close_notifier: Arc, @@ -810,15 +805,15 @@ pub struct NicCtx { } impl NicCtx { - pub fn new( + pub(crate) fn new( global_ctx: ArcGlobalCtx, - peer_manager: &Arc, - peer_packet_receiver: Arc>, + packet_plane: Arc, + peer_packet_receiver: Arc>, close_notifier: Arc, ) -> Self { NicCtx { global_ctx: global_ctx.clone(), - peer_mgr: Arc::downgrade(peer_manager), + packet_plane, peer_packet_receiver, close_notifier, @@ -870,102 +865,17 @@ impl NicCtx { Ok(()) } - async fn do_forward_nic_to_peers_ipv4(ret: ZCPacket, mgr: &PeerManager) { - if let Some(ipv4) = Ipv4Packet::new(ret.payload()) { - if ipv4.get_version() != 4 { - tracing::info!("[USER_PACKET] not ipv4 packet: {:?}", ipv4); - return; - } - let dst_ipv4 = ipv4.get_destination(); - let src_ipv4 = ipv4.get_source(); - let my_ipv4 = mgr.get_global_ctx().get_ipv4().map(|x| x.address()); - tracing::trace!( - ?ret, - ?src_ipv4, - ?dst_ipv4, - "[USER_PACKET] recv new packet from tun device and forward to peers." - ); - - // Subnet A is proxied as 10.0.0.0/24, and Subnet B is also proxied as 10.0.0.0/24. - // - // Subnet A has received a route advertised by Subnet B. As a result, A can reach - // the physical subnet 10.0.0.0/24 directly and has also added a virtual route for - // the same subnet 10.0.0.0/24. However, the physical route has a higher priority - // (lower metric) than the virtual one. - // - // When A sends a UDP packet to a non-existent IP within this subnet, the packet - // cannot be delivered on the physical network and is instead routed to the virtual - // network interface. - // - // The virtual interface receives the packet and forwards it to itself, which triggers - // the subnet proxy logic. The subnet proxy then attempts to send another packet to - // the same destination address, causing the same process to repeat and creating an - // infinite loop. Therefore, we must avoid re-sending packets back to ourselves - // when the subnet proxy itself is the originator of the packet. - // - // However, there is a special scenario to consider: when A acts as a gateway, - // packets from devices behind A may be forwarded by the OS to the ET (e.g., an - // eBPF or tunneling component), which happens to proxy the subnet. In this case, - // the packet’s source IP is not A’s own IP, and we must allow such packets to be - // sent to the virtual interface (i.e., "sent to ourselves") to maintain correct - // forwarding behavior. Thus, loop prevention should only apply when the source IP - // belongs to the local host. - let send_ret = mgr - .send_msg_by_ip(ret, IpAddr::V4(dst_ipv4), Some(src_ipv4) == my_ipv4) - .await; - if send_ret.is_err() { - tracing::trace!(?send_ret, "[USER_PACKET] send_msg failed") - } - } else { - tracing::warn!(?ret, "[USER_PACKET] not ipv4 packet"); - } - } - - async fn do_forward_nic_to_peers_ipv6(ret: ZCPacket, mgr: &PeerManager) { - if let Some(ipv6) = Ipv6Packet::new(ret.payload()) { - if ipv6.get_version() != 6 { - tracing::info!("[USER_PACKET] not ipv6 packet: {:?}", ipv6); - return; - } - let src_ipv6 = ipv6.get_source(); - let dst_ipv6 = ipv6.get_destination(); - let is_local_src = mgr.get_global_ctx().is_ip_local_ipv6(&src_ipv6); - tracing::trace!( - ?ret, - ?src_ipv6, - ?dst_ipv6, - "[USER_PACKET] recv new packet from tun device and forward to peers." - ); - - if src_ipv6.is_unicast_link_local() && !is_local_src { - // do not route link local packet to other nodes unless the address is assigned by user - return; - } - - // TODO: use zero-copy - let send_ret = mgr - .send_msg_by_ip(ret, IpAddr::V6(dst_ipv6), is_local_src) - .await; - if send_ret.is_err() { - tracing::trace!(?send_ret, "[USER_PACKET] send_msg failed") - } - } else { - tracing::warn!(?ret, "[USER_PACKET] not ipv6 packet"); - } - } - - async fn do_forward_nic_to_peers(ret: ZCPacket, mgr: &PeerManager) { + async fn do_forward_nic_to_peers(ret: ZCPacket, packet_plane: &CorePacketPlane) { let payload = ret.payload(); if payload.is_empty() { return; } - - match payload[0] >> 4 { - 4 => Self::do_forward_nic_to_peers_ipv4(ret, mgr).await, - 6 => Self::do_forward_nic_to_peers_ipv6(ret, mgr).await, - _ => { - tracing::warn!(?ret, "[USER_PACKET] unknown IP version"); - } + tracing::trace!( + ?ret, + "[USER_PACKET] recv new packet from tun device and forward to peers." + ); + if let Err(error) = packet_plane.send_ip_packet(payload.to_vec()).await { + tracing::trace!(?error, "[USER_PACKET] send_msg failed"); } } @@ -974,9 +884,7 @@ impl NicCtx { mut stream: Pin>, ) -> Result<(), Error> { // read from nic and write to corresponding tunnel - let Some(mgr) = self.peer_mgr.upgrade() else { - return Err(anyhow::anyhow!("peer manager not available").into()); - }; + let packet_plane = self.packet_plane.clone(); let close_notifier = self.close_notifier.clone(); self.tasks.spawn(async move { while let Some(ret) = stream.next().await { @@ -984,7 +892,7 @@ impl NicCtx { tracing::error!("read from nic failed: {:?}", ret); break; } - Self::do_forward_nic_to_peers(ret.unwrap(), mgr.as_ref()).await; + Self::do_forward_nic_to_peers(ret.unwrap(), packet_plane.as_ref()).await; } close_notifier.notify_one(); tracing::error!("nic closed when recving from it"); @@ -999,12 +907,12 @@ impl NicCtx { self.tasks.spawn(async move { // unlock until coroutine finished let mut channel = channel.lock().await; - while let Ok(packet) = recv_packet_from_chan(&mut channel).await { + while let Some(packet) = channel.recv().await { tracing::trace!( "[USER_PACKET] forward packet from peers to nic. packet: {:?}", packet ); - let ret = sink.send(packet).await; + let ret = sink.send(ZCPacket::new_with_payload(&packet)).await; if ret.is_err() { tracing::error!(?ret, "do_forward_tunnel_to_nic sink error"); } @@ -1020,12 +928,11 @@ impl NicCtx { return; } - let Some(peer_manager) = self.peer_mgr.upgrade() else { - tracing::warn!("peer manager is dropped, skip Windows UDP broadcast relay"); - return; - }; - - match super::windows_udp_broadcast::start(peer_manager, virtual_ipv4) { + match super::windows_udp_broadcast::start( + self.packet_plane.clone(), + self.global_ctx.clone(), + virtual_ipv4, + ) { Ok(handle) => { self.windows_udp_broadcast_relay = Some(handle); tracing::info!("Windows UDP broadcast relay started"); @@ -1129,9 +1036,7 @@ impl NicCtx { } async fn run_proxy_cidrs_route_updater(&mut self) -> Result<(), Error> { - let Some(peer_mgr) = self.peer_mgr.upgrade() else { - return Err(anyhow::anyhow!("peer manager not available").into()); - }; + let packet_plane = self.packet_plane.clone(); let global_ctx = self.global_ctx.clone(); let net_ns = self.global_ctx.net_ns.clone(); let nic = self.nic.lock().await; @@ -1143,19 +1048,17 @@ impl NicCtx { let mut cur_proxy_cidrs = BTreeSet::::new(); // Initial sync: get current proxy_cidrs state and apply routes - let (_, added, removed) = ProxyCidrsMonitor::diff_proxy_cidrs( - peer_mgr.as_ref(), - &global_ctx, - &cur_proxy_cidrs, - ) - .await; + let Some(diff) = packet_plane.proxy_cidr_diff(&cur_proxy_cidrs).await else { + tracing::error!("proxy CIDR monitor host is unavailable"); + return; + }; Self::apply_route_changes( &ifcfg, &ifname, &net_ns, &mut cur_proxy_cidrs, - added, - removed, + diff.added, + diff.removed, ) .await; @@ -1172,13 +1075,12 @@ impl NicCtx { ); event_receiver = event_receiver.resubscribe(); // Full sync after lagged to recover consistent state - let (_, added, removed) = ProxyCidrsMonitor::diff_proxy_cidrs( - peer_mgr.as_ref(), - &global_ctx, - &cur_proxy_cidrs, - ) - .await; - GlobalCtxEvent::ProxyCidrsUpdated(added, removed) + let Some(diff) = packet_plane.proxy_cidr_diff(&cur_proxy_cidrs).await + else { + tracing::error!("proxy CIDR monitor host is unavailable"); + return; + }; + GlobalCtxEvent::ProxyCidrsUpdated(diff.added, diff.removed) } }; @@ -1204,9 +1106,7 @@ impl NicCtx { } async fn run_public_ipv6_route_updater(&mut self) -> Result<(), Error> { - let Some(peer_mgr) = self.peer_mgr.upgrade() else { - return Err(anyhow::anyhow!("peer manager not available").into()); - }; + let packet_plane = self.packet_plane.clone(); let global_ctx = self.global_ctx.clone(); let net_ns = self.global_ctx.net_ns.clone(); let nic = self.nic.lock().await; @@ -1216,7 +1116,7 @@ impl NicCtx { self.tasks.spawn(async move { let mut cur_routes = BTreeSet::::new(); - let initial_routes = peer_mgr.list_public_ipv6_routes().await; + let initial_routes = packet_plane.public_ipv6_routes().await; let initial_added = initial_routes.iter().copied().collect::>(); Self::apply_public_ipv6_route_changes( &ifcfg, @@ -1234,7 +1134,7 @@ impl NicCtx { Err(tokio::sync::broadcast::error::RecvError::Closed) => break, Err(tokio::sync::broadcast::error::RecvError::Lagged(_)) => { event_receiver = event_receiver.resubscribe(); - let latest = peer_mgr.list_public_ipv6_routes().await; + let latest = packet_plane.public_ipv6_routes().await; let added = latest.difference(&cur_routes).copied().collect::>(); let removed = cur_routes.difference(&latest).copied().collect::>(); GlobalCtxEvent::PublicIpv6RoutesUpdated(added, removed) @@ -1262,15 +1162,13 @@ impl NicCtx { } async fn run_public_ipv6_addr_updater(&mut self) -> Result<(), Error> { - let Some(peer_mgr) = self.peer_mgr.upgrade() else { - return Err(anyhow::anyhow!("peer manager not available").into()); - }; + let packet_plane = self.packet_plane.clone(); let global_ctx = self.global_ctx.clone(); let nic = self.nic.clone(); let mut event_receiver = global_ctx.subscribe(); self.tasks.spawn(async move { - let mut current_addr = peer_mgr.get_my_public_ipv6_addr().await; + let mut current_addr = packet_plane.public_ipv6_addr().await; if let Some(addr) = current_addr { let nic = nic.lock().await; if let Err(err) = nic.link_up().await { @@ -1293,7 +1191,7 @@ impl NicCtx { Err(tokio::sync::broadcast::error::RecvError::Closed) => break, Err(tokio::sync::broadcast::error::RecvError::Lagged(_)) => { event_receiver = event_receiver.resubscribe(); - let latest = peer_mgr.get_my_public_ipv6_addr().await; + let latest = packet_plane.public_ipv6_addr().await; GlobalCtxEvent::PublicIpv6Changed(current_addr, latest) } }; diff --git a/easytier/src/instance/windows_udp_broadcast.rs b/easytier/src/instance/windows_udp_broadcast.rs index b49ee9e4..4eaab566 100644 --- a/easytier/src/instance/windows_udp_broadcast.rs +++ b/easytier/src/instance/windows_udp_broadcast.rs @@ -1,443 +1,12 @@ use std::net::Ipv4Addr; -use cidr::Ipv4Inet; -use pnet::packet::{ - ip::IpNextHeaderProtocols, - ipv4::{self, Ipv4Flags, Ipv4Packet, MutableIpv4Packet}, - udp::{self, MutableUdpPacket, UdpPacket}, -}; +use easytier_core::gateway::udp_broadcast::PhysicalInterface; -#[cfg(any(windows, test))] -use { - crate::{ - common::global_ctx::GlobalCtxEvent, - common::stats_manager::{CounterHandle, LabelSet, LabelType, MetricName}, - peers::peer_manager::PeerManager, - tunnel::packet_def::ZCPacket, - }, - anyhow::Context, - network_interface::{Addr, NetworkInterface, NetworkInterfaceConfig}, - socket2::{Domain, Protocol, SockAddr, Socket, Type}, - std::{ - io, - net::{IpAddr, SocketAddrV4, UdpSocket as StdUdpSocket}, - sync::Arc, - }, - tokio_util::task::AbortOnDropHandle, -}; +#[cfg(all(windows, feature = "tun"))] +mod runtime; +#[cfg(all(windows, feature = "tun"))] +pub(crate) use runtime::start; -#[cfg(all(windows, any(target_arch = "x86_64", target_arch = "x86")))] -use windivert::{ - WinDivert, - error::WinDivertError, - layer, - packet::WinDivertPacket, - prelude::{WinDivertFlags, WinDivertShutdownMode}, -}; - -#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] -struct PhysicalInterface { - addr: Ipv4Addr, - directed_broadcast: Ipv4Addr, -} - -impl PhysicalInterface { - fn from_ip_and_prefix(addr: Ipv4Addr, prefix: u8) -> Option { - if should_ignore_interface_addr(addr) || prefix > 30 { - return None; - } - - Some(Self { - addr, - directed_broadcast: directed_broadcast(addr, prefix)?, - }) - } -} - -#[derive(Debug, Clone)] -struct BroadcastRelayConfig { - virtual_ipv4: Ipv4Inet, - physical_interfaces: Vec, -} - -impl BroadcastRelayConfig { - fn new(virtual_ipv4: Ipv4Inet, physical_interfaces: Vec) -> Self { - Self { - virtual_ipv4, - physical_interfaces, - } - } - - fn is_physical_source(&self, addr: Ipv4Addr) -> bool { - self.physical_interfaces - .iter() - .any(|iface| iface.addr == addr) - } - - fn normalize_destination(&self, dst: Ipv4Addr) -> Option { - if dst.is_broadcast() || dst.is_multicast() { - return Some(dst); - } - - self.physical_interfaces - .iter() - .any(|iface| iface.directed_broadcast == dst) - .then_some(self.virtual_ipv4.last_address()) - } -} - -#[derive(Debug, Clone, PartialEq, Eq)] -struct NormalizedPacket { - packet: Vec, - destination: Ipv4Addr, -} - -#[derive(Debug, Clone, Copy, PartialEq, Eq)] -struct UdpPacketSummary { - src: Ipv4Addr, - dst: Ipv4Addr, - src_port: u16, - dst_port: u16, - ip_len: usize, - udp_len: usize, - payload_len: usize, -} - -impl UdpPacketSummary { - fn parse(packet: &[u8]) -> Option { - let ipv4_packet = Ipv4Packet::new(packet)?; - if ipv4_packet.get_version() != 4 - || ipv4_packet.get_next_level_protocol() != IpNextHeaderProtocols::Udp - { - return None; - } - - let header_len = usize::from(ipv4_packet.get_header_length()) * 4; - let total_len = usize::from(ipv4_packet.get_total_length()); - if header_len < Ipv4Packet::minimum_packet_size() - || total_len < header_len + UdpPacket::minimum_packet_size() - || total_len > packet.len() - { - return None; - } - - let udp_packet = UdpPacket::new(&packet[header_len..total_len])?; - let udp_len = usize::from(udp_packet.get_length()); - if udp_len < UdpPacket::minimum_packet_size() || header_len + udp_len != total_len { - return None; - } - - Some(Self { - src: ipv4_packet.get_source(), - dst: ipv4_packet.get_destination(), - src_port: udp_packet.get_source(), - dst_port: udp_packet.get_destination(), - ip_len: total_len, - udp_len, - payload_len: udp_len - UdpPacket::minimum_packet_size(), - }) - } -} - -#[derive(Debug, Clone, Copy, PartialEq, Eq)] -struct ParsedUdpBroadcastPacket { - header_len: usize, - udp_len: usize, - normalized_destination: Ipv4Addr, - summary: UdpPacketSummary, -} - -#[cfg(any(windows, test))] -#[derive(Clone)] -struct BroadcastRelayStats { - packets_captured: CounterHandle, - packets_ignored: CounterHandle, - packets_forwarded: CounterHandle, - packets_forward_failed: CounterHandle, -} - -#[cfg(any(windows, test))] -impl BroadcastRelayStats { - fn new(peer_manager: &PeerManager) -> Self { - let global_ctx = peer_manager.get_global_ctx(); - let label_set = - LabelSet::new().with_label_type(LabelType::NetworkName(global_ctx.get_network_name())); - let stats_manager = global_ctx.stats_manager(); - - Self { - packets_captured: stats_manager.get_counter( - MetricName::UdpBroadcastRelayPacketsCaptured, - label_set.clone(), - ), - packets_ignored: stats_manager.get_counter( - MetricName::UdpBroadcastRelayPacketsIgnored, - label_set.clone(), - ), - packets_forwarded: stats_manager.get_counter( - MetricName::UdpBroadcastRelayPacketsForwarded, - label_set.clone(), - ), - packets_forward_failed: stats_manager - .get_counter(MetricName::UdpBroadcastRelayPacketsForwardFailed, label_set), - } - } - - fn record_captured(&self) { - self.packets_captured.inc(); - } - - fn record_ignored(&self) { - self.packets_ignored.inc(); - } - - fn record_forwarded(&self) { - self.packets_forwarded.inc(); - } - - fn record_forward_failed(&self) { - self.packets_forward_failed.inc(); - } -} - -fn should_ignore_interface_addr(addr: Ipv4Addr) -> bool { - addr.is_unspecified() || addr.is_loopback() || addr.is_multicast() || addr.is_broadcast() -} - -fn prefix_len_from_netmask(mask: Ipv4Addr) -> Option { - let raw = u32::from(mask); - let prefix = raw.count_ones() as u8; - let expected = if prefix == 0 { - 0 - } else { - u32::MAX << (32 - prefix) - }; - (raw == expected).then_some(prefix) -} - -fn directed_broadcast(addr: Ipv4Addr, prefix: u8) -> Option { - if prefix > 32 { - return None; - } - - let mask = if prefix == 0 { - 0 - } else { - u32::MAX << (32 - prefix) - }; - Some(Ipv4Addr::from(u32::from(addr) | !mask)) -} - -fn parse_udp_broadcast( - packet: &[u8], - config: &BroadcastRelayConfig, -) -> Result { - let ipv4_packet = Ipv4Packet::new(packet).ok_or("malformed_ipv4")?; - if ipv4_packet.get_version() != 4 - || ipv4_packet.get_next_level_protocol() != IpNextHeaderProtocols::Udp - { - return Err("not_udp_ipv4"); - } - - if ipv4_packet.get_fragment_offset() != 0 - || ipv4_packet.get_flags() & Ipv4Flags::MoreFragments != 0 - { - return Err("fragmented"); - } - - let header_len = usize::from(ipv4_packet.get_header_length()) * 4; - let total_len = usize::from(ipv4_packet.get_total_length()); - if header_len < Ipv4Packet::minimum_packet_size() - || total_len < header_len + UdpPacket::minimum_packet_size() - || total_len > packet.len() - { - return Err("bad_ipv4_length"); - } - - let src = ipv4_packet.get_source(); - let dst = ipv4_packet.get_destination(); - if should_ignore_interface_addr(src) { - return Err("ignored_source"); - } - if src == config.virtual_ipv4.address() { - return Err("virtual_source_duplicate"); - } - if !config.is_physical_source(src) { - return Err("non_physical_source"); - } - - let normalized_destination = config - .normalize_destination(dst) - .ok_or("unsupported_destination")?; - if normalized_destination.is_loopback() { - return Err("loopback_destination"); - } - - let udp_packet = UdpPacket::new(&packet[header_len..total_len]).ok_or("malformed_udp")?; - let udp_len = usize::from(udp_packet.get_length()); - if udp_len < UdpPacket::minimum_packet_size() || header_len + udp_len != total_len { - return Err("bad_udp_length"); - } - - Ok(ParsedUdpBroadcastPacket { - header_len, - udp_len, - normalized_destination, - summary: UdpPacketSummary { - src, - dst, - src_port: udp_packet.get_source(), - dst_port: udp_packet.get_destination(), - ip_len: total_len, - udp_len, - payload_len: udp_len - UdpPacket::minimum_packet_size(), - }, - }) -} - -fn log_ignored_udp_packet(packet: &[u8], reason: &'static str) { - if let Some(summary) = UdpPacketSummary::parse(packet) { - tracing::debug!( - src = %summary.src, - dst = %summary.dst, - src_port = summary.src_port, - dst_port = summary.dst_port, - ip_len = summary.ip_len, - udp_len = summary.udp_len, - payload_len = summary.payload_len, - reason, - "ignored Windows UDP broadcast packet" - ); - } else { - tracing::debug!( - packet_len = packet.len(), - reason, - "ignored malformed Windows UDP raw packet" - ); - } -} - -fn normalize_udp_broadcast_packet( - packet: &[u8], - config: &BroadcastRelayConfig, -) -> Option { - let parsed = match parse_udp_broadcast(packet, config) { - Ok(parsed) => parsed, - Err(reason) => { - if tracing::enabled!(tracing::Level::DEBUG) { - log_ignored_udp_packet(packet, reason); - } - return None; - } - }; - let header_len = parsed.header_len; - let udp_len = parsed.udp_len; - let destination = parsed.normalized_destination; - let summary = parsed.summary; - let packet_len = header_len + udp_len; - let virtual_ipv4 = config.virtual_ipv4.address(); - let mut normalized = packet[..packet_len].to_vec(); - - { - let mut ipv4_packet = MutableIpv4Packet::new(&mut normalized)?; - ipv4_packet.set_source(virtual_ipv4); - ipv4_packet.set_destination(destination); - ipv4_packet.set_total_length(packet_len as u16); - ipv4_packet.set_checksum(0); - } - - { - let mut udp_packet = MutableUdpPacket::new(&mut normalized[header_len..packet_len])?; - udp_packet.set_checksum(0); - let checksum = udp::ipv4_checksum(&udp_packet.to_immutable(), &virtual_ipv4, &destination); - udp_packet.set_checksum(checksum); - } - - { - let mut ipv4_packet = MutableIpv4Packet::new(&mut normalized)?; - let checksum = ipv4::checksum(&ipv4_packet.to_immutable()); - ipv4_packet.set_checksum(checksum); - } - - tracing::debug!( - src = %summary.src, - dst = %summary.dst, - src_port = summary.src_port, - dst_port = summary.dst_port, - ip_len = summary.ip_len, - udp_len = summary.udp_len, - payload_len = summary.payload_len, - normalized_src = %virtual_ipv4, - normalized_dst = %destination, - "normalized Windows UDP broadcast packet" - ); - - Some(NormalizedPacket { - packet: normalized, - destination, - }) -} - -#[cfg(any(windows, test))] -fn log_captured_udp_packet(packet: &[u8]) { - if let Some(summary) = UdpPacketSummary::parse(packet) { - tracing::debug!( - src = %summary.src, - dst = %summary.dst, - src_port = summary.src_port, - dst_port = summary.dst_port, - ip_len = summary.ip_len, - udp_len = summary.udp_len, - payload_len = summary.payload_len, - "captured Windows UDP broadcast candidate" - ); - } else { - tracing::debug!( - packet_len = packet.len(), - "captured malformed Windows UDP broadcast candidate" - ); - } -} - -#[cfg(any(windows, test))] -fn collect_physical_interfaces(virtual_ipv4: Ipv4Inet) -> anyhow::Result> { - let mut ret = Vec::new(); - for iface in NetworkInterface::show().context("failed to list Windows network interfaces")? { - if iface.internal { - continue; - } - - for addr in iface.addr { - let Addr::V4(v4) = addr else { - continue; - }; - if v4.ip == virtual_ipv4.address() { - continue; - } - - let Some(netmask) = v4.netmask else { - continue; - }; - let Some(prefix) = prefix_len_from_netmask(netmask) else { - tracing::debug!( - iface = %iface.name, - ip = %v4.ip, - mask = %netmask, - "ignoring interface with non-contiguous IPv4 netmask" - ); - continue; - }; - let Some(physical) = PhysicalInterface::from_ip_and_prefix(v4.ip, prefix) else { - continue; - }; - if !ret.contains(&physical) { - ret.push(physical); - } - } - } - Ok(ret) -} - -#[cfg(any(windows, test))] fn join_addr_equals(field: &str, addrs: &[Ipv4Addr]) -> String { addrs .iter() @@ -446,17 +15,16 @@ fn join_addr_equals(field: &str, addrs: &[Ipv4Addr]) -> String { .join(" or ") } -#[cfg(any(windows, test))] fn build_windivert_udp_filter(physical_interfaces: &[PhysicalInterface]) -> String { let mut src_addrs = Vec::new(); let mut directed_broadcasts = Vec::new(); for iface in physical_interfaces { - if !src_addrs.contains(&iface.addr) { - src_addrs.push(iface.addr); + if !src_addrs.contains(&iface.address()) { + src_addrs.push(iface.address()); } - if !directed_broadcasts.contains(&iface.directed_broadcast) { - directed_broadcasts.push(iface.directed_broadcast); + if !directed_broadcasts.contains(&iface.directed_broadcast()) { + directed_broadcasts.push(iface.directed_broadcast()); } } @@ -478,601 +46,9 @@ fn build_windivert_udp_filter(physical_interfaces: &[PhysicalInterface]) -> Stri ) } -#[cfg(any(windows, test))] -fn open_raw_udp_socket() -> io::Result { - let socket = Socket::new(Domain::IPV4, Type::RAW, Some(Protocol::UDP))?; - // Match ubihazard/broadcast: use one raw UDP listener on loopback, then - // inspect the IPv4 header to identify the real physical source interface. - socket.bind(&SockAddr::from(SocketAddrV4::new(Ipv4Addr::LOCALHOST, 0)))?; - socket.set_nonblocking(true)?; - Ok(socket) -} - -#[cfg(windows)] -fn socket2_into_udp_socket(socket: Socket) -> StdUdpSocket { - use std::os::windows::io::{FromRawSocket, IntoRawSocket}; - - // The raw socket handle came from socket2 and is transferred exactly once. - unsafe { StdUdpSocket::from_raw_socket(socket.into_raw_socket()) } -} - -#[cfg(all(not(windows), unix))] -fn socket2_into_udp_socket(socket: Socket) -> StdUdpSocket { - use std::os::fd::{FromRawFd, IntoRawFd}; - - // The raw socket fd came from socket2 and is transferred exactly once. - unsafe { StdUdpSocket::from_raw_fd(socket.into_raw_fd()) } -} - -#[cfg(any(windows, test))] -struct RawUdpCaptureSocket { - socket: tokio::net::UdpSocket, - buf: Vec, -} - -#[cfg(any(windows, test))] -impl RawUdpCaptureSocket { - const MAX_PACKET_LEN: usize = 65_535; - - fn open() -> anyhow::Result { - let socket = open_raw_udp_socket().with_context(|| { - "failed to open Windows raw UDP broadcast listener; administrator privileges are required" - })?; - let socket = socket2_into_udp_socket(socket); - let socket = tokio::net::UdpSocket::from_std(socket) - .context("failed to register Windows raw UDP broadcast listener with Tokio")?; - - Ok(Self { - socket, - buf: vec![0; Self::MAX_PACKET_LEN], - }) - } - - async fn recv(&mut self) -> io::Result<&[u8]> { - let len = self.socket.recv(&mut self.buf).await?; - Ok(&self.buf[..len]) - } -} - -#[cfg(all(windows, any(target_arch = "x86_64", target_arch = "x86")))] -struct WinDivertCaptureReader { - inner: std::cell::UnsafeCell>, -} - -#[cfg(all(windows, any(target_arch = "x86_64", target_arch = "x86")))] -unsafe impl Send for WinDivertCaptureReader {} - -#[cfg(all(windows, any(target_arch = "x86_64", target_arch = "x86")))] -unsafe impl Sync for WinDivertCaptureReader {} - -#[cfg(all(windows, any(target_arch = "x86_64", target_arch = "x86")))] -impl WinDivertCaptureReader { - fn new(inner: WinDivert) -> Self { - Self { - inner: std::cell::UnsafeCell::new(inner), - } - } - - fn recv<'a>( - &self, - buffer: Option<&'a mut [u8]>, - ) -> Result, WinDivertError> { - let inner = unsafe { &*self.inner.get() }; - inner.recv(buffer) - } - - fn shutdown(&self) -> anyhow::Result<()> { - let inner = unsafe { &mut *self.inner.get() }; - inner - .shutdown(WinDivertShutdownMode::Recv) - .with_context(|| "WinDivert UDP broadcast capture shutdown failed")?; - Ok(()) - } - - fn close(&self) -> anyhow::Result<()> { - let inner = unsafe { &mut *self.inner.get() }; - inner - .close(windivert::CloseAction::Nothing) - .with_context(|| "WinDivert UDP broadcast capture close failed")?; - Ok(()) - } -} - -#[cfg(all(windows, any(target_arch = "x86_64", target_arch = "x86")))] -impl Drop for WinDivertCaptureReader { - fn drop(&mut self) { - if let Err(err) = self.close() { - tracing::error!(?err, "WinDivert UDP broadcast capture close failed"); - } - } -} - -#[cfg(all(windows, any(target_arch = "x86_64", target_arch = "x86")))] -struct WinDivertCaptureSocket { - rx: tokio::sync::mpsc::Receiver>, - reader: Arc, - buf: Vec, -} - -#[cfg(all(windows, any(target_arch = "x86_64", target_arch = "x86")))] -impl WinDivertCaptureSocket { - const CHANNEL_CAPACITY: usize = 1024; - const MAX_PACKET_LEN: usize = 65_535; - - fn open(config: &BroadcastRelayConfig) -> anyhow::Result { - let filter = build_windivert_udp_filter(&config.physical_interfaces); - tracing::debug!( - filter = %filter, - "opening WinDivert UDP broadcast capture backend" - ); - - let flags = WinDivertFlags::default().set_sniff(); - let reader = WinDivert::network(&filter, 0, flags) - .map_err(io::Error::other) - .with_context(|| "failed to open WinDivert UDP broadcast capture")?; - let reader = Arc::new(WinDivertCaptureReader::new(reader)); - let reader_clone = reader.clone(); - let (tx, rx) = tokio::sync::mpsc::channel(Self::CHANNEL_CAPACITY); - - std::thread::Builder::new() - .name("easytier-udp-broadcast-windivert".to_owned()) - .spawn(move || { - let mut buffer = vec![0; Self::MAX_PACKET_LEN]; - loop { - match reader_clone.recv(Some(&mut buffer)) { - Ok(packet) => { - if tx.blocking_send(packet.data.to_vec()).is_err() { - break; - } - } - Err(err) => { - tracing::warn!(?err, "WinDivert UDP broadcast capture receive failed"); - break; - } - } - } - }) - .with_context(|| "failed to spawn WinDivert UDP broadcast capture thread")?; - - Ok(Self { - rx, - reader, - buf: Vec::new(), - }) - } - - async fn recv(&mut self) -> io::Result<&[u8]> { - self.buf = self.rx.recv().await.ok_or_else(|| { - io::Error::new( - io::ErrorKind::BrokenPipe, - "WinDivert UDP broadcast capture stopped", - ) - })?; - Ok(&self.buf) - } -} - -#[cfg(all(windows, any(target_arch = "x86_64", target_arch = "x86")))] -impl Drop for WinDivertCaptureSocket { - fn drop(&mut self) { - if let Err(err) = self.reader.shutdown() { - tracing::debug!(?err, "WinDivert UDP broadcast capture shutdown failed"); - } - } -} - -#[cfg(any(windows, test))] -enum CaptureSocket { - Raw(RawUdpCaptureSocket), - #[cfg(all(windows, any(target_arch = "x86_64", target_arch = "x86")))] - WinDivert(WinDivertCaptureSocket), -} - -#[cfg(any(windows, test))] -impl CaptureSocket { - async fn recv(&mut self) -> io::Result<&[u8]> { - match self { - Self::Raw(socket) => socket.recv().await, - #[cfg(all(windows, any(target_arch = "x86_64", target_arch = "x86")))] - Self::WinDivert(socket) => socket.recv().await, - } - } - - fn backend_name(&self) -> &'static str { - match self { - Self::Raw(_) => "raw_socket", - #[cfg(all(windows, any(target_arch = "x86_64", target_arch = "x86")))] - Self::WinDivert(_) => "windivert", - } - } - - fn fallback_to_raw(&mut self) -> anyhow::Result { - #[cfg(all(windows, any(target_arch = "x86_64", target_arch = "x86")))] - { - if matches!(self, Self::WinDivert(_)) { - *self = Self::Raw(RawUdpCaptureSocket::open()?); - return Ok(true); - } - } - - Ok(false) - } -} - -#[cfg(all(windows, any(target_arch = "x86_64", target_arch = "x86")))] -fn open_capture_socket(config: &BroadcastRelayConfig) -> anyhow::Result { - match WinDivertCaptureSocket::open(config) { - Ok(socket) => Ok(CaptureSocket::WinDivert(socket)), - Err(err) => { - tracing::warn!( - ?err, - "WinDivert UDP broadcast capture unavailable; falling back to raw socket" - ); - RawUdpCaptureSocket::open().map(CaptureSocket::Raw) - } - } -} - -#[cfg(all( - any(windows, test), - not(all(windows, any(target_arch = "x86_64", target_arch = "x86"))) -))] -fn open_capture_socket(_config: &BroadcastRelayConfig) -> anyhow::Result { - RawUdpCaptureSocket::open().map(CaptureSocket::Raw) -} - -#[cfg(any(windows, test))] -fn issue_start_result_event( - peer_manager: &PeerManager, - capture_backend: Option<&str>, - error: Option, -) { - peer_manager - .get_global_ctx() - .issue_event(GlobalCtxEvent::UdpBroadcastRelayStartResult { - capture_backend: capture_backend.map(str::to_owned), - error, - }); -} - -#[cfg(any(windows, test))] -async fn forward_normalized_packet( - peer_manager: &PeerManager, - normalized: NormalizedPacket, - stats: &BroadcastRelayStats, -) { - let packet = ZCPacket::new_with_payload(&normalized.packet); - let ret = peer_manager - .send_msg_by_ip(packet, IpAddr::V4(normalized.destination), true) - .await; - - let summary = UdpPacketSummary::parse(&normalized.packet); - match ret { - Ok(_) => { - stats.record_forwarded(); - - if let Some(summary) = summary { - tracing::debug!( - src = %summary.src, - dst = %summary.dst, - src_port = summary.src_port, - dst_port = summary.dst_port, - ip_len = summary.ip_len, - udp_len = summary.udp_len, - payload_len = summary.payload_len, - peer_dst = %normalized.destination, - broadcast = true, - "forwarded Windows UDP broadcast packet" - ); - } else { - tracing::debug!( - packet_len = normalized.packet.len(), - peer_dst = %normalized.destination, - broadcast = true, - "forwarded Windows UDP broadcast packet" - ); - } - } - Err(err) => { - stats.record_forward_failed(); - - if let Some(summary) = summary { - tracing::debug!( - src = %summary.src, - dst = %summary.dst, - src_port = summary.src_port, - dst_port = summary.dst_port, - ip_len = summary.ip_len, - udp_len = summary.udp_len, - payload_len = summary.payload_len, - peer_dst = %normalized.destination, - broadcast = true, - ?err, - "failed to forward Windows UDP broadcast packet" - ); - } else { - tracing::debug!( - packet_len = normalized.packet.len(), - peer_dst = %normalized.destination, - broadcast = true, - ?err, - "failed to forward Windows UDP broadcast packet" - ); - } - } - } -} - -#[cfg(any(windows, test))] -async fn capture_loop( - peer_manager: Arc, - config: BroadcastRelayConfig, - mut socket: CaptureSocket, - stats: BroadcastRelayStats, -) { - let mut capture_backend = socket.backend_name(); - - loop { - let normalized = match socket.recv().await { - Ok(packet) => { - stats.record_captured(); - if tracing::enabled!(tracing::Level::DEBUG) { - log_captured_udp_packet(packet); - } - let normalized = normalize_udp_broadcast_packet(packet, &config); - if normalized.is_none() { - stats.record_ignored(); - } - normalized - } - Err(err) => { - tracing::warn!( - ?err, - capture_backend, - "Windows UDP broadcast capture receive failed" - ); - match socket.fallback_to_raw() { - Ok(true) => { - let old_backend = capture_backend; - capture_backend = socket.backend_name(); - tracing::warn!( - old_backend, - new_backend = capture_backend, - "Windows UDP broadcast capture backend fell back" - ); - } - Ok(false) => {} - Err(fallback_err) => { - tracing::error!( - ?fallback_err, - "Windows UDP broadcast raw socket fallback failed; stopping relay" - ); - break; - } - } - continue; - } - }; - - if let Some(normalized) = normalized { - forward_normalized_packet(&peer_manager, normalized, &stats).await; - } - } -} - -#[cfg(any(windows, test))] -pub(crate) fn start( - peer_manager: Arc, - virtual_ipv4: Ipv4Inet, -) -> anyhow::Result> { - let physical_interfaces = match collect_physical_interfaces(virtual_ipv4) { - Ok(interfaces) => interfaces, - Err(err) => { - issue_start_result_event(&peer_manager, None, Some(format!("{err:#}"))); - return Err(err); - } - }; - if physical_interfaces.is_empty() { - let msg = "no physical IPv4 interface is available for UDP broadcast relay"; - issue_start_result_event(&peer_manager, None, Some(msg.to_owned())); - anyhow::bail!(msg); - } - - let config = BroadcastRelayConfig::new(virtual_ipv4, physical_interfaces); - let socket = match open_capture_socket(&config) { - Ok(socket) => socket, - Err(err) => { - issue_start_result_event(&peer_manager, None, Some(format!("{err:#}"))); - return Err(err); - } - }; - let capture_backend = socket.backend_name(); - issue_start_result_event(&peer_manager, Some(capture_backend), None); - - tracing::debug!( - virtual_ipv4 = %config.virtual_ipv4, - physical_interfaces = ?config.physical_interfaces, - capture_backend, - "starting Windows UDP broadcast relay" - ); - - let stats = BroadcastRelayStats::new(&peer_manager); - let task = tokio::spawn(capture_loop(peer_manager, config, socket, stats)); - Ok(AbortOnDropHandle::new(task)) -} - #[cfg(test)] mod tests { use super::*; - use pnet::packet::{MutablePacket, Packet}; - - fn config() -> BroadcastRelayConfig { - BroadcastRelayConfig::new( - "10.144.144.1/24".parse().unwrap(), - vec![PhysicalInterface::from_ip_and_prefix(Ipv4Addr::new(192, 168, 1, 7), 24).unwrap()], - ) - } - - fn build_udp_packet(src: Ipv4Addr, dst: Ipv4Addr, payload: &[u8]) -> Vec { - let mut packet = vec![0; 20 + 8 + payload.len()]; - { - let mut ipv4_packet = MutableIpv4Packet::new(&mut packet).unwrap(); - ipv4_packet.set_version(4); - ipv4_packet.set_header_length(5); - ipv4_packet.set_total_length((20 + 8 + payload.len()) as u16); - ipv4_packet.set_ttl(64); - ipv4_packet.set_next_level_protocol(IpNextHeaderProtocols::Udp); - ipv4_packet.set_source(src); - ipv4_packet.set_destination(dst); - } - - { - let mut udp_packet = MutableUdpPacket::new(&mut packet[20..]).unwrap(); - udp_packet.set_source(12345); - udp_packet.set_destination(37020); - udp_packet.set_length((8 + payload.len()) as u16); - udp_packet.payload_mut().copy_from_slice(payload); - let checksum = udp::ipv4_checksum(&udp_packet.to_immutable(), &src, &dst); - udp_packet.set_checksum(checksum); - } - - { - let mut ipv4_packet = MutableIpv4Packet::new(&mut packet).unwrap(); - let checksum = ipv4::checksum(&ipv4_packet.to_immutable()); - ipv4_packet.set_checksum(checksum); - } - - packet - } - - fn assert_valid_checksums(packet: &[u8]) { - let ipv4_packet = Ipv4Packet::new(packet).unwrap(); - assert_eq!(ipv4::checksum(&ipv4_packet), ipv4_packet.get_checksum()); - let udp_packet = UdpPacket::new(ipv4_packet.payload()).unwrap(); - assert_eq!( - udp::ipv4_checksum( - &udp_packet, - &ipv4_packet.get_source(), - &ipv4_packet.get_destination() - ), - udp_packet.get_checksum() - ); - } - - #[test] - fn windows_udp_broadcast_rewrites_limited_broadcast() { - let packet = build_udp_packet(Ipv4Addr::new(192, 168, 1, 7), Ipv4Addr::BROADCAST, b"hello"); - - let normalized = normalize_udp_broadcast_packet(&packet, &config()).unwrap(); - let ipv4_packet = Ipv4Packet::new(&normalized.packet).unwrap(); - - assert_eq!(normalized.destination, Ipv4Addr::BROADCAST); - assert_eq!(ipv4_packet.get_source(), Ipv4Addr::new(10, 144, 144, 1)); - assert_eq!(ipv4_packet.get_destination(), Ipv4Addr::BROADCAST); - assert_eq!(&ipv4_packet.payload()[8..], b"hello"); - assert_valid_checksums(&normalized.packet); - } - - #[test] - fn windows_udp_broadcast_rewrites_directed_broadcast() { - let packet = build_udp_packet( - Ipv4Addr::new(192, 168, 1, 7), - Ipv4Addr::new(192, 168, 1, 255), - b"directed", - ); - - let normalized = normalize_udp_broadcast_packet(&packet, &config()).unwrap(); - let ipv4_packet = Ipv4Packet::new(&normalized.packet).unwrap(); - - assert_eq!(normalized.destination, Ipv4Addr::new(10, 144, 144, 255)); - assert_eq!(ipv4_packet.get_source(), Ipv4Addr::new(10, 144, 144, 1)); - assert_eq!( - ipv4_packet.get_destination(), - Ipv4Addr::new(10, 144, 144, 255) - ); - assert_eq!(&ipv4_packet.payload()[8..], b"directed"); - assert_valid_checksums(&normalized.packet); - } - - #[test] - fn windows_udp_broadcast_preserves_multicast_destination() { - let multicast = Ipv4Addr::new(239, 255, 255, 250); - let packet = build_udp_packet(Ipv4Addr::new(192, 168, 1, 7), multicast, b"multicast"); - - let normalized = normalize_udp_broadcast_packet(&packet, &config()).unwrap(); - let ipv4_packet = Ipv4Packet::new(&normalized.packet).unwrap(); - - assert_eq!(normalized.destination, multicast); - assert_eq!(ipv4_packet.get_source(), Ipv4Addr::new(10, 144, 144, 1)); - assert_eq!(ipv4_packet.get_destination(), multicast); - assert_eq!(&ipv4_packet.payload()[8..], b"multicast"); - assert_valid_checksums(&normalized.packet); - } - - #[test] - fn windows_udp_broadcast_rejects_malformed_packets() { - assert!(normalize_udp_broadcast_packet(&[], &config()).is_none()); - - let mut packet = - build_udp_packet(Ipv4Addr::new(192, 168, 1, 7), Ipv4Addr::BROADCAST, b"bad"); - packet[2..4].copy_from_slice(&10u16.to_be_bytes()); - assert!(normalize_udp_broadcast_packet(&packet, &config()).is_none()); - } - - #[test] - fn windows_udp_broadcast_rejects_fragments() { - let mut packet = build_udp_packet( - Ipv4Addr::new(192, 168, 1, 7), - Ipv4Addr::BROADCAST, - b"fragment", - ); - { - let mut ipv4_packet = MutableIpv4Packet::new(&mut packet).unwrap(); - ipv4_packet.set_flags(Ipv4Flags::MoreFragments); - } - - assert!(normalize_udp_broadcast_packet(&packet, &config()).is_none()); - } - - #[test] - fn windows_udp_broadcast_rejects_non_broadcast_destinations() { - let packet = build_udp_packet( - Ipv4Addr::new(192, 168, 1, 7), - Ipv4Addr::new(192, 168, 1, 10), - b"unicast", - ); - - assert!(normalize_udp_broadcast_packet(&packet, &config()).is_none()); - } - - #[test] - fn windows_udp_broadcast_rejects_virtual_source_duplicates() { - let packet = build_udp_packet(Ipv4Addr::new(10, 144, 144, 1), Ipv4Addr::BROADCAST, b"loop"); - - assert!(normalize_udp_broadcast_packet(&packet, &config()).is_none()); - } - - #[test] - fn windows_udp_broadcast_detects_directed_broadcast_from_prefix() { - let physical = - PhysicalInterface::from_ip_and_prefix(Ipv4Addr::new(172, 16, 5, 10), 20).unwrap(); - assert_eq!(physical.directed_broadcast, Ipv4Addr::new(172, 16, 15, 255)); - assert_eq!( - prefix_len_from_netmask(Ipv4Addr::new(255, 255, 240, 0)), - Some(20) - ); - assert_eq!(prefix_len_from_netmask(Ipv4Addr::new(255, 0, 255, 0)), None); - } - - #[test] - fn windows_udp_broadcast_keeps_link_local_interfaces() { - let physical = - PhysicalInterface::from_ip_and_prefix(Ipv4Addr::new(169, 254, 13, 10), 16).unwrap(); - assert_eq!( - physical.directed_broadcast, - Ipv4Addr::new(169, 254, 255, 255) - ); - } #[test] fn windows_udp_broadcast_windivert_filter_is_constrained() { diff --git a/easytier/src/instance/windows_udp_broadcast/capture_raw.rs b/easytier/src/instance/windows_udp_broadcast/capture_raw.rs new file mode 100644 index 00000000..90758b03 --- /dev/null +++ b/easytier/src/instance/windows_udp_broadcast/capture_raw.rs @@ -0,0 +1,21 @@ +use super::{BroadcastRelayConfig, RawUdpCaptureSocket}; + +pub(super) struct CaptureSocket(RawUdpCaptureSocket); + +impl CaptureSocket { + pub(super) async fn recv(&mut self) -> std::io::Result<&[u8]> { + self.0.recv().await + } + + pub(super) fn backend_name(&self) -> &'static str { + "raw_socket" + } + + pub(super) fn fallback_to_raw(&mut self) -> anyhow::Result { + Ok(false) + } +} + +pub(super) fn open_capture_socket(_config: &BroadcastRelayConfig) -> anyhow::Result { + RawUdpCaptureSocket::open().map(CaptureSocket) +} diff --git a/easytier/src/instance/windows_udp_broadcast/capture_windivert.rs b/easytier/src/instance/windows_udp_broadcast/capture_windivert.rs new file mode 100644 index 00000000..ebb074a1 --- /dev/null +++ b/easytier/src/instance/windows_udp_broadcast/capture_windivert.rs @@ -0,0 +1,173 @@ +use std::{io, sync::Arc}; + +use anyhow::Context; +use windivert::{ + WinDivert, + error::WinDivertError, + layer, + packet::WinDivertPacket, + prelude::{WinDivertFlags, WinDivertShutdownMode}, +}; + +use super::{BroadcastRelayConfig, RawUdpCaptureSocket}; +use crate::instance::windows_udp_broadcast::build_windivert_udp_filter; + +struct WinDivertCaptureReader { + inner: std::cell::UnsafeCell>, +} + +unsafe impl Send for WinDivertCaptureReader {} +unsafe impl Sync for WinDivertCaptureReader {} + +impl WinDivertCaptureReader { + fn new(inner: WinDivert) -> Self { + Self { + inner: std::cell::UnsafeCell::new(inner), + } + } + + fn recv<'a>( + &self, + buffer: Option<&'a mut [u8]>, + ) -> Result, WinDivertError> { + let inner = unsafe { &*self.inner.get() }; + inner.recv(buffer) + } + + fn shutdown(&self) -> anyhow::Result<()> { + let inner = unsafe { &mut *self.inner.get() }; + inner + .shutdown(WinDivertShutdownMode::Recv) + .with_context(|| "WinDivert UDP broadcast capture shutdown failed")?; + Ok(()) + } + + fn close(&self) -> anyhow::Result<()> { + let inner = unsafe { &mut *self.inner.get() }; + inner + .close(windivert::CloseAction::Nothing) + .with_context(|| "WinDivert UDP broadcast capture close failed")?; + Ok(()) + } +} + +impl Drop for WinDivertCaptureReader { + fn drop(&mut self) { + if let Err(err) = self.close() { + tracing::error!(?err, "WinDivert UDP broadcast capture close failed"); + } + } +} + +pub(super) struct WinDivertCaptureSocket { + rx: tokio::sync::mpsc::Receiver>, + reader: Arc, + buf: Vec, +} + +impl WinDivertCaptureSocket { + const CHANNEL_CAPACITY: usize = 1024; + const MAX_PACKET_LEN: usize = 65_535; + + fn open(config: &BroadcastRelayConfig) -> anyhow::Result { + let filter = build_windivert_udp_filter(config.physical_interfaces()); + tracing::debug!( + filter = %filter, + "opening WinDivert UDP broadcast capture backend" + ); + + let flags = WinDivertFlags::default().set_sniff(); + let reader = WinDivert::network(&filter, 0, flags) + .map_err(io::Error::other) + .with_context(|| "failed to open WinDivert UDP broadcast capture")?; + let reader = Arc::new(WinDivertCaptureReader::new(reader)); + let reader_clone = reader.clone(); + let (tx, rx) = tokio::sync::mpsc::channel(Self::CHANNEL_CAPACITY); + + std::thread::Builder::new() + .name("easytier-udp-broadcast-windivert".to_owned()) + .spawn(move || { + let mut buffer = vec![0; Self::MAX_PACKET_LEN]; + loop { + match reader_clone.recv(Some(&mut buffer)) { + Ok(packet) => { + if tx.blocking_send(packet.data.to_vec()).is_err() { + break; + } + } + Err(err) => { + tracing::warn!(?err, "WinDivert UDP broadcast capture receive failed"); + break; + } + } + } + }) + .with_context(|| "failed to spawn WinDivert UDP broadcast capture thread")?; + + Ok(Self { + rx, + reader, + buf: Vec::new(), + }) + } + + async fn recv(&mut self) -> io::Result<&[u8]> { + self.buf = self.rx.recv().await.ok_or_else(|| { + io::Error::new( + io::ErrorKind::BrokenPipe, + "WinDivert UDP broadcast capture stopped", + ) + })?; + Ok(&self.buf) + } +} + +impl Drop for WinDivertCaptureSocket { + fn drop(&mut self) { + if let Err(err) = self.reader.shutdown() { + tracing::debug!(?err, "WinDivert UDP broadcast capture shutdown failed"); + } + } +} + +pub(super) enum CaptureSocket { + Raw(RawUdpCaptureSocket), + WinDivert(WinDivertCaptureSocket), +} + +impl CaptureSocket { + pub(super) async fn recv(&mut self) -> io::Result<&[u8]> { + match self { + Self::Raw(socket) => socket.recv().await, + Self::WinDivert(socket) => socket.recv().await, + } + } + + pub(super) fn backend_name(&self) -> &'static str { + match self { + Self::Raw(_) => "raw_socket", + Self::WinDivert(_) => "windivert", + } + } + + pub(super) fn fallback_to_raw(&mut self) -> anyhow::Result { + if matches!(self, Self::WinDivert(_)) { + *self = Self::Raw(RawUdpCaptureSocket::open()?); + return Ok(true); + } + Ok(false) + } +} + +pub(super) fn open_capture_socket(config: &BroadcastRelayConfig) -> anyhow::Result { + match WinDivertCaptureSocket::open(config) { + Ok(socket) => Ok(CaptureSocket::WinDivert(socket)), + Err(err) => { + tracing::warn!( + ?err, + "WinDivert UDP broadcast capture unavailable; falling back to raw socket" + ); + RawUdpCaptureSocket::open().map(CaptureSocket::Raw) + } + } +} diff --git a/easytier/src/instance/windows_udp_broadcast/runtime.rs b/easytier/src/instance/windows_udp_broadcast/runtime.rs new file mode 100644 index 00000000..883ddcb4 --- /dev/null +++ b/easytier/src/instance/windows_udp_broadcast/runtime.rs @@ -0,0 +1,359 @@ +use std::net::Ipv4Addr; + +use cidr::Ipv4Inet; +use easytier_core::gateway::udp_broadcast::PhysicalInterface; +use easytier_core::gateway::udp_broadcast::{ + BroadcastRelayConfig, NormalizedPacket, UdpBroadcastPacketRejection, UdpPacketSummary, + normalize_udp_broadcast_packet, +}; + +use { + crate::common::global_ctx::{ArcGlobalCtx, GlobalCtxEvent}, + anyhow::Context, + easytier_core::{gateway::udp_broadcast::UdpBroadcastRelayStats, instance::CorePacketPlane}, + network_interface::{Addr, NetworkInterface, NetworkInterfaceConfig}, + socket2::{Domain, Protocol, SockAddr, Socket, Type}, + std::{ + io, + net::{SocketAddrV4, UdpSocket as StdUdpSocket}, + sync::Arc, + }, + tokio_util::task::AbortOnDropHandle, +}; + +#[cfg(not(any(target_arch = "x86_64", target_arch = "x86")))] +#[path = "capture_raw.rs"] +mod capture; +#[cfg(any(target_arch = "x86_64", target_arch = "x86"))] +#[path = "capture_windivert.rs"] +mod capture; + +use capture::{CaptureSocket, open_capture_socket}; + +fn log_ignored_udp_packet(packet: &[u8], rejection: UdpBroadcastPacketRejection) { + let reason = rejection.reason(); + if let Some(summary) = UdpPacketSummary::parse(packet) { + tracing::debug!( + src = %summary.src, + dst = %summary.dst, + src_port = summary.src_port, + dst_port = summary.dst_port, + ip_len = summary.ip_len, + udp_len = summary.udp_len, + payload_len = summary.payload_len, + reason, + "ignored Windows UDP broadcast packet" + ); + } else { + tracing::debug!( + packet_len = packet.len(), + reason, + "ignored malformed Windows UDP raw packet" + ); + } +} + +fn log_normalized_udp_packet( + packet: &[u8], + config: &BroadcastRelayConfig, + normalized: &NormalizedPacket, +) { + let Some(summary) = UdpPacketSummary::parse(packet) else { + return; + }; + + tracing::debug!( + src = %summary.src, + dst = %summary.dst, + src_port = summary.src_port, + dst_port = summary.dst_port, + ip_len = summary.ip_len, + udp_len = summary.udp_len, + payload_len = summary.payload_len, + normalized_src = %config.virtual_ipv4().address(), + normalized_dst = %normalized.destination, + "normalized Windows UDP broadcast packet" + ); +} + +fn log_captured_udp_packet(packet: &[u8]) { + if let Some(summary) = UdpPacketSummary::parse(packet) { + tracing::debug!( + src = %summary.src, + dst = %summary.dst, + src_port = summary.src_port, + dst_port = summary.dst_port, + ip_len = summary.ip_len, + udp_len = summary.udp_len, + payload_len = summary.payload_len, + "captured Windows UDP broadcast candidate" + ); + } else { + tracing::debug!( + packet_len = packet.len(), + "captured malformed Windows UDP broadcast candidate" + ); + } +} + +fn collect_physical_interfaces(virtual_ipv4: Ipv4Inet) -> anyhow::Result> { + let mut ret = Vec::new(); + for iface in NetworkInterface::show().context("failed to list Windows network interfaces")? { + for addr in iface.addr { + let Addr::V4(v4) = addr else { + continue; + }; + let physical = match PhysicalInterface::from_observation( + v4.ip, + v4.netmask, + iface.internal, + virtual_ipv4.address(), + ) { + Ok(Some(physical)) => physical, + Ok(None) => continue, + Err(non_contiguous) => { + tracing::debug!( + iface = %iface.name, + ip = %v4.ip, + mask = %non_contiguous.netmask(), + "ignoring interface with non-contiguous IPv4 netmask" + ); + continue; + } + }; + ret.push(physical); + } + } + Ok(ret) +} + +fn open_raw_udp_socket() -> io::Result { + let socket = Socket::new(Domain::IPV4, Type::RAW, Some(Protocol::UDP))?; + // Match ubihazard/broadcast: use one raw UDP listener on loopback, then + // inspect the IPv4 header to identify the real physical source interface. + socket.bind(&SockAddr::from(SocketAddrV4::new(Ipv4Addr::LOCALHOST, 0)))?; + socket.set_nonblocking(true)?; + Ok(socket) +} + +fn socket2_into_udp_socket(socket: Socket) -> StdUdpSocket { + use std::os::windows::io::{FromRawSocket, IntoRawSocket}; + + // The raw socket handle came from socket2 and is transferred exactly once. + unsafe { StdUdpSocket::from_raw_socket(socket.into_raw_socket()) } +} + +struct RawUdpCaptureSocket { + socket: tokio::net::UdpSocket, + buf: Vec, +} + +impl RawUdpCaptureSocket { + const MAX_PACKET_LEN: usize = 65_535; + + fn open() -> anyhow::Result { + let socket = open_raw_udp_socket().with_context(|| { + "failed to open Windows raw UDP broadcast listener; administrator privileges are required" + })?; + let socket = socket2_into_udp_socket(socket); + let socket = tokio::net::UdpSocket::from_std(socket) + .context("failed to register Windows raw UDP broadcast listener with Tokio")?; + + Ok(Self { + socket, + buf: vec![0; Self::MAX_PACKET_LEN], + }) + } + + async fn recv(&mut self) -> io::Result<&[u8]> { + let len = self.socket.recv(&mut self.buf).await?; + Ok(&self.buf[..len]) + } +} + +fn issue_start_result_event( + global_ctx: &ArcGlobalCtx, + capture_backend: Option<&str>, + error: Option, +) { + global_ctx.issue_event(GlobalCtxEvent::UdpBroadcastRelayStartResult { + capture_backend: capture_backend.map(str::to_owned), + error, + }); +} + +async fn forward_normalized_packet( + packet_plane: &CorePacketPlane, + normalized: NormalizedPacket, + stats: &UdpBroadcastRelayStats, +) { + let ret = packet_plane + .send_local_ip_packet(normalized.packet.clone()) + .await; + + let summary = UdpPacketSummary::parse(&normalized.packet); + match ret { + Ok(_) => { + stats.record_forwarded(); + + if let Some(summary) = summary { + tracing::debug!( + src = %summary.src, + dst = %summary.dst, + src_port = summary.src_port, + dst_port = summary.dst_port, + ip_len = summary.ip_len, + udp_len = summary.udp_len, + payload_len = summary.payload_len, + peer_dst = %normalized.destination, + broadcast = true, + "forwarded Windows UDP broadcast packet" + ); + } else { + tracing::debug!( + packet_len = normalized.packet.len(), + peer_dst = %normalized.destination, + broadcast = true, + "forwarded Windows UDP broadcast packet" + ); + } + } + Err(err) => { + stats.record_forward_failed(); + + if let Some(summary) = summary { + tracing::debug!( + src = %summary.src, + dst = %summary.dst, + src_port = summary.src_port, + dst_port = summary.dst_port, + ip_len = summary.ip_len, + udp_len = summary.udp_len, + payload_len = summary.payload_len, + peer_dst = %normalized.destination, + broadcast = true, + ?err, + "failed to forward Windows UDP broadcast packet" + ); + } else { + tracing::debug!( + packet_len = normalized.packet.len(), + peer_dst = %normalized.destination, + broadcast = true, + ?err, + "failed to forward Windows UDP broadcast packet" + ); + } + } + } +} + +async fn capture_loop( + packet_plane: Arc, + config: BroadcastRelayConfig, + mut socket: CaptureSocket, + stats: UdpBroadcastRelayStats, +) { + let mut capture_backend = socket.backend_name(); + + loop { + let normalized = match socket.recv().await { + Ok(packet) => { + stats.record_captured(); + if tracing::enabled!(tracing::Level::DEBUG) { + log_captured_udp_packet(packet); + } + let normalized = match normalize_udp_broadcast_packet(packet, &config) { + Ok(normalized) => { + if tracing::enabled!(tracing::Level::DEBUG) { + log_normalized_udp_packet(packet, &config, &normalized); + } + Some(normalized) + } + Err(rejection) => { + if tracing::enabled!(tracing::Level::DEBUG) { + log_ignored_udp_packet(packet, rejection); + } + None + } + }; + if normalized.is_none() { + stats.record_ignored(); + } + normalized + } + Err(err) => { + tracing::warn!( + ?err, + capture_backend, + "Windows UDP broadcast capture receive failed" + ); + match socket.fallback_to_raw() { + Ok(true) => { + let old_backend = capture_backend; + capture_backend = socket.backend_name(); + tracing::warn!( + old_backend, + new_backend = capture_backend, + "Windows UDP broadcast capture backend fell back" + ); + } + Ok(false) => {} + Err(fallback_err) => { + tracing::error!( + ?fallback_err, + "Windows UDP broadcast raw socket fallback failed; stopping relay" + ); + break; + } + } + continue; + } + }; + + if let Some(normalized) = normalized { + forward_normalized_packet(&packet_plane, normalized, &stats).await; + } + } +} + +pub(crate) fn start( + packet_plane: Arc, + global_ctx: ArcGlobalCtx, + virtual_ipv4: Ipv4Inet, +) -> anyhow::Result> { + let physical_interfaces = match collect_physical_interfaces(virtual_ipv4) { + Ok(interfaces) => interfaces, + Err(err) => { + issue_start_result_event(&global_ctx, None, Some(format!("{err:#}"))); + return Err(err); + } + }; + let config = BroadcastRelayConfig::new(virtual_ipv4, physical_interfaces); + if config.physical_interfaces().is_empty() { + let msg = "no physical IPv4 interface is available for UDP broadcast relay"; + issue_start_result_event(&global_ctx, None, Some(msg.to_owned())); + anyhow::bail!(msg); + } + + let socket = match open_capture_socket(&config) { + Ok(socket) => socket, + Err(err) => { + issue_start_result_event(&global_ctx, None, Some(format!("{err:#}"))); + return Err(err); + } + }; + let capture_backend = socket.backend_name(); + issue_start_result_event(&global_ctx, Some(capture_backend), None); + + tracing::debug!( + virtual_ipv4 = %config.virtual_ipv4(), + physical_interfaces = ?config.physical_interfaces(), + capture_backend, + "starting Windows UDP broadcast relay" + ); + + let stats = packet_plane.udp_broadcast_relay_stats(); + let task = tokio::spawn(capture_loop(packet_plane, config, socket, stats)); + Ok(AbortOnDropHandle::new(task)) +} diff --git a/easytier/src/instance_manager.rs b/easytier/src/instance_manager.rs deleted file mode 100644 index a28f310a..00000000 --- a/easytier/src/instance_manager.rs +++ /dev/null @@ -1,848 +0,0 @@ -#[cfg(feature = "ffi-dataplane")] -use crate::launcher::{DataPlaneTcpListener, DataPlaneTcpStream, DataPlaneUdpSocket}; -use dashmap::DashMap; -use std::fmt::{Display, Formatter}; -use std::{collections::BTreeMap, path::PathBuf, sync::Arc}; -use tokio_util::task::AbortOnDropHandle; - -use crate::{ - common::{ - config::{ConfigFileControl, ConfigLoader, ConfigSource, TomlConfigLoader}, - global_ctx::{EventBusSubscriber, GlobalCtxEvent}, - log, - }, - launcher::{NetworkInstance, NetworkInstanceRunningInfo}, - proto::{self}, - rpc_service::InstanceRpcService, -}; - -pub(crate) struct DaemonGuard { - guard: Option>, - stop_check_notifier: Arc, -} -impl Drop for DaemonGuard { - fn drop(&mut self) { - drop(self.guard.take()); - self.stop_check_notifier.notify_one(); - } -} - -pub struct NetworkInstanceManager { - instance_map: Arc>, - instance_stop_tasks: Arc>>, - stop_check_notifier: Arc, - instance_error_messages: Arc>, - config_dir: Option, - guard_counter: Arc<()>, - remote_mutation_lock: Arc>, -} - -impl Default for NetworkInstanceManager { - fn default() -> Self { - Self::new() - } -} - -impl NetworkInstanceManager { - pub fn new() -> Self { - NetworkInstanceManager { - instance_map: Arc::new(DashMap::new()), - instance_stop_tasks: Arc::new(DashMap::new()), - stop_check_notifier: Arc::new(tokio::sync::Notify::new()), - instance_error_messages: Arc::new(DashMap::new()), - config_dir: None, - guard_counter: Arc::new(()), - remote_mutation_lock: Arc::new(tokio::sync::Mutex::new(())), - } - } - - pub fn with_config_path(mut self, config_dir: Option) -> Self { - self.config_dir = config_dir; - self - } - - pub fn remote_mutation_lock(&self) -> Arc> { - self.remote_mutation_lock.clone() - } - - fn start_instance_task(&self, instance_id: uuid::Uuid) -> Result<(), anyhow::Error> { - if tokio::runtime::Handle::try_current().is_err() { - return Err(anyhow::anyhow!( - "tokio runtime not found, cannot start instance task" - )); - } - - let instance = self - .instance_map - .get(&instance_id) - .ok_or_else(|| anyhow::anyhow!("instance {} not found", instance_id))?; - let instance_stop_notifier = instance.get_stop_notifier(); - let instance_event_receiver = instance.subscribe_event(); - - let instance_map = self.instance_map.clone(); - let instance_stop_tasks = self.instance_stop_tasks.clone(); - let instance_error_messages = self.instance_error_messages.clone(); - - let stop_check_notifier = self.stop_check_notifier.clone(); - self.instance_stop_tasks.insert( - instance_id, - AbortOnDropHandle::new(tokio::spawn(async move { - let Some(instance_stop_notifier) = instance_stop_notifier else { - return; - }; - let _t = instance_event_receiver - .map(|event| AbortOnDropHandle::new(handle_event(instance_id, event))); - instance_stop_notifier.notified().await; - if let Some(instance) = instance_map.get(&instance_id) - && let Some(error) = instance.get_latest_error_msg() - { - log::error!(%error, "instance {} stopped", instance_id); - instance_error_messages.insert(instance_id, error); - } - stop_check_notifier.notify_one(); - instance_stop_tasks.remove(&instance_id); - instance_stop_tasks.shrink_to_fit(); - })), - ); - Ok(()) - } - - pub fn run_network_instance( - &self, - cfg: TomlConfigLoader, - watch_event: bool, - config_file_control: ConfigFileControl, - ) -> Result { - let instance_id = cfg.get_id(); - if self.instance_map.contains_key(&instance_id) { - anyhow::bail!("instance {} already exists", instance_id); - } - - let mut instance = NetworkInstance::new(cfg, config_file_control); - instance.start()?; - - self.instance_map.insert(instance_id, instance); - if watch_event { - self.start_instance_task(instance_id)?; - } - Ok(instance_id) - } - - pub fn retain_network_instance( - &self, - instance_ids: Vec, - ) -> Result, anyhow::Error> { - self.instance_map.retain(|k, _| instance_ids.contains(k)); - self.instance_map.shrink_to_fit(); - self.instance_error_messages - .retain(|k, _| instance_ids.contains(k)); - self.instance_error_messages.shrink_to_fit(); - Ok(self.list_network_instance_ids()) - } - - pub fn delete_network_instance( - &self, - instance_ids: Vec, - ) -> Result, anyhow::Error> { - self.instance_map.retain(|k, _| !instance_ids.contains(k)); - self.instance_map.shrink_to_fit(); - self.instance_error_messages - .retain(|k, _| !instance_ids.contains(k)); - self.instance_error_messages.shrink_to_fit(); - Ok(self.list_network_instance_ids()) - } - - pub async fn collect_network_infos( - &self, - ) -> Result, anyhow::Error> { - let mut ret = BTreeMap::new(); - for instance in self.instance_map.iter() { - if let Ok(info) = instance.get_running_info().await { - ret.insert(*instance.key(), info); - } - } - for v in self.instance_error_messages.iter() { - ret.insert( - *v.key(), - NetworkInstanceRunningInfo { - error_msg: Some(v.value().clone()), - ..Default::default() - }, - ); - } - Ok(ret) - } - - pub fn collect_network_infos_sync( - &self, - ) -> Result, anyhow::Error> { - tokio::runtime::Runtime::new()?.block_on(self.collect_network_infos()) - } - - #[cfg(feature = "ffi-dataplane")] - pub async fn data_plane_tcp_connect( - &self, - instance_id: &uuid::Uuid, - dst_addr: std::net::SocketAddr, - timeout: std::time::Duration, - ) -> Result { - let instance = self - .instance_map - .get(instance_id) - .ok_or_else(|| anyhow::anyhow!("instance {} not found", instance_id))?; - instance.data_plane_tcp_connect(dst_addr, timeout).await - } - - #[cfg(feature = "ffi-dataplane")] - pub async fn data_plane_tcp_bind( - &self, - instance_id: &uuid::Uuid, - local_port: u16, - timeout: std::time::Duration, - ) -> Result { - let instance = self - .instance_map - .get(instance_id) - .ok_or_else(|| anyhow::anyhow!("instance {} not found", instance_id))?; - instance.data_plane_tcp_bind(local_port, timeout).await - } - - #[cfg(feature = "ffi-dataplane")] - pub async fn data_plane_udp_bind( - &self, - instance_id: &uuid::Uuid, - local_port: u16, - timeout: std::time::Duration, - ) -> Result { - let instance = self - .instance_map - .get(instance_id) - .ok_or_else(|| anyhow::anyhow!("instance {} not found", instance_id))?; - instance.data_plane_udp_bind(local_port, timeout).await - } - - #[cfg(feature = "ffi-dataplane")] - pub fn data_plane_wait_runtime_handle( - &self, - instance_id: &uuid::Uuid, - timeout: std::time::Duration, - ) -> Option { - self.instance_map - .get(instance_id) - .and_then(|inst| inst.wait_runtime_handle(timeout)) - } - - pub async fn get_network_info( - &self, - instance_id: &uuid::Uuid, - ) -> Option { - if let Some(err_msg) = self.instance_error_messages.get(instance_id) { - return Some(NetworkInstanceRunningInfo { - error_msg: Some(err_msg.value().clone()), - ..Default::default() - }); - } - self.instance_map - .get(instance_id)? - .get_running_info() - .await - .ok() - } - - pub fn list_network_instance_ids(&self) -> Vec { - self.instance_map.iter().map(|item| *item.key()).collect() - } - - pub fn get_instance_name(&self, instance_id: &uuid::Uuid) -> Option { - self.instance_map - .get(instance_id) - .map(|instance| instance.value().get_inst_name()) - } - - pub fn get_network_name(&self, instance_id: &uuid::Uuid) -> Option { - self.instance_map - .get(instance_id) - .map(|instance| instance.value().get_network_name()) - } - - pub fn iter(&self) -> dashmap::iter::Iter<'_, uuid::Uuid, NetworkInstance> { - self.instance_map.iter() - } - - pub fn get_instance_config_control( - &self, - instance_id: &uuid::Uuid, - ) -> Option { - self.instance_map - .get(instance_id) - .map(|instance| instance.value().get_config_file_control().clone()) - } - - pub fn get_instance_config(&self, instance_id: &uuid::Uuid) -> Option { - self.instance_map - .get(instance_id) - .map(|instance| instance.value().get_config()) - } - - pub fn get_instance_network_config_source( - &self, - instance_id: &uuid::Uuid, - ) -> Option { - self.instance_map - .get(instance_id) - .map(|instance| instance.value().get_network_config_source()) - } - - pub fn get_instance_service( - &self, - instance_id: &uuid::Uuid, - ) -> Option> { - self.instance_map - .get(instance_id) - .and_then(|instance| instance.value().get_api_service()) - } - - pub fn set_tun_fd(&self, instance_id: &uuid::Uuid, fd: i32) -> Result<(), anyhow::Error> { - let sender = self - .instance_map - .get(instance_id) - .ok_or_else(|| anyhow::anyhow!("instance not found"))? - .get_tun_fd_sender() - .ok_or_else(|| anyhow::anyhow!("tun fd sender not found"))?; - - sender - .try_send(Some(fd)) - .map_err(|e| anyhow::anyhow!("failed to send tun fd: {}", e))?; - - Ok(()) - } - - pub fn get_config_dir(&self) -> Option<&PathBuf> { - self.config_dir.as_ref() - } - - pub(crate) fn register_daemon(&self) -> DaemonGuard { - DaemonGuard { - guard: Some(self.guard_counter.clone()), - stop_check_notifier: self.stop_check_notifier.clone(), - } - } - - pub(crate) fn notify_stop_check(&self) { - self.stop_check_notifier.notify_one(); - } - - pub async fn wait(&self) { - loop { - let local_instance_running = self - .instance_map - .iter() - .any(|item| item.value().is_easytier_running()); - let daemon_running = Arc::strong_count(&self.guard_counter) > 1; - - if !local_instance_running && !daemon_running { - break; - } - - self.stop_check_notifier.notified().await; - } - } -} - -macro_rules! event { - ($lvl:ident, category: $cat:expr, $($args:tt)+) => { - event!(@impl $lvl, concat!("INSTANCE::", $cat), $($args)+) - }; - - ($lvl:ident, $($args:tt)+) => { - event!(@impl $lvl, "INSTANCE", $($args)+) - }; - - (@impl $lvl:ident, $cat:expr, $($args:tt)+) => { - log::$lvl!( - category: $cat, - $($args)+ - ); - }; -} - -#[tracing::instrument] -fn handle_event( - instance_id: uuid::Uuid, - mut events: EventBusSubscriber, -) -> tokio::task::JoinHandle<()> { - tokio::spawn(async move { - loop { - if let Ok(e) = events.recv().await { - match e { - GlobalCtxEvent::PeerAdded(peer_id) => { - event!(info, peer_id, "[{}] new peer added", instance_id); - } - - GlobalCtxEvent::PeerRemoved(peer_id) => { - event!(info, peer_id, "[{}] peer removed", instance_id); - } - - GlobalCtxEvent::PeerConnAdded(conn_info) => { - event!( - info, - category: "CONNECTION", - %conn_info, - "[{}] new peer connection added", - instance_id, - ); - } - - GlobalCtxEvent::PeerConnRemoved(conn_info) => { - event!( - info, - category: "CONNECTION", - %conn_info, - "[{}] peer connection removed", - instance_id, - ); - } - - GlobalCtxEvent::ListenerAddFailed(listener, msg) => { - event!(warn, %listener, msg, "[{}] listener add failed", instance_id); - } - - GlobalCtxEvent::ListenerAcceptFailed(listener, msg) => { - event!(warn, %listener, msg, "[{}] listener accept failed", instance_id); - } - - GlobalCtxEvent::ListenerAdded(listener) => { - if listener.scheme() == "ring" { - continue; - } - event!( - info, - %listener, - "[{}] new listener added", - instance_id - ); - } - - GlobalCtxEvent::ConnectionAccepted(local, remote) => { - event!(info, category: "CONNECTION", local, remote, "[{}] new connection accepted", instance_id); - } - - GlobalCtxEvent::ConnectionError(local, remote, err) => { - event!(info, category: "CONNECTION", local, remote, err, "[{}] connection error", instance_id); - } - - GlobalCtxEvent::ListenerPortMappingEstablished { - local_listener, - mapped_listener, - backend, - } => { - event!( - info, - %local_listener, - %mapped_listener, - backend, - "[{}] listener port mapping established", - instance_id - ); - } - - GlobalCtxEvent::TunDeviceReady(dev) => { - event!(info, dev, "[{}] tun device ready", instance_id); - } - - GlobalCtxEvent::TunDeviceError(err) => { - event!(error, %err, "[{}] tun device error", instance_id); - } - - GlobalCtxEvent::Connecting(dst) => { - event!(info, category: "CONNECTION", %dst, "[{}] connecting to peer", instance_id); - } - - GlobalCtxEvent::ConnectError(dst, ip_version, error) => { - event!( - info, - category: "CONNECTION", - dst, - ip_version, - %error, - "[{}] connect to peer error", - instance_id - ); - } - - GlobalCtxEvent::VpnPortalStarted(portal) => { - event!(info, portal, "[{}] vpn portal started", instance_id); - } - - GlobalCtxEvent::VpnPortalClientConnected(portal, client_addr) => { - event!( - info, - portal, - client_addr, - "[{}] vpn portal client connected", - instance_id - ); - } - - GlobalCtxEvent::VpnPortalClientDisconnected(portal, client_addr) => { - event!( - info, - portal, - client_addr, - "[{}] vpn portal client disconnected", - instance_id - ); - } - - GlobalCtxEvent::DhcpIpv4Changed(old, new) => { - event!(info, ?old, ?new, "[{}] dhcp ip changed", instance_id); - } - - GlobalCtxEvent::DhcpIpv4Conflicted(ip) => { - event!(info, ?ip, "[{}] dhcp ip conflict", instance_id); - } - - GlobalCtxEvent::PublicIpv6Changed(old, new) => { - event!(info, ?old, ?new, "[{}] public ipv6 changed", instance_id); - } - - GlobalCtxEvent::PublicIpv6RoutesUpdated(added, removed) => { - event!( - info, - ?added, - ?removed, - "[{}] public ipv6 routes updated", - instance_id - ); - } - - GlobalCtxEvent::PortForwardAdded(cfg) => { - event!( - info, - local = %cfg.bind_addr.unwrap(), - remote = %cfg.dst_addr.unwrap(), - proto = %cfg.socket_type().as_str_name(), - "[{}] port forward added", - instance_id, - ); - } - - GlobalCtxEvent::ConfigPatched(patch) => { - event!(info, ?patch, "[{}] config patched", instance_id); - } - - GlobalCtxEvent::ProxyCidrsUpdated(added, removed) => { - event!( - info, - ?added, - ?removed, - "[{}] proxy CIDRs updated", - instance_id - ); - } - - GlobalCtxEvent::UdpBroadcastRelayStartResult { - capture_backend, - error, - } => { - if let Some(error) = error { - event!( - warn, - ?capture_backend, - %error, - "[{}] UDP broadcast relay start failed", - instance_id - ); - } else { - event!( - info, - ?capture_backend, - "[{}] UDP broadcast relay started", - instance_id - ); - } - } - - GlobalCtxEvent::CredentialChanged => { - event!(info, "[{}] credential changed", instance_id); - } - } - } else { - events = events.resubscribe(); - } - } - }) -} - -impl Display for proto::api::instance::PeerConnInfo { - fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result { - f.debug_struct("PeerConnInfo") - .field("my_peer_id", &self.my_peer_id) - .field("dst_peer_id", &self.peer_id) - .field("tunnel_info", &self.tunnel) - .finish() - } -} - -#[cfg(test)] -mod tests { - use super::*; - use crate::common::config::*; - - #[tokio::test] - #[serial_test::serial] - async fn it_works() { - let manager = NetworkInstanceManager::new(); - let cfg_str = r#" - listeners = [] - "#; - - let port = crate::utils::find_free_tcp_port(10012..65534).expect("no free tcp port found"); - - let instance_id1 = manager - .run_network_instance( - TomlConfigLoader::new_from_str(cfg_str) - .inspect(|c| { - c.set_listeners(vec![format!("tcp://0.0.0.0:{}", port).parse().unwrap()]); - }) - .unwrap(), - true, - ConfigFileControl::STATIC_CONFIG, - ) - .unwrap(); - let instance_id2 = manager - .run_network_instance( - TomlConfigLoader::new_from_str(cfg_str).unwrap(), - true, - ConfigFileControl::STATIC_CONFIG, - ) - .unwrap(); - let instance_id3 = manager - .run_network_instance( - TomlConfigLoader::new_from_str(cfg_str).unwrap(), - false, - ConfigFileControl::STATIC_CONFIG, - ) - .unwrap(); - let instance_id4 = manager - .run_network_instance( - TomlConfigLoader::new_from_str(cfg_str).unwrap(), - true, - ConfigFileControl::STATIC_CONFIG, - ) - .unwrap(); - let instance_id5 = manager - .run_network_instance( - TomlConfigLoader::new_from_str(cfg_str).unwrap(), - false, - ConfigFileControl::STATIC_CONFIG, - ) - .unwrap(); - - tokio::time::sleep(std::time::Duration::from_secs(1)).await; // to make instance actually started - - assert!(!crate::utils::check_tcp_available(port)); - - assert!(manager.instance_map.contains_key(&instance_id1)); - assert!(manager.instance_map.contains_key(&instance_id2)); - assert!(manager.instance_map.contains_key(&instance_id3)); - assert!(manager.instance_map.contains_key(&instance_id4)); - assert!(manager.instance_map.contains_key(&instance_id5)); - assert_eq!(manager.list_network_instance_ids().len(), 5); - assert_eq!(manager.instance_stop_tasks.len(), 3); // FFI and GUI instance does not have a stop task - - manager - .delete_network_instance(vec![instance_id3, instance_id4, instance_id5]) - .unwrap(); - assert!(!manager.instance_map.contains_key(&instance_id3)); - assert!(!manager.instance_map.contains_key(&instance_id4)); - assert!(!manager.instance_map.contains_key(&instance_id5)); - assert_eq!(manager.list_network_instance_ids().len(), 2); - } - - #[test] - #[serial_test::serial] - fn test_no_tokio_runtime() { - let manager = NetworkInstanceManager::new(); - let cfg_str = r#" - listeners = [] - "#; - - let port = crate::utils::find_free_tcp_port(10012..65534).expect("no free tcp port found"); - - assert!( - manager - .run_network_instance( - TomlConfigLoader::new_from_str(cfg_str).unwrap(), - true, - ConfigFileControl::STATIC_CONFIG - ) - .is_err() - ); - assert!( - manager - .run_network_instance( - TomlConfigLoader::new_from_str(cfg_str).unwrap(), - true, - ConfigFileControl::STATIC_CONFIG - ) - .is_err() - ); - assert!( - manager - .run_network_instance( - TomlConfigLoader::new_from_str(cfg_str) - .inspect(|c| { - c.set_listeners(vec![ - format!("tcp://0.0.0.0:{}", port).parse().unwrap(), - ]); - }) - .unwrap(), - false, - ConfigFileControl::STATIC_CONFIG - ) - .is_ok() - ); - assert!( - manager - .run_network_instance( - TomlConfigLoader::new_from_str(cfg_str).unwrap(), - true, - ConfigFileControl::STATIC_CONFIG - ) - .is_err() - ); - assert!( - manager - .run_network_instance( - TomlConfigLoader::new_from_str(cfg_str).unwrap(), - false, - ConfigFileControl::STATIC_CONFIG - ) - .is_ok() - ); - - std::thread::sleep(std::time::Duration::from_secs(1)); // wait instance actually started - - assert!(!crate::utils::check_tcp_available(port)); - - assert_eq!(manager.list_network_instance_ids().len(), 5); - assert_eq!( - manager - .instance_map - .iter() - .map(|item| item.is_easytier_running()) - .filter(|x| *x) - .count(), - 5 - ); // stop tasks failed not affect instance running status - assert_eq!(manager.instance_stop_tasks.len(), 0); - } - - #[tokio::test] - #[serial_test::serial] - async fn test_single_instance_failed() { - let free_tcp_port = - crate::utils::find_free_tcp_port(10012..65534).expect("no free tcp port found"); - - // Test with event watching enabled (for CLI/File/RPC usage) - instance should auto-stop on error - for watch_event in [true] { - let _port_holder = - std::net::TcpListener::bind(format!("0.0.0.0:{}", free_tcp_port)).unwrap(); - - let cfg_str = format!( - r#" - listeners = ["tcp://0.0.0.0:{}"] - "#, - free_tcp_port - ); - - let manager = NetworkInstanceManager::new(); - manager - .run_network_instance( - TomlConfigLoader::new_from_str(cfg_str.as_str()).unwrap(), - watch_event, - ConfigFileControl::STATIC_CONFIG, - ) - .unwrap(); - - tokio::select! { - _ = manager.wait() => { - assert_eq!(manager.list_network_instance_ids().len(), 1); - } - _ = tokio::time::sleep(std::time::Duration::from_secs(5)) => { - panic!("instance manager with single failed instance({:?}) should not running", watch_event); - } - } - } - - // Test without event watching (for FFI usage) - instance should remain even if failed - { - let watch_event = false; - let _port_holder = - std::net::TcpListener::bind(format!("0.0.0.0:{}", free_tcp_port)).unwrap(); - - let cfg_str = format!( - r#" - listeners = ["tcp://0.0.0.0:{}"] - "#, - free_tcp_port - ); - - let manager = NetworkInstanceManager::new(); - manager - .run_network_instance( - TomlConfigLoader::new_from_str(cfg_str.as_str()).unwrap(), - watch_event, - ConfigFileControl::STATIC_CONFIG, - ) - .unwrap(); - - assert_eq!(manager.list_network_instance_ids().len(), 1); - } - } - - #[tokio::test] - #[serial_test::serial] - async fn test_multiple_instances_one_failed() { - let free_tcp_port = - crate::utils::find_free_tcp_port(10012..65534).expect("no free tcp port found"); - - let manager = NetworkInstanceManager::new(); - let cfg_str = format!( - r#" - listeners = ["tcp://0.0.0.0:{}"] - [flags] - enable_ipv6 = false - "#, - free_tcp_port - ); - - manager - .run_network_instance( - TomlConfigLoader::new_from_str(cfg_str.as_str()).unwrap(), - true, - ConfigFileControl::STATIC_CONFIG, - ) - .unwrap(); - - tokio::time::sleep(std::time::Duration::from_secs(1)).await; - - manager - .run_network_instance( - TomlConfigLoader::new_from_str(cfg_str.as_str()).unwrap(), - true, - ConfigFileControl::STATIC_CONFIG, - ) - .unwrap(); - - tokio::select! { - _ = manager.wait() => { - panic!("instance manager with multiple instances one failed should still running"); - } - _ = tokio::time::sleep(std::time::Duration::from_secs(2)) => { - assert_eq!(manager.list_network_instance_ids().len(), 2); - } - } - } -} diff --git a/easytier/src/launcher.rs b/easytier/src/launcher.rs deleted file mode 100644 index 2e59b751..00000000 --- a/easytier/src/launcher.rs +++ /dev/null @@ -1,1674 +0,0 @@ -use crate::common::config::{ - ConfigFileControl, ConfigSource, PortForwardConfig, parse_mapped_listener_urls, - process_secure_mode_cfg, -}; -#[cfg(feature = "ffi-dataplane")] -use crate::gateway::socks5::Socks5Server; -#[cfg(feature = "ffi-dataplane")] -pub use crate::gateway::socks5::{DataPlaneTcpListener, DataPlaneTcpStream, DataPlaneUdpSocket}; -use crate::proto::api::{self, manage}; -use crate::proto::rpc_types::controller::BaseController; -use crate::rpc_service::InstanceRpcService; -use crate::{ - common::{ - config::{ - ConfigLoader, NetworkIdentity, PeerConfig, TomlConfigLoader, VpnPortalConfig, - gen_default_flags, - }, - constants::EASYTIER_VERSION, - global_ctx::{EventBusSubscriber, GlobalCtxEvent}, - }, - instance::instance::Instance, - proto::api::instance::list_peer_route_pair, -}; -use anyhow::Context; -use chrono::{DateTime, Local}; -use std::{ - collections::VecDeque, - net::SocketAddr, - sync::{Arc, Mutex, RwLock, atomic::AtomicBool}, -}; -use tokio::{ - sync::{broadcast, mpsc}, - task::JoinSet, -}; - -pub type MyNodeInfo = crate::proto::api::manage::MyNodeInfo; - -type ArcMutApiService = Arc>>>; -type TunFd = Option; - -#[derive(serde::Serialize, Clone)] -pub struct Event { - time: DateTime, - event: GlobalCtxEvent, -} - -struct EasyTierData { - events: RwLock>, - tun_fd: (mpsc::Sender, Mutex>>), - event_subscriber: RwLock>, - instance_stop_notifier: Arc, - #[cfg(feature = "ffi-dataplane")] - data_plane: tokio::sync::watch::Sender>>, - #[cfg(feature = "ffi-dataplane")] - runtime_handle: ( - parking_lot::Mutex>, - parking_lot::Condvar, - ), -} - -impl Default for EasyTierData { - fn default() -> Self { - let (tx, _) = broadcast::channel(16); - let (sender, receiver) = mpsc::channel(16); - Self { - event_subscriber: RwLock::new(tx), - events: RwLock::new(VecDeque::new()), - tun_fd: (sender, Mutex::new(Some(receiver))), - instance_stop_notifier: Arc::new(tokio::sync::Notify::new()), - #[cfg(feature = "ffi-dataplane")] - data_plane: tokio::sync::watch::channel(None).0, - #[cfg(feature = "ffi-dataplane")] - runtime_handle: (parking_lot::Mutex::new(None), parking_lot::Condvar::new()), - } - } -} - -pub struct EasyTierLauncher { - instance_alive: Arc, - stop_flag: Arc, - thread_handle: Option>, - api_service: ArcMutApiService, - running_cfg: String, - error_msg: Arc>>, - data: Arc, -} - -impl EasyTierLauncher { - pub fn new() -> Self { - let instance_alive = Arc::new(AtomicBool::new(false)); - Self { - instance_alive, - thread_handle: None, - api_service: Arc::new(RwLock::new(None)), - error_msg: Arc::new(RwLock::new(None)), - running_cfg: String::new(), - stop_flag: Arc::new(AtomicBool::new(false)), - data: Arc::new(EasyTierData::default()), - } - } - - async fn handle_easytier_event(event: GlobalCtxEvent, data: &EasyTierData) { - let mut events = data.events.write().unwrap(); - let _ = data.event_subscriber.read().unwrap().send(event.clone()); - events.push_front(Event { - time: chrono::Local::now(), - event, - }); - if events.len() > 20 { - events.pop_back(); - } - } - - #[cfg(mobile)] - async fn run_routine_for_mobile( - instance: &Instance, - data: &EasyTierData, - tasks: &mut JoinSet<()>, - ) { - let global_ctx = instance.get_global_ctx(); - let peer_mgr = instance.get_peer_manager(); - let nic_ctx = instance.get_nic_ctx(); - let peer_packet_receiver = instance.get_peer_packet_receiver(); - let mut tun_fd_receiver = data.tun_fd.1.lock().unwrap().take().unwrap(); - - tasks.spawn(async move { - loop { - let Some(tun_fd) = tun_fd_receiver.recv().await.flatten() else { - return; - }; - let res = Instance::setup_nic_ctx_for_mobile( - nic_ctx.clone(), - global_ctx.clone(), - peer_mgr.clone(), - peer_packet_receiver.clone(), - tun_fd, - ) - .await; - } - }); - } - - async fn easytier_routine( - cfg: TomlConfigLoader, - stop_signal: Arc, - api_service: ArcMutApiService, - data: Arc, - ) -> Result<(), anyhow::Error> { - let mut instance = Instance::new(cfg); - let mut tasks = JoinSet::new(); - - // Subscribe to global context events - let global_ctx = instance.get_global_ctx(); - let data_c = data.clone(); - tasks.spawn(async move { - let mut receiver = global_ctx.subscribe(); - loop { - match receiver.recv().await { - Ok(event) => { - Self::handle_easytier_event(event.clone(), &data_c).await; - } - Err(broadcast::error::RecvError::Closed) => { - break; - } - Err(broadcast::error::RecvError::Lagged(_)) => { - // do nothing currently - receiver = receiver.resubscribe(); - } - } - } - }); - - #[cfg(mobile)] - Self::run_routine_for_mobile(&instance, &data, &mut tasks).await; - - if let Err(err) = instance.run().await { - tasks.abort_all(); - drop(tasks); - instance.clear_resources().await; - return Err(err.into()); - } - - #[cfg(feature = "ffi-dataplane")] - data.data_plane - .send_replace(Some(instance.get_socks5_server())); - - api_service - .write() - .unwrap() - .replace(Arc::new(instance.get_api_rpc_service())); - drop(api_service); - - stop_signal.notified().await; - - tasks.abort_all(); - drop(tasks); - - instance.clear_resources().await; - drop(instance); - - Ok(()) - } - - pub fn start(&mut self, cfg_generator: F) - where - F: FnOnce() -> Result + Send + Sync, - { - let error_msg = self.error_msg.clone(); - let cfg = match cfg_generator() { - Err(e) => { - error_msg.write().unwrap().replace(e.to_string()); - return; - } - Ok(cfg) => cfg, - }; - - self.running_cfg = cfg.dump(); - - let stop_flag = self.stop_flag.clone(); - - let instance_alive = self.instance_alive.clone(); - instance_alive.store(true, std::sync::atomic::Ordering::Relaxed); - - let data = self.data.clone(); - let api_service = self.api_service.clone(); - - self.thread_handle = Some(std::thread::spawn(move || { - let rt = if cfg.get_flags().multi_thread { - let worker_threads = 2.max(cfg.get_flags().multi_thread_count as usize); - tokio::runtime::Builder::new_multi_thread() - .worker_threads(worker_threads) - .enable_all() - .build() - } else { - tokio::runtime::Builder::new_current_thread() - .enable_all() - .build() - } - .unwrap(); - - #[cfg(feature = "ffi-dataplane")] - { - let (lock, cvar) = &data.runtime_handle; - *lock.lock() = Some(rt.handle().clone()); - cvar.notify_all(); - } - - let stop_notifier = Arc::new(tokio::sync::Notify::new()); - - let stop_notifier_clone = stop_notifier.clone(); - rt.spawn(async move { - while !stop_flag.load(std::sync::atomic::Ordering::Relaxed) { - tokio::time::sleep(std::time::Duration::from_millis(100)).await; - } - stop_notifier_clone.notify_one(); - }); - - let notifier = data.instance_stop_notifier.clone(); - let ret = rt.block_on(Self::easytier_routine( - cfg, - stop_notifier, - api_service, - data, - )); - if let Err(e) = ret { - error_msg.write().unwrap().replace(format!("{:?}", e)); - } - instance_alive.store(false, std::sync::atomic::Ordering::Relaxed); - notifier.notify_one(); - rt.shutdown_background(); - })); - } - - pub fn error_msg(&self) -> Option { - self.error_msg.read().unwrap().clone() - } - - pub fn running(&self) -> bool { - self.instance_alive - .load(std::sync::atomic::Ordering::Relaxed) - } - - pub fn get_events(&self) -> Vec { - let events = self.data.events.read().unwrap(); - events.iter().cloned().collect() - } - - pub fn get_api_service(&self) -> Option> { - match self.api_service.read() { - Ok(guard) => guard.clone(), - Err(e) => { - tracing::error!("Failed to acquire read lock for api_service: {:?}", e); - None - } - } - } - - #[cfg(feature = "ffi-dataplane")] - pub fn get_data_plane(&self) -> Option> { - self.data.data_plane.borrow().clone() - } - - /// Waits up to `deadline` for the data-plane server to be published. - #[cfg(feature = "ffi-dataplane")] - pub async fn wait_data_plane( - &self, - deadline: tokio::time::Instant, - ) -> Option> { - let mut rx = self.data.data_plane.subscribe(); - loop { - if let Some(server) = rx.borrow_and_update().clone() { - return Some(server); - } - if tokio::time::timeout_at(deadline, rx.changed()) - .await - .is_err() - { - return None; - } - } - } - - /// Blocks up to `timeout` for the runtime handle to be published. - #[cfg(feature = "ffi-dataplane")] - pub fn wait_runtime_handle( - &self, - timeout: std::time::Duration, - ) -> Option { - let (lock, cvar) = &self.data.runtime_handle; - let mut guard = lock.lock(); - cvar.wait_while_for(&mut guard, |h| h.is_none(), timeout); - guard.clone() - } -} - -impl Default for EasyTierLauncher { - fn default() -> Self { - Self::new() - } -} - -impl Drop for EasyTierLauncher { - fn drop(&mut self) { - self.stop_flag - .store(true, std::sync::atomic::Ordering::Relaxed); - if let Some(handle) = self.thread_handle.take() - && let Err(e) = handle.join() - { - println!("Error when joining thread: {:?}", e); - } - } -} - -pub type NetworkInstanceRunningInfo = crate::proto::api::manage::NetworkInstanceRunningInfo; - -pub struct NetworkInstance { - config: TomlConfigLoader, - launcher: Option, - config_file_control: ConfigFileControl, -} - -impl NetworkInstance { - pub fn new(config: TomlConfigLoader, config_file_control: ConfigFileControl) -> Self { - Self { - config, - launcher: None, - config_file_control, - } - } - - pub fn is_easytier_running(&self) -> bool { - self.launcher.is_some() && self.launcher.as_ref().unwrap().running() - } - - pub async fn get_running_info(&self) -> anyhow::Result { - let launcher = self.launcher.as_ref().ok_or_else(|| { - anyhow::anyhow!("instance is not running, please start the instance first") - })?; - let api_service = self.get_api_service().ok_or_else(|| { - anyhow::anyhow!("failed to get api service, instance may not be running") - })?; - let ctrl = BaseController::default(); - - let peers = api_service - .get_peer_manage_service() - .list_peer(ctrl.clone(), api::instance::ListPeerRequest::default()) - .await? - .peer_infos; - let my_info = api_service - .get_peer_manage_service() - .show_node_info(ctrl.clone(), api::instance::ShowNodeInfoRequest::default()) - .await? - .node_info - .ok_or_else(|| anyhow::anyhow!("failed to get my node info"))?; - let vpn_portal_cfg = api_service - .get_vpn_portal_service() - .get_vpn_portal_info( - ctrl.clone(), - api::instance::GetVpnPortalInfoRequest::default(), - ) - .await? - .vpn_portal_info - .map(|i| i.client_config); - let routes = api_service - .get_peer_manage_service() - .list_route(ctrl.clone(), api::instance::ListRouteRequest::default()) - .await? - .routes; - let peer_route_pairs = list_peer_route_pair(peers.clone(), routes.clone()); - let foreign_network_summary = api_service - .get_peer_manage_service() - .get_foreign_network_summary( - ctrl.clone(), - api::instance::GetForeignNetworkSummaryRequest::default(), - ) - .await? - .summary; - let dev_name = api_service - .get_config_service() - .get_config(ctrl.clone(), api::config::GetConfigRequest::default()) - .await? - .config - .ok_or_else(|| anyhow::anyhow!("failed to get config"))? - .dev_name - .unwrap_or_else(|| "".to_string()); - - Ok(NetworkInstanceRunningInfo { - dev_name, - my_node_info: Some(MyNodeInfo { - virtual_ipv4: my_info - .ipv4_addr - .parse::() - .ok() - .map(Into::into), - hostname: my_info.hostname, - version: EASYTIER_VERSION.to_string(), - ips: my_info.ip_list, - stun_info: my_info.stun_info, - listeners: my_info - .listeners - .into_iter() - .map(|s| s.parse::().unwrap().into()) - .collect(), - vpn_portal_cfg, - peer_id: my_info.peer_id, - }), - events: launcher - .get_events() - .iter() - .map(|e| serde_json::to_string(e).unwrap()) - .collect(), - routes, - peers, - peer_route_pairs, - running: launcher.running(), - error_msg: launcher.error_msg(), - foreign_network_summary, - }) - } - - pub fn get_inst_name(&self) -> String { - self.config.get_inst_name() - } - - pub fn get_network_name(&self) -> String { - self.config.get_network_identity().network_name - } - - pub fn get_tun_fd_sender(&self) -> Option> { - self.launcher - .as_ref() - .map(|launcher| launcher.data.tun_fd.0.clone()) - } - - pub fn start(&mut self) -> Result { - if self.is_easytier_running() { - return Ok(self.subscribe_event().unwrap()); - } - - let launcher = EasyTierLauncher::new(); - self.launcher = Some(launcher); - let ev = self.subscribe_event().unwrap(); - - self.launcher - .as_mut() - .unwrap() - .start(|| Ok(self.config.clone())); - - Ok(ev) - } - - pub fn subscribe_event(&self) -> Option> { - self.launcher - .as_ref() - .map(|launcher| launcher.data.event_subscriber.read().unwrap().subscribe()) - } - - pub fn get_stop_notifier(&self) -> Option> { - self.launcher - .as_ref() - .map(|launcher| launcher.data.instance_stop_notifier.clone()) - } - - pub fn get_config_file_control(&self) -> &ConfigFileControl { - &self.config_file_control - } - - pub fn get_config(&self) -> TomlConfigLoader { - self.config.clone() - } - - pub fn get_network_config_source(&self) -> ConfigSource { - self.config.get_network_config_source() - } - - pub fn get_latest_error_msg(&self) -> Option { - if let Some(launcher) = self.launcher.as_ref() { - launcher.error_msg.read().unwrap().clone() - } else { - None - } - } - - pub fn get_api_service(&self) -> Option> { - self.launcher - .as_ref() - .and_then(|launcher| launcher.get_api_service()) - } - - /// Waits up to `timeout` for the data-plane server to come up, returning it - /// together with the deadline so the caller can spend the remaining budget - /// on the actual operation. - #[cfg(feature = "ffi-dataplane")] - async fn wait_data_plane( - &self, - timeout: std::time::Duration, - ) -> anyhow::Result<(Arc, tokio::time::Instant)> { - let deadline = tokio::time::Instant::now() + timeout; - let launcher = self - .launcher - .as_ref() - .ok_or_else(|| anyhow::anyhow!("data plane is not ready"))?; - let server = launcher - .wait_data_plane(deadline) - .await - .ok_or_else(|| anyhow::anyhow!("data plane is not ready"))?; - Ok((server, deadline)) - } - - #[cfg(feature = "ffi-dataplane")] - pub async fn data_plane_tcp_connect( - &self, - dst_addr: SocketAddr, - timeout: std::time::Duration, - ) -> anyhow::Result { - let (server, deadline) = self.wait_data_plane(timeout).await?; - server - .data_plane_tcp_connect(dst_addr, deadline - tokio::time::Instant::now()) - .await - .map_err(Into::into) - } - - #[cfg(feature = "ffi-dataplane")] - pub async fn data_plane_tcp_bind( - &self, - local_port: u16, - timeout: std::time::Duration, - ) -> anyhow::Result { - let (server, deadline) = self.wait_data_plane(timeout).await?; - server - .data_plane_tcp_bind(local_port, deadline - tokio::time::Instant::now()) - .await - .map_err(Into::into) - } - - #[cfg(feature = "ffi-dataplane")] - pub async fn data_plane_udp_bind( - &self, - local_port: u16, - timeout: std::time::Duration, - ) -> anyhow::Result { - let (server, deadline) = self.wait_data_plane(timeout).await?; - server - .data_plane_udp_bind(local_port, deadline - tokio::time::Instant::now()) - .await - .map_err(Into::into) - } - - #[cfg(feature = "ffi-dataplane")] - pub fn wait_runtime_handle( - &self, - timeout: std::time::Duration, - ) -> Option { - self.launcher - .as_ref() - .and_then(|launcher| launcher.wait_runtime_handle(timeout)) - } -} - -pub fn add_proxy_network_to_config( - proxy_network: &str, - cfg: &TomlConfigLoader, -) -> Result<(), anyhow::Error> { - let parts: Vec<&str> = proxy_network.split("->").collect(); - let real_cidr = parts[0] - .parse() - .with_context(|| format!("failed to parse proxy network: {}", parts[0]))?; - - if parts.len() > 2 { - return Err(anyhow::anyhow!( - "invalid proxy network format: {}, support format: or ->, example: - 10.0.0.0/24 or 10.0.0.0/24->192.168.0.0/24", - proxy_network - )); - } - - let mapped_cidr = if parts.len() == 2 { - Some( - parts[1] - .parse() - .with_context(|| format!("failed to parse mapped network: {}", parts[1]))?, - ) - } else { - None - }; - cfg.add_proxy_cidr(real_cidr, mapped_cidr)?; - Ok(()) -} - -pub type NetworkingMethod = crate::proto::api::manage::NetworkingMethod; -pub type NetworkConfig = crate::proto::api::manage::NetworkConfig; - -impl NetworkConfig { - fn parse_peer(peer: &manage::NetworkPeerConfig) -> Result, anyhow::Error> { - let uri = peer.uri.trim(); - if uri.is_empty() { - return Ok(None); - } - - Ok(Some(PeerConfig { - uri: uri - .parse() - .with_context(|| format!("failed to parse peer uri: {}", uri))?, - peer_public_key: peer.peer_public_key.clone(), - })) - } - - fn parse_peers(peers: &[manage::NetworkPeerConfig]) -> Result, anyhow::Error> { - let mut ret = Vec::new(); - for peer in peers { - if let Some(peer) = Self::parse_peer(peer)? { - ret.push(peer); - } - } - Ok(ret) - } - - fn parse_peer_urls(peer_urls: &[String]) -> Result, anyhow::Error> { - let mut peers = vec![]; - for peer_url in peer_urls.iter() { - let peer_url = peer_url.trim(); - if peer_url.is_empty() { - continue; - } - peers.push(PeerConfig { - uri: peer_url - .parse() - .with_context(|| format!("failed to parse peer uri: {}", peer_url))?, - peer_public_key: None, - }); - } - Ok(peers) - } - - pub fn gen_config(&self) -> Result { - let cfg = TomlConfigLoader::default(); - cfg.set_id( - self.instance_id - .clone() - .unwrap_or(uuid::Uuid::new_v4().to_string()) - .parse() - .with_context(|| format!("failed to parse instance id: {:?}", self.instance_id))?, - ); - cfg.set_hostname(self.hostname.clone()); - cfg.set_dhcp(self.dhcp.unwrap_or_default()); - cfg.set_inst_name(self.network_name.clone().unwrap_or_default()); - - // The web UI does not expose credential inputs directly, but imported/saved - // NetworkConfig objects still need to preserve credential-mode instances via - // secure_mode.local_private_key + empty network_secret. - let credential_secret = if self.network_secret.is_some() { - None - } else { - self.secure_mode - .as_ref() - .and_then(|mode| mode.local_private_key.clone()) - .filter(|s| !s.is_empty()) - }; - - if credential_secret.is_some() { - cfg.set_network_identity(NetworkIdentity::new_credential( - self.network_name.clone().unwrap_or_default(), - )); - } else { - cfg.set_network_identity(NetworkIdentity::new( - self.network_name.clone().unwrap_or_default(), - self.network_secret.clone().unwrap_or_default(), - )); - } - - if !cfg.get_dhcp() { - let virtual_ipv4 = self.virtual_ipv4.clone().unwrap_or_default(); - if !virtual_ipv4.is_empty() { - let ip = format!("{}/{}", virtual_ipv4, self.network_length.unwrap_or(24)) - .parse() - .with_context(|| { - format!( - "failed to parse ipv4 inet address: {}, {:?}", - virtual_ipv4, self.network_length - ) - })?; - cfg.set_ipv4(Some(ip)); - } - } - - match NetworkingMethod::try_from(self.networking_method.unwrap_or_default()) - .unwrap_or_default() - { - NetworkingMethod::PublicServer => { - let peers = Self::parse_peers(&self.peers)?; - if peers.is_empty() { - let public_server_url = self.public_server_url.clone().unwrap_or_default(); - cfg.set_peers(vec![PeerConfig { - uri: public_server_url.parse().with_context(|| { - format!("failed to parse public server uri: {}", public_server_url) - })?, - peer_public_key: None, - }]); - } else { - cfg.set_peers(peers); - } - } - NetworkingMethod::Manual => { - let mut peers = Self::parse_peers(&self.peers)?; - if peers.is_empty() { - peers = Self::parse_peer_urls(&self.peer_urls)?; - } - if !peers.is_empty() { - cfg.set_peers(peers); - } - } - NetworkingMethod::Standalone => {} - } - - let mut listener_urls = vec![]; - for listener_url in self.listener_urls.iter() { - if listener_url.is_empty() { - continue; - } - listener_urls.push( - listener_url - .parse() - .with_context(|| format!("failed to parse listener uri: {}", listener_url))?, - ); - } - cfg.set_listeners(listener_urls); - - for n in self.proxy_cidrs.iter() { - add_proxy_network_to_config(n, &cfg)?; - } - - if !self.port_forwards.is_empty() { - cfg.set_port_forwards( - self.port_forwards - .iter() - .filter(|pf| !pf.bind_ip.is_empty() && !pf.dst_ip.is_empty()) - .filter_map(|pf| { - let bind_addr = - format!("{}:{}", pf.bind_ip, pf.bind_port).parse::(); - let dst_addr = - format!("{}:{}", pf.dst_ip, pf.dst_port).parse::(); - - match (bind_addr, dst_addr) { - (Ok(bind_addr), Ok(dst_addr)) => Some(PortForwardConfig { - bind_addr, - dst_addr, - proto: pf.proto.clone(), - }), - _ => None, - } - }) - .collect::>(), - ); - } - - if self.enable_vpn_portal.unwrap_or_default() { - let cidr = format!( - "{}/{}", - self.vpn_portal_client_network_addr - .clone() - .unwrap_or_default(), - self.vpn_portal_client_network_len.unwrap_or(24) - ); - cfg.set_vpn_portal_config(VpnPortalConfig { - client_cidr: cidr - .parse() - .with_context(|| format!("failed to parse vpn portal client cidr: {}", cidr))?, - wireguard_listen: format!( - "0.0.0.0:{}", - self.vpn_portal_listen_port.unwrap_or_default() - ) - .parse() - .with_context(|| { - format!( - "failed to parse vpn portal wireguard listen port. {:?}", - self.vpn_portal_listen_port - ) - })?, - }); - } - - if self.enable_manual_routes.unwrap_or_default() { - let mut routes = Vec::::with_capacity(self.routes.len()); - for route in self.routes.iter() { - routes.push( - route - .parse() - .with_context(|| format!("failed to parse route: {}", route))?, - ); - } - cfg.set_routes(Some(routes)); - } - - if !self.exit_nodes.is_empty() { - let mut exit_nodes = Vec::::with_capacity(self.exit_nodes.len()); - for node in self.exit_nodes.iter() { - exit_nodes.push( - node.parse() - .with_context(|| format!("failed to parse exit node: {}", node))?, - ); - } - cfg.set_exit_nodes(exit_nodes); - } - - if self.enable_socks5.unwrap_or_default() - && let Some(socks5_port) = self.socks5_port - { - cfg.set_socks5_portal(Some( - format!("socks5://0.0.0.0:{}", socks5_port).parse().unwrap(), - )); - } - - if !self.mapped_listeners.is_empty() { - let mapped_listeners = parse_mapped_listener_urls(&self.mapped_listeners)?; - cfg.set_mapped_listeners(Some(mapped_listeners)); - } - - if let Some(credential_file) = self - .credential_file - .as_ref() - .filter(|path| !path.is_empty()) - { - cfg.set_credential_file(Some(credential_file.into())); - } - - if let Some(credential_secret) = credential_secret { - cfg.set_secure_mode(Some(process_secure_mode_cfg( - crate::proto::common::SecureModeConfig { - enabled: true, - local_private_key: Some(credential_secret), - local_public_key: None, - }, - )?)); - } else { - cfg.set_secure_mode( - self.secure_mode - .clone() - .map(process_secure_mode_cfg) - .transpose()?, - ); - } - - let mut flags = gen_default_flags(); - if let Some(latency_first) = self.latency_first { - flags.latency_first = latency_first; - } - - if let Some(dev_name) = self.dev_name.clone() { - flags.dev_name = dev_name; - } - - if let Some(use_smoltcp) = self.use_smoltcp { - flags.use_smoltcp = use_smoltcp; - } - - if let Some(ipv6_public_addr_provider) = self.ipv6_public_addr_provider { - cfg.set_ipv6_public_addr_provider(ipv6_public_addr_provider); - } - - if let Some(ipv6_public_addr_auto) = self.ipv6_public_addr_auto { - cfg.set_ipv6_public_addr_auto(ipv6_public_addr_auto); - } - - if let Some(ipv6_public_addr_prefix) = self - .ipv6_public_addr_prefix - .as_ref() - .filter(|prefix| !prefix.is_empty()) - { - cfg.set_ipv6_public_addr_prefix(Some(ipv6_public_addr_prefix.parse().with_context( - || format!("failed to parse ipv6 public address prefix: {ipv6_public_addr_prefix}"), - )?)); - } - - if let Some(disable_ipv6) = self.disable_ipv6 { - flags.enable_ipv6 = !disable_ipv6; - } - - if let Some(enable_kcp_proxy) = self.enable_kcp_proxy { - flags.enable_kcp_proxy = enable_kcp_proxy; - } - - if let Some(disable_kcp_input) = self.disable_kcp_input { - flags.disable_kcp_input = disable_kcp_input; - } - - if let Some(enable_quic_proxy) = self.enable_quic_proxy { - flags.enable_quic_proxy = enable_quic_proxy; - } - - if let Some(disable_quic_input) = self.disable_quic_input { - flags.disable_quic_input = disable_quic_input; - } - - if let Some(disable_p2p) = self.disable_p2p { - flags.disable_p2p = disable_p2p; - } - - if let Some(p2p_only) = self.p2p_only { - flags.p2p_only = p2p_only; - } - - if let Some(lazy_p2p) = self.lazy_p2p { - flags.lazy_p2p = lazy_p2p; - } - - if let Some(bind_device) = self.bind_device { - flags.bind_device = bind_device; - } - - if self.socket_mark.is_some() { - flags.socket_mark = self.socket_mark; - } - - if let Some(no_tun) = self.no_tun { - flags.no_tun = no_tun; - } - - if let Some(enable_exit_node) = self.enable_exit_node { - flags.enable_exit_node = enable_exit_node; - } - - if let Some(relay_all_peer_rpc) = self.relay_all_peer_rpc { - flags.relay_all_peer_rpc = relay_all_peer_rpc; - } - - if let Some(need_p2p) = self.need_p2p { - flags.need_p2p = need_p2p; - } - - if let Some(multi_thread) = self.multi_thread { - flags.multi_thread = multi_thread; - } - - if let Some(proxy_forward_by_system) = self.proxy_forward_by_system { - flags.proxy_forward_by_system = proxy_forward_by_system; - } - - if let Some(disable_encryption) = self.disable_encryption { - flags.enable_encryption = !disable_encryption; - } - - if self.enable_relay_network_whitelist.unwrap_or_default() { - if !self.relay_network_whitelist.is_empty() { - flags.relay_network_whitelist = self.relay_network_whitelist.join(" "); - } else { - flags.relay_network_whitelist = "".to_string(); - } - } - - if let Some(disable_tcp_hole_punching) = self.disable_tcp_hole_punching { - flags.disable_tcp_hole_punching = disable_tcp_hole_punching; - } - - if let Some(disable_udp_hole_punching) = self.disable_udp_hole_punching { - flags.disable_udp_hole_punching = disable_udp_hole_punching; - } - - if let Some(disable_upnp) = self.disable_upnp { - flags.disable_upnp = disable_upnp; - } - - if let Some(disable_relay_data) = self.disable_relay_data { - flags.disable_relay_data = disable_relay_data; - } - - if let Some(enable_udp_broadcast_relay) = self.enable_udp_broadcast_relay { - flags.enable_udp_broadcast_relay = enable_udp_broadcast_relay; - } - - if let Some(disable_sym_hole_punching) = self.disable_sym_hole_punching { - flags.disable_sym_hole_punching = disable_sym_hole_punching; - } - - if let Some(enable_magic_dns) = self.enable_magic_dns { - flags.accept_dns = enable_magic_dns; - } - - if let Some(mtu) = self.mtu { - flags.mtu = mtu as u32; - } - - if let Some(instance_recv_bps_limit) = self.instance_recv_bps_limit { - flags.instance_recv_bps_limit = instance_recv_bps_limit; - } - - if let Some(enable_private_mode) = self.enable_private_mode { - flags.private_mode = enable_private_mode; - } - - if let Some(encryption_algorithm) = self.encryption_algorithm.clone() { - flags.encryption_algorithm = encryption_algorithm; - } - - if let Some(acl) = self.acl.as_ref() - && !acl.is_empty() - { - cfg.set_acl(Some(acl.clone())); - } - - if let Some(data_compress_algo) = self.data_compress_algo { - if data_compress_algo < 1 { - flags.data_compress_algo = 1; - } else { - flags.data_compress_algo = data_compress_algo - } - } - - cfg.set_flags(flags); - Ok(cfg) - } - - pub fn new_from_config(config: impl ConfigLoader) -> Result { - let default_config = TomlConfigLoader::default(); - - let mut result = Self { - ..Default::default() - }; - - result.instance_id = Some(config.get_id().to_string()); - if config.get_hostname() != default_config.get_hostname() { - result.hostname = Some(config.get_hostname()); - } - - result.dhcp = Some(config.get_dhcp()); - - let network_identity = config.get_network_identity(); - result.network_name = Some(network_identity.network_name.clone()); - result.network_secret = network_identity.network_secret; - - if let Some(ipv4) = config.get_ipv4() { - result.virtual_ipv4 = Some(ipv4.address().to_string()); - result.network_length = Some(ipv4.network_length() as i32); - } - - if config.get_ipv6_public_addr_provider() != default_config.get_ipv6_public_addr_provider() - { - result.ipv6_public_addr_provider = Some(config.get_ipv6_public_addr_provider()); - } - if config.get_ipv6_public_addr_auto() != default_config.get_ipv6_public_addr_auto() { - result.ipv6_public_addr_auto = Some(config.get_ipv6_public_addr_auto()); - } - result.ipv6_public_addr_prefix = config - .get_ipv6_public_addr_prefix() - .map(|prefix| prefix.to_string()); - - let peers = config.get_peers(); - result.networking_method = Some(NetworkingMethod::Manual as i32); - if !peers.is_empty() { - result.peer_urls = peers.iter().map(|p| p.uri.to_string()).collect(); - result.peers = peers - .iter() - .map(|p| manage::NetworkPeerConfig { - uri: p.uri.to_string(), - peer_public_key: p.peer_public_key.clone(), - }) - .collect(); - } - - result.listener_urls = config - .get_listeners() - .unwrap_or_default() - .iter() - .map(|l| l.to_string()) - .collect(); - - result.proxy_cidrs = config - .get_proxy_cidrs() - .iter() - .map(|c| { - if let Some(mapped) = c.mapped_cidr { - format!("{}->{}", c.cidr, mapped) - } else { - c.cidr.to_string() - } - }) - .collect(); - - let port_forwards = config.get_port_forwards(); - if !port_forwards.is_empty() { - result.port_forwards = port_forwards - .iter() - .map(|f| manage::PortForwardConfig { - proto: f.proto.clone(), - bind_ip: f.bind_addr.ip().to_string(), - bind_port: f.bind_addr.port() as u32, - dst_ip: f.dst_addr.ip().to_string(), - dst_port: f.dst_addr.port() as u32, - }) - .collect(); - } - - if let Some(vpn_config) = config.get_vpn_portal_config() { - result.enable_vpn_portal = Some(true); - - let cidr = vpn_config.client_cidr; - result.vpn_portal_client_network_addr = Some(cidr.first_address().to_string()); - result.vpn_portal_client_network_len = Some(cidr.network_length() as i32); - - result.vpn_portal_listen_port = Some(vpn_config.wireguard_listen.port() as i32); - } - - if let Some(routes) = config.get_routes() - && !routes.is_empty() - { - result.enable_manual_routes = Some(true); - result.routes = routes.iter().map(|r| r.to_string()).collect(); - } - - let exit_nodes = config.get_exit_nodes(); - if !exit_nodes.is_empty() { - result.exit_nodes = exit_nodes.iter().map(|n| n.to_string()).collect(); - } - - if let Some(socks5_portal) = config.get_socks5_portal() { - result.enable_socks5 = Some(true); - result.socks5_port = socks5_portal.port().map(|p| p as i32); - } - - let mapped_listeners = config.get_mapped_listeners(); - if !mapped_listeners.is_empty() { - result.mapped_listeners = mapped_listeners.iter().map(|l| l.to_string()).collect(); - } - - result.secure_mode = config.get_secure_mode(); - result.credential_file = config - .get_credential_file() - .map(|path| path.to_string_lossy().into_owned()); - let flags = config.get_flags(); - let default_flags = default_config.get_flags(); - result.latency_first = Some(flags.latency_first); - result.dev_name = Some(flags.dev_name.clone()); - result.use_smoltcp = Some(flags.use_smoltcp); - result.disable_ipv6 = Some(!flags.enable_ipv6); - result.enable_kcp_proxy = Some(flags.enable_kcp_proxy); - result.disable_kcp_input = Some(flags.disable_kcp_input); - result.enable_quic_proxy = Some(flags.enable_quic_proxy); - result.disable_quic_input = Some(flags.disable_quic_input); - result.disable_p2p = Some(flags.disable_p2p); - result.p2p_only = Some(flags.p2p_only); - result.lazy_p2p = Some(flags.lazy_p2p); - result.bind_device = Some(flags.bind_device); - result.socket_mark = flags.socket_mark; - result.no_tun = Some(flags.no_tun); - result.enable_exit_node = Some(flags.enable_exit_node); - result.relay_all_peer_rpc = Some(flags.relay_all_peer_rpc); - result.need_p2p = Some(flags.need_p2p); - result.multi_thread = Some(flags.multi_thread); - result.proxy_forward_by_system = Some(flags.proxy_forward_by_system); - result.disable_encryption = Some(!flags.enable_encryption); - result.disable_tcp_hole_punching = Some(flags.disable_tcp_hole_punching); - result.disable_udp_hole_punching = Some(flags.disable_udp_hole_punching); - result.disable_upnp = Some(flags.disable_upnp); - result.disable_relay_data = Some(flags.disable_relay_data); - result.enable_udp_broadcast_relay = Some(flags.enable_udp_broadcast_relay); - result.disable_sym_hole_punching = Some(flags.disable_sym_hole_punching); - result.enable_magic_dns = Some(flags.accept_dns); - result.mtu = Some(flags.mtu as i32); - result.data_compress_algo = (flags.data_compress_algo != default_flags.data_compress_algo) - .then_some(flags.data_compress_algo); - result.encryption_algorithm = (flags.encryption_algorithm - != default_flags.encryption_algorithm) - .then_some(flags.encryption_algorithm.clone()); - result.instance_recv_bps_limit = - (flags.instance_recv_bps_limit != u64::MAX).then_some(flags.instance_recv_bps_limit); - result.enable_private_mode = Some(flags.private_mode); - - result.acl = config.get_acl(); - - if flags.relay_network_whitelist == "*" { - result.enable_relay_network_whitelist = Some(false); - } else { - result.enable_relay_network_whitelist = Some(true); - if flags.relay_network_whitelist.is_empty() { - result.relay_network_whitelist = vec![]; - } else { - result.relay_network_whitelist = flags - .relay_network_whitelist - .split_whitespace() - .map(|s| s.to_string()) - .collect(); - } - } - - Ok(result) - } -} - -#[cfg(test)] -mod tests { - use crate::{ - common::config::{ConfigLoader, process_secure_mode_cfg}, - proto::common::{CompressionAlgoPb, SecureModeConfig}, - }; - use base64::prelude::{BASE64_STANDARD, Engine as _}; - use rand::Rng; - use std::net::{IpAddr, Ipv4Addr, Ipv6Addr}; - - fn gen_default_config() -> crate::common::config::TomlConfigLoader { - let config = crate::common::config::TomlConfigLoader::default(); - config.set_id(uuid::Uuid::new_v4()); - config.set_dhcp(false); - config.set_inst_name("default".to_string()); - config.set_listeners(vec![]); - config - } - - #[test] - fn test_network_config_conversion_basic() -> Result<(), anyhow::Error> { - let config = gen_default_config(); - - let network_config = super::NetworkConfig::new_from_config(&config)?; - - let generated_config = network_config.gen_config()?; - - let config_str = config.dump(); - let generated_config_str = generated_config.dump(); - - assert_eq!( - config_str, - generated_config_str, - "Generated config does not match original config:\nOriginal:\n{}\n\nGenerated:\n{}\nNetwork Config: {}\n", - config_str, - generated_config_str, - serde_json::to_string(&network_config).unwrap() - ); - Ok(()) - } - - #[test] - fn network_config_dump_preserves_web_flags() -> Result<(), anyhow::Error> { - let network_config = super::NetworkConfig { - instance_id: Some(uuid::Uuid::new_v4().to_string()), - dhcp: Some(true), - network_name: Some("demo".to_string()), - network_secret: Some("secret".to_string()), - networking_method: Some(crate::proto::api::manage::NetworkingMethod::Manual as i32), - peer_urls: vec!["tcp://1.2.3.4:11010".to_string()], - listener_urls: vec!["tcp://0.0.0.0:11010".to_string()], - dev_name: Some("et_test".to_string()), - enable_quic_proxy: Some(true), - disable_tcp_hole_punching: Some(true), - disable_sym_hole_punching: Some(true), - ..Default::default() - }; - - let dumped = network_config.gen_config()?.dump(); - - assert!(dumped.contains("dev_name = \"et_test\"")); - assert!(dumped.contains("enable_quic_proxy = true")); - assert!(dumped.contains("disable_tcp_hole_punching = true")); - assert!(dumped.contains("disable_sym_hole_punching = true")); - Ok(()) - } - - #[test] - fn test_network_config_conversion_preserves_peer_public_key() -> Result<(), anyhow::Error> { - let peer_url = "tcp://1.2.3.4:11010"; - let peer_public_key = BASE64_STANDARD.encode([9u8; 32]); - let config = gen_default_config(); - config.set_peers(vec![crate::common::config::PeerConfig { - uri: peer_url.parse()?, - peer_public_key: Some(peer_public_key.clone()), - }]); - - let network_config = super::NetworkConfig::new_from_config(&config)?; - - assert_eq!(network_config.peer_urls, vec![peer_url.to_string()]); - assert_eq!(network_config.peers.len(), 1); - assert_eq!(network_config.peers[0].uri, peer_url); - assert_eq!( - network_config.peers[0].peer_public_key.as_deref(), - Some(peer_public_key.as_str()) - ); - - let generated_config = network_config.gen_config()?; - assert_eq!(generated_config.get_peers(), config.get_peers()); - Ok(()) - } - - #[test] - fn network_config_gen_config_trims_legacy_peer_urls() -> Result<(), anyhow::Error> { - let network_config = super::NetworkConfig { - instance_id: Some(uuid::Uuid::new_v4().to_string()), - dhcp: Some(true), - networking_method: Some(crate::proto::api::manage::NetworkingMethod::Manual as i32), - peer_urls: vec![ - " tcp://1.2.3.4:11010 ".to_string(), - " ".to_string(), - "\tudp://5.6.7.8:11010\n".to_string(), - ], - ..Default::default() - }; - - let generated_config = network_config.gen_config()?; - let peers = generated_config.get_peers(); - - assert_eq!(peers.len(), 2); - assert_eq!(peers[0].uri.as_str(), "tcp://1.2.3.4:11010"); - assert_eq!(peers[1].uri.as_str(), "udp://5.6.7.8:11010"); - Ok(()) - } - - #[test] - fn test_network_config_conversion_random() -> Result<(), anyhow::Error> { - let mut rng = rand::thread_rng(); - - for _ in 0..100 { - let config = gen_default_config(); - - config.set_id(uuid::Uuid::new_v4()); - - config.set_dhcp(rng.gen_bool(0.5)); - - if rng.gen_bool(0.7) { - let hostname = format!("host-{}", rng.r#gen::()); - config.set_hostname(Some(hostname)); - } - - config.set_network_identity(crate::common::config::NetworkIdentity::new( - format!("network-{}", rng.r#gen::()), - format!("secret-{}", rng.r#gen::()), - )); - config.set_inst_name(config.get_network_identity().network_name.clone()); - - if !config.get_dhcp() { - let addr = Ipv4Addr::new( - rng.gen_range(1..254), - rng.gen_range(0..255), - rng.gen_range(0..255), - rng.gen_range(1..254), - ); - let prefix_len = rng.gen_range(1..31); - let ipv4 = format!("{}/{}", addr, prefix_len).parse().unwrap(); - config.set_ipv4(Some(ipv4)); - } - - let peer_count = rng.gen_range(0..3); - let mut peers = Vec::new(); - for _ in 0..peer_count { - let port = rng.gen_range(10000..60000); - let protocol = if rng.gen_bool(0.5) { "tcp" } else { "udp" }; - let uri = format!("{}://127.0.0.1:{}", protocol, port) - .parse() - .unwrap(); - peers.push(crate::common::config::PeerConfig { - uri, - peer_public_key: None, - }); - } - config.set_peers(peers); - - if rng.gen_bool(0.7) { - let listener_count = rng.gen_range(0..3); - let mut listeners = Vec::new(); - for _ in 0..listener_count { - let port = rng.gen_range(10000..60000); - let protocol = if rng.gen_bool(0.5) { "tcp" } else { "udp" }; - listeners.push(format!("{}://0.0.0.0:{}", protocol, port).parse().unwrap()); - } - config.set_listeners(listeners); - } - - if rng.gen_bool(0.6) { - let proxy_count = rng.gen_range(0..3); - for _ in 0..proxy_count { - let network = format!( - "{}.{}.{}.0/{}", - rng.gen_range(1..254), - rng.gen_range(0..255), - rng.gen_range(0..255), - rng.gen_range(24..30) - ) - .parse::() - .unwrap(); - - let mapped_network = if rng.gen_bool(0.5) { - Some( - format!( - "{}.{}.{}.0/{}", - rng.gen_range(1..254), - rng.gen_range(0..255), - rng.gen_range(0..255), - network.network_length() - ) - .parse::() - .unwrap(), - ) - } else { - None - }; - config.add_proxy_cidr(network, mapped_network).unwrap(); - } - } - - if rng.gen_bool(0.5) { - let vpn_network = format!( - "{}.{}.{}.0/{}", - rng.gen_range(10..173), - rng.gen_range(0..255), - rng.gen_range(0..255), - rng.gen_range(24..30) - ); - let vpn_port = rng.gen_range(10000..60000); - config.set_vpn_portal_config(crate::common::config::VpnPortalConfig { - client_cidr: vpn_network.parse().unwrap(), - wireguard_listen: format!("0.0.0.0:{}", vpn_port).parse().unwrap(), - }); - } - - if rng.gen_bool(0.6) { - let route_count = rng.gen_range(1..3); - let mut routes = Vec::new(); - for _ in 0..route_count { - let route = format!( - "{}.{}.{}.0/{}", - rng.gen_range(1..254), - rng.gen_range(0..255), - rng.gen_range(0..255), - rng.gen_range(24..30) - ); - routes.push(route.parse().unwrap()); - } - config.set_routes(Some(routes)); - } - - if rng.gen_bool(0.4) { - let node_count = rng.gen_range(1..3); - let mut nodes = Vec::new(); - for _ in 0..node_count { - let ip = Ipv4Addr::new( - rng.gen_range(1..254), - rng.gen_range(0..255), - rng.gen_range(0..255), - rng.gen_range(1..254), - ); - nodes.push(IpAddr::V4(ip)); - // gen ipv6 - let ip = Ipv6Addr::new( - rng.gen_range(0..65535), - rng.gen_range(0..65535), - rng.gen_range(0..65535), - rng.gen_range(0..65535), - rng.gen_range(0..65535), - rng.gen_range(0..65535), - rng.gen_range(0..65535), - rng.gen_range(0..65535), - ); - nodes.push(IpAddr::V6(ip)); - } - config.set_exit_nodes(nodes); - } - - if rng.gen_bool(0.5) { - let socks5_port = rng.gen_range(10000..60000); - config.set_socks5_portal(Some( - format!("socks5://0.0.0.0:{}", socks5_port).parse().unwrap(), - )); - } - - if rng.gen_bool(0.4) { - let count = rng.gen_range(1..3); - let mut mapped_listeners = Vec::new(); - for _ in 0..count { - let port = rng.gen_range(10000..60000); - mapped_listeners.push(format!("tcp://0.0.0.0:{}", port).parse().unwrap()); - } - config.set_mapped_listeners(Some(mapped_listeners)); - } - - if rng.gen_bool(0.3) { - config.set_secure_mode(Some(SecureModeConfig { - enabled: true, - local_private_key: None, - local_public_key: None, - })); - } - - if rng.gen_bool(0.9) { - let mut flags = crate::common::config::gen_default_flags(); - flags.latency_first = rng.gen_bool(0.5); - flags.dev_name = format!("etun{}", rng.gen_range(0..10)); - flags.use_smoltcp = rng.gen_bool(0.3); - flags.enable_ipv6 = rng.gen_bool(0.8); - flags.enable_kcp_proxy = rng.gen_bool(0.5); - flags.disable_kcp_input = rng.gen_bool(0.3); - flags.enable_quic_proxy = rng.gen_bool(0.5); - flags.disable_quic_input = rng.gen_bool(0.3); - flags.disable_p2p = rng.gen_bool(0.2); - flags.p2p_only = rng.gen_bool(0.2); - flags.lazy_p2p = rng.gen_bool(0.3); - flags.bind_device = rng.gen_bool(0.3); - flags.no_tun = rng.gen_bool(0.1); - flags.enable_exit_node = rng.gen_bool(0.4); - flags.relay_all_peer_rpc = rng.gen_bool(0.5); - flags.need_p2p = rng.gen_bool(0.3); - flags.multi_thread = rng.gen_bool(0.7); - flags.proxy_forward_by_system = rng.gen_bool(0.3); - flags.enable_encryption = rng.gen_bool(0.8); - flags.disable_tcp_hole_punching = rng.gen_bool(0.2); - flags.disable_udp_hole_punching = rng.gen_bool(0.2); - flags.disable_upnp = rng.gen_bool(0.2); - flags.enable_udp_broadcast_relay = rng.gen_bool(0.2); - flags.accept_dns = rng.gen_bool(0.6); - flags.mtu = rng.gen_range(1200..1500); - flags.private_mode = rng.gen_bool(0.3); - - if rng.gen_bool(0.4) { - flags.relay_network_whitelist = (0..rng.gen_range(1..3)) - .map(|_| { - format!( - "{}.{}.0.0/16", - rng.gen_range(10..192), - rng.gen_range(0..255) - ) - }) - .collect::>() - .join(" "); - } - - config.set_flags(flags); - } - - if let Some(secure_mode) = config.get_secure_mode() { - config.set_secure_mode(Some(process_secure_mode_cfg(secure_mode)?)); - } - - let network_config = super::NetworkConfig::new_from_config(&config)?; - let generated_config = network_config.gen_config()?; - generated_config.set_peers(generated_config.get_peers()); // Ensure peers field is not None - - let config_str = config.dump(); - let generated_config_str = generated_config.dump(); - - assert_eq!( - config_str, - generated_config_str, - "Generated config does not match original config:\nOriginal:\n{}\n\nGenerated:\n{}\nNetwork Config: {}\n", - config_str, - generated_config_str, - serde_json::to_string(&network_config).unwrap() - ); - } - - Ok(()) - } - - #[test] - fn test_network_config_conversion_credential_mode() -> Result<(), anyhow::Error> { - let private_key = x25519_dalek::StaticSecret::from([7u8; 32]); - let public_key = x25519_dalek::PublicKey::from(&private_key); - let credential_secret = BASE64_STANDARD.encode(private_key.as_bytes()); - let credential_file = "/tmp/easytier-credentials.json".to_string(); - - let config = gen_default_config(); - config.set_network_identity(crate::common::config::NetworkIdentity::new_credential( - "credential-net".to_string(), - )); - config.set_inst_name("credential-net".to_string()); - config.set_credential_file(Some(credential_file.clone().into())); - config.set_secure_mode(Some(SecureModeConfig { - enabled: true, - local_private_key: Some(credential_secret.clone()), - local_public_key: Some(BASE64_STANDARD.encode(public_key.as_bytes())), - })); - - let network_config = super::NetworkConfig::new_from_config(&config)?; - assert_eq!( - network_config.credential_file.as_deref(), - Some(credential_file.as_str()) - ); - assert_eq!(network_config.network_secret, None); - assert_eq!( - network_config - .secure_mode - .as_ref() - .and_then(|mode| mode.local_private_key.as_deref()), - Some(credential_secret.as_str()) - ); - - let generated_config = network_config.gen_config()?; - assert_eq!( - generated_config.get_network_identity().network_secret, - None, - "credential mode should not be converted back into network_secret mode" - ); - assert_eq!( - generated_config - .get_credential_file() - .map(|path| path.to_string_lossy().into_owned()), - Some(credential_file) - ); - assert_eq!( - generated_config - .get_secure_mode() - .and_then(|mode| mode.local_private_key), - Some(credential_secret) - ); - - Ok(()) - } - - #[test] - fn test_network_config_conversion_preserves_runtime_algorithm_flags() - -> Result<(), anyhow::Error> { - let config = gen_default_config(); - let mut flags = config.get_flags(); - flags.data_compress_algo = CompressionAlgoPb::Zstd.into(); - flags.encryption_algorithm = "managed-test-algo".to_string(); - config.set_flags(flags.clone()); - - let network_config = super::NetworkConfig::new_from_config(&config)?; - - assert_eq!( - network_config.data_compress_algo, - Some(CompressionAlgoPb::Zstd as i32) - ); - assert_eq!( - network_config.encryption_algorithm.as_deref(), - Some("managed-test-algo") - ); - - let generated_config = network_config.gen_config()?; - assert_eq!( - generated_config.get_flags().data_compress_algo, - flags.data_compress_algo - ); - assert_eq!( - generated_config.get_flags().encryption_algorithm, - flags.encryption_algorithm - ); - - Ok(()) - } -} diff --git a/easytier/src/lib.rs b/easytier/src/lib.rs index 91c0a042..4fe5e823 100644 --- a/easytier/src/lib.rs +++ b/easytier/src/lib.rs @@ -1,32 +1,26 @@ -#![allow(dead_code)] - use std::io; use clap::Command; use clap_complete::{Generator, Shell}; -// Re-export `Instant` at the crate root so public APIs that expose it -// (e.g. `Route::get_peer_info_last_update_time`) reference a deliberate -// public type rather than leaking an inaccessible one. -pub use quanta::Instant; - mod arch; mod gateway; +mod host_runtime; pub mod instance; -mod peer_center; mod vpn_portal; pub mod common; -pub mod connector; +#[cfg(feature = "management")] pub mod core; -pub mod instance_manager; -pub mod launcher; -pub mod peers; pub mod proto; +#[cfg(feature = "management-rpc")] pub mod rpc_service; +#[cfg(feature = "management")] pub mod service_manager; +pub(crate) mod socket; pub mod tunnel; pub mod utils; +#[cfg(feature = "management")] pub mod web_client; #[cfg(test)] diff --git a/easytier/src/peer_center/mod.rs b/easytier/src/peer_center/mod.rs deleted file mode 100644 index 690f2074..00000000 --- a/easytier/src/peer_center/mod.rs +++ /dev/null @@ -1,52 +0,0 @@ -// peer_center is used to collect peer info into one peer node. -// the center node is selected with the following rules: -// 1. has smallest peer id -// 2. TODO: has allow_to_be_center peer feature -// peer center is not guaranteed to be stable and can be changed when peer enter or leave. -// it's used to reduce the cost to exchange infos between peers. - -use std::collections::BTreeMap; - -use crate::proto::api::instance::PeerInfo; -use crate::proto::peer_rpc::{DirectConnectedPeerInfo, PeerInfoForGlobalMap}; - -pub mod instance; -mod server; - -#[derive(thiserror::Error, Debug, serde::Deserialize, serde::Serialize)] -pub enum Error { - #[error("Digest not match, need provide full peer info to center server.")] - DigestMismatch, - #[error("Not center server")] - NotCenterServer, - #[error("Instance shutdown")] - Shutdown, -} - -pub type Digest = u64; - -impl From> for PeerInfoForGlobalMap { - fn from(peers: Vec) -> Self { - let mut peer_map = BTreeMap::new(); - for peer in peers { - let Some(min_lat) = peer - .conns - .iter() - .map(|conn| conn.stats.as_ref().unwrap().latency_us) - .min() - else { - continue; - }; - - let dp_info = DirectConnectedPeerInfo { - latency_ms: std::cmp::max(1, (min_lat as u32 / 1000) as i32), - }; - - // sort conn info so hash result is stable - peer_map.insert(peer.peer_id, dp_info); - } - PeerInfoForGlobalMap { - direct_peers: peer_map, - } - } -} diff --git a/easytier/src/peers/credential_manager.rs b/easytier/src/peers/credential_manager.rs deleted file mode 100644 index c6e35645..00000000 --- a/easytier/src/peers/credential_manager.rs +++ /dev/null @@ -1,667 +0,0 @@ -use std::{ - collections::HashMap, - path::PathBuf, - sync::Mutex, - time::{Duration, SystemTime, UNIX_EPOCH}, -}; - -use base64::Engine; -use base64::engine::general_purpose::STANDARD as BASE64_STANDARD; -use serde::{Deserialize, Serialize}; -use x25519_dalek::{PublicKey, StaticSecret}; - -use crate::proto::peer_rpc::{TrustedCredentialPubkey, TrustedCredentialPubkeyProof}; - -fn default_true() -> bool { - true -} - -fn current_unix_timestamp() -> i64 { - SystemTime::now() - .duration_since(UNIX_EPOCH) - .unwrap() - .as_secs() as i64 -} - -#[derive(Debug, Clone, Serialize, Deserialize)] -struct CredentialEntry { - pubkey: String, - #[serde(default)] - secret: String, - groups: Vec, - allow_relay: bool, - allowed_proxy_cidrs: Vec, - #[serde(default = "default_true")] - reusable: bool, - expiry_unix: i64, - created_at_unix: i64, -} - -impl CredentialEntry { - fn is_active_at(&self, now: i64) -> bool { - self.expiry_unix > now - } - - fn to_trusted_credential(&self) -> Option { - Some(TrustedCredentialPubkey { - pubkey: CredentialManager::decode_pubkey_b64(&self.pubkey)?, - groups: self.groups.clone(), - allow_relay: self.allow_relay, - expiry_unix: self.expiry_unix, - allowed_proxy_cidrs: self.allowed_proxy_cidrs.clone(), - reusable: Some(self.reusable), - }) - } - - fn to_api_credential_info( - &self, - credential_id: &str, - ) -> crate::proto::api::instance::CredentialInfo { - crate::proto::api::instance::CredentialInfo { - credential_id: credential_id.to_string(), - groups: self.groups.clone(), - allow_relay: self.allow_relay, - expiry_unix: self.expiry_unix, - allowed_proxy_cidrs: self.allowed_proxy_cidrs.clone(), - reusable: Some(self.reusable), - } - } -} - -pub struct CredentialManager { - credentials: Mutex>, - storage_path: Option, -} - -impl CredentialManager { - pub fn new(storage_path: Option) -> Self { - let mgr = CredentialManager { - credentials: Mutex::new(HashMap::new()), - storage_path, - }; - mgr.load_from_disk(); - mgr - } - - pub fn generate_credential( - &self, - groups: Vec, - allow_relay: bool, - allowed_proxy_cidrs: Vec, - ttl: Duration, - ) -> (String, String) { - self.generate_credential_with_options( - groups, - allow_relay, - allowed_proxy_cidrs, - ttl, - None, - true, - ) - } - - pub fn generate_credential_with_id( - &self, - groups: Vec, - allow_relay: bool, - allowed_proxy_cidrs: Vec, - ttl: Duration, - credential_id: Option, - ) -> (String, String) { - self.generate_credential_with_options( - groups, - allow_relay, - allowed_proxy_cidrs, - ttl, - credential_id, - true, - ) - } - - pub fn generate_credential_with_options( - &self, - groups: Vec, - allow_relay: bool, - allowed_proxy_cidrs: Vec, - ttl: Duration, - credential_id: Option, - reusable: bool, - ) -> (String, String) { - self.remove_expired_credentials(); - - let mut credentials = self.credentials.lock().unwrap(); - let id = if let Some(id) = credential_id - .map(|x| x.trim().to_string()) - .filter(|x| !x.is_empty()) - { - if let Some(existing) = credentials.get(&id) - && !existing.secret.is_empty() - { - return (id, existing.secret.clone()); - } - id - } else { - uuid::Uuid::new_v4().to_string() - }; - - let (entry, secret) = - Self::build_entry(groups, allow_relay, allowed_proxy_cidrs, reusable, ttl); - credentials.insert(id.clone(), entry); - drop(credentials); - self.save_to_disk(); - (id, secret) - } - - fn build_entry( - groups: Vec, - allow_relay: bool, - allowed_proxy_cidrs: Vec, - reusable: bool, - ttl: Duration, - ) -> (CredentialEntry, String) { - let private = StaticSecret::random_from_rng(rand::rngs::OsRng); - let public = PublicKey::from(&private); - let pubkey = BASE64_STANDARD.encode(public.as_bytes()); - let secret = BASE64_STANDARD.encode(private.as_bytes()); - - let now = SystemTime::now() - .duration_since(UNIX_EPOCH) - .unwrap() - .as_secs() as i64; - let expiry_unix = now + ttl.as_secs() as i64; - - let entry = CredentialEntry { - pubkey, - secret: secret.clone(), - groups, - allow_relay, - allowed_proxy_cidrs, - reusable, - expiry_unix, - created_at_unix: now, - }; - (entry, secret) - } - - pub fn revoke_credential(&self, credential_id: &str) -> bool { - let removed = self - .credentials - .lock() - .unwrap() - .remove(credential_id) - .is_some(); - if removed { - self.save_to_disk(); - } - removed - } - - pub fn remove_expired_credentials(&self) -> bool { - self.remove_expired_credentials_at(current_unix_timestamp()) - } - - fn remove_expired_credentials_at(&self, now: i64) -> bool { - let removed = { - let mut credentials = self.credentials.lock().unwrap(); - let before = credentials.len(); - credentials.retain(|_, entry| entry.is_active_at(now)); - before != credentials.len() - }; - - if removed { - self.save_to_disk(); - } - - removed - } - - pub fn get_trusted_pubkeys(&self, network_secret: &str) -> Vec { - let now = current_unix_timestamp(); - - self.credentials - .lock() - .unwrap() - .values() - .filter(|entry| entry.is_active_at(now)) - .filter_map(|entry| { - entry.to_trusted_credential().map(|credential| { - TrustedCredentialPubkeyProof::new_signed(credential, network_secret) - }) - }) - .collect() - } - - pub fn is_pubkey_trusted(&self, pubkey: &[u8]) -> bool { - let now = current_unix_timestamp(); - - let encoded = BASE64_STANDARD.encode(pubkey); - self.credentials - .lock() - .unwrap() - .values() - .any(|entry| entry.pubkey == encoded && entry.is_active_at(now)) - } - - pub fn list_credentials(&self) -> Vec { - let now = current_unix_timestamp(); - - self.credentials - .lock() - .unwrap() - .iter() - .filter(|(_, entry)| entry.is_active_at(now)) - .map(|(id, entry)| entry.to_api_credential_info(id)) - .collect() - } - - fn save_to_disk(&self) { - let Some(path) = &self.storage_path else { - return; - }; - let creds = self.credentials.lock().unwrap(); - if let Ok(json) = serde_json::to_string_pretty(&*creds) - && let Err(e) = std::fs::write(path, json) - { - tracing::warn!(?e, "failed to save credentials to disk"); - } - } - - fn load_from_disk(&self) { - let Some(path) = &self.storage_path else { - return; - }; - let Ok(data) = std::fs::read_to_string(path) else { - return; - }; - match serde_json::from_str::>(&data) { - Ok(loaded) => { - *self.credentials.lock().unwrap() = loaded; - tracing::info!("loaded credentials from {}", path.display()); - } - Err(e) => { - tracing::warn!(?e, "failed to parse credentials file"); - } - } - } - - fn decode_pubkey_b64(s: &str) -> Option> { - let decoded = BASE64_STANDARD.decode(s).ok()?; - if decoded.len() != 32 { - return None; - } - Some(decoded) - } -} - -#[cfg(test)] -mod tests { - use super::*; - - #[test] - fn test_generate_and_revoke() { - let mgr = CredentialManager::new(None); - let (id, secret) = mgr.generate_credential( - vec!["guest".to_string()], - false, - vec![], - Duration::from_secs(3600), - ); - - assert!(!id.is_empty()); - assert!(!secret.is_empty()); - assert!(uuid::Uuid::parse_str(&id).is_ok()); - - let privkey_bytes: [u8; 32] = BASE64_STANDARD.decode(&secret).unwrap().try_into().unwrap(); - let private = StaticSecret::from(privkey_bytes); - let pubkey_bytes = PublicKey::from(&private).as_bytes().to_vec(); - assert!(mgr.is_pubkey_trusted(&pubkey_bytes)); - - let trusted = mgr.get_trusted_pubkeys("sec"); - assert_eq!(trusted.len(), 1); - assert_eq!( - trusted[0].credential.as_ref().unwrap().groups, - vec!["guest".to_string()] - ); - assert_eq!(trusted[0].credential.as_ref().unwrap().reusable, Some(true)); - - assert!(mgr.revoke_credential(&id)); - assert!(!mgr.is_pubkey_trusted(&pubkey_bytes)); - assert!(mgr.get_trusted_pubkeys("sec").is_empty()); - } - - #[test] - fn test_expired_credential() { - let mgr = CredentialManager::new(None); - // TTL of 0 seconds - immediately expired - let (_, secret) = mgr.generate_credential(vec![], false, vec![], Duration::from_secs(0)); - - let privkey_bytes: [u8; 32] = BASE64_STANDARD.decode(&secret).unwrap().try_into().unwrap(); - let private = StaticSecret::from(privkey_bytes); - let pubkey_bytes = PublicKey::from(&private).as_bytes().to_vec(); - assert!(!mgr.is_pubkey_trusted(&pubkey_bytes)); - assert!(mgr.get_trusted_pubkeys("sec").is_empty()); - } - - #[test] - fn test_list_credentials() { - let mgr = CredentialManager::new(None); - mgr.generate_credential( - vec!["a".to_string()], - true, - vec!["10.0.0.0/24".to_string()], - Duration::from_secs(3600), - ); - mgr.generate_credential(vec![], false, vec![], Duration::from_secs(3600)); - - let list = mgr.list_credentials(); - assert_eq!(list.len(), 2); - assert!(list.iter().all(|item| item.reusable == Some(true))); - } - - #[test] - fn test_keypair_validity() { - // Verify the generated private key can derive the same public key - let mgr = CredentialManager::new(None); - let (id, secret) = - mgr.generate_credential(vec![], false, vec![], Duration::from_secs(3600)); - - let privkey_bytes: [u8; 32] = BASE64_STANDARD.decode(&secret).unwrap().try_into().unwrap(); - let private = StaticSecret::from(privkey_bytes); - let derived_public = PublicKey::from(&private); - assert!(uuid::Uuid::parse_str(&id).is_ok()); - assert!(mgr.is_pubkey_trusted(derived_public.as_bytes())); - } - - #[test] - fn test_revoke_nonexistent() { - let mgr = CredentialManager::new(None); - assert!(!mgr.revoke_credential("nonexistent_id")); - } - - #[test] - fn test_multiple_credentials_independent() { - let mgr = CredentialManager::new(None); - let (id1, secret1) = mgr.generate_credential( - vec!["group1".to_string()], - false, - vec![], - Duration::from_secs(3600), - ); - let (_id2, secret2) = mgr.generate_credential( - vec!["group2".to_string()], - true, - vec!["10.0.0.0/8".to_string()], - Duration::from_secs(3600), - ); - - let sk1: [u8; 32] = BASE64_STANDARD - .decode(&secret1) - .unwrap() - .try_into() - .unwrap(); - let sk2: [u8; 32] = BASE64_STANDARD - .decode(&secret2) - .unwrap() - .try_into() - .unwrap(); - let pk1 = PublicKey::from(&StaticSecret::from(sk1)) - .as_bytes() - .to_vec(); - let pk2 = PublicKey::from(&StaticSecret::from(sk2)) - .as_bytes() - .to_vec(); - - assert!(mgr.is_pubkey_trusted(&pk1)); - assert!(mgr.is_pubkey_trusted(&pk2)); - - // Revoke first, second should still be trusted - mgr.revoke_credential(&id1); - assert!(!mgr.is_pubkey_trusted(&pk1)); - assert!(mgr.is_pubkey_trusted(&pk2)); - - let trusted = mgr.get_trusted_pubkeys("sec"); - assert_eq!(trusted.len(), 1); - assert_eq!( - trusted[0].credential.as_ref().unwrap().groups, - vec!["group2".to_string()] - ); - assert!(trusted[0].credential.as_ref().unwrap().allow_relay); - assert_eq!( - trusted[0].credential.as_ref().unwrap().allowed_proxy_cidrs, - vec!["10.0.0.0/8".to_string()] - ); - assert_eq!(trusted[0].credential.as_ref().unwrap().reusable, Some(true)); - } - - #[test] - fn test_trusted_pubkeys_include_metadata() { - let mgr = CredentialManager::new(None); - let (_, secret) = mgr.generate_credential( - vec!["admin".to_string(), "ops".to_string()], - true, - vec!["192.168.0.0/16".to_string(), "10.0.0.0/8".to_string()], - Duration::from_secs(7200), - ); - - let trusted = mgr.get_trusted_pubkeys("sec"); - assert_eq!(trusted.len(), 1); - let tc = &trusted[0]; - assert_eq!( - tc.credential.as_ref().unwrap().groups, - vec!["admin".to_string(), "ops".to_string()] - ); - assert!(tc.credential.as_ref().unwrap().allow_relay); - assert_eq!( - tc.credential.as_ref().unwrap().allowed_proxy_cidrs, - vec!["192.168.0.0/16".to_string(), "10.0.0.0/8".to_string()] - ); - assert_eq!(tc.credential.as_ref().unwrap().reusable, Some(true)); - assert!(tc.credential.as_ref().unwrap().expiry_unix > 0); - assert!(tc.verify_credential_hmac("sec")); - assert!( - tc.credential - .as_ref() - .map(|x| !x.pubkey.is_empty()) - .unwrap_or(false) - ); - - let sk: [u8; 32] = BASE64_STANDARD.decode(&secret).unwrap().try_into().unwrap(); - let pk = PublicKey::from(&StaticSecret::from(sk)).as_bytes().to_vec(); - assert_eq!(tc.credential.as_ref().unwrap().pubkey, pk); - } - - #[test] - fn test_unknown_pubkey_not_trusted() { - let mgr = CredentialManager::new(None); - mgr.generate_credential(vec![], false, vec![], Duration::from_secs(3600)); - - let random_key = [42u8; 32]; - assert!(!mgr.is_pubkey_trusted(&random_key)); - } - - #[test] - fn test_persistence_roundtrip() { - let dir = tempfile::tempdir().unwrap(); - let path = dir.path().join("creds.json"); - - // Create and save - { - let mgr = CredentialManager::new(Some(path.clone())); - mgr.generate_credential( - vec!["persist_group".to_string()], - true, - vec!["10.0.0.0/24".to_string()], - Duration::from_secs(3600), - ); - assert_eq!(mgr.list_credentials().len(), 1); - } - - // Load from disk - { - let mgr = CredentialManager::new(Some(path)); - let list = mgr.list_credentials(); - assert_eq!(list.len(), 1); - assert_eq!(list[0].groups, vec!["persist_group".to_string()]); - assert!(list[0].allow_relay); - assert_eq!(list[0].reusable, Some(true)); - } - } - - #[test] - fn test_list_credentials_filters_expired() { - let mgr = CredentialManager::new(None); - mgr.generate_credential(vec![], false, vec![], Duration::from_secs(3600)); - mgr.generate_credential(vec![], false, vec![], Duration::from_secs(0)); // expired - - let list = mgr.list_credentials(); - assert_eq!(list.len(), 1); - } - - #[test] - fn test_remove_expired_credentials_removes_and_persists() { - let dir = tempfile::tempdir().unwrap(); - let path = dir.path().join("creds.json"); - let mgr = CredentialManager::new(Some(path.clone())); - mgr.generate_credential_with_id( - vec!["active".to_string()], - false, - vec![], - Duration::from_secs(3600), - Some("active-id".to_string()), - ); - mgr.generate_credential_with_id( - vec!["expired".to_string()], - false, - vec![], - Duration::from_secs(0), - Some("expired-id".to_string()), - ); - - assert!(mgr.remove_expired_credentials()); - assert_eq!(mgr.list_credentials().len(), 1); - - let reloaded = CredentialManager::new(Some(path)); - let list = reloaded.list_credentials(); - assert_eq!(list.len(), 1); - assert_eq!(list[0].credential_id, "active-id"); - } - - #[test] - fn test_generate_with_specified_id_reuses_existing_result() { - let mgr = CredentialManager::new(None); - let fixed_id = "fixed-credential-id".to_string(); - let (id1, secret1) = mgr.generate_credential_with_id( - vec!["group-a".to_string()], - false, - vec!["10.0.0.0/24".to_string()], - Duration::from_secs(3600), - Some(fixed_id.clone()), - ); - let (id2, secret2) = mgr.generate_credential_with_id( - vec!["group-b".to_string()], - true, - vec!["192.168.0.0/16".to_string()], - Duration::from_secs(7200), - Some(fixed_id.clone()), - ); - - assert_eq!(id1, fixed_id); - assert_eq!(id2, fixed_id); - assert_eq!(secret1, secret2); - - let list = mgr.list_credentials(); - assert_eq!(list.len(), 1); - assert_eq!(list[0].credential_id, fixed_id); - assert_eq!(list[0].groups, vec!["group-a".to_string()]); - assert!(!list[0].allow_relay); - assert_eq!(list[0].allowed_proxy_cidrs, vec!["10.0.0.0/24".to_string()]); - assert_eq!(list[0].reusable, Some(true)); - } - - #[test] - fn test_generate_with_specified_id_replaces_expired_existing_result() { - let mgr = CredentialManager::new(None); - let fixed_id = "fixed-credential-id".to_string(); - let (id1, secret1) = mgr.generate_credential_with_id( - vec!["expired".to_string()], - false, - vec![], - Duration::from_secs(0), - Some(fixed_id.clone()), - ); - let (id2, secret2) = mgr.generate_credential_with_id( - vec!["fresh".to_string()], - true, - vec!["10.0.0.0/24".to_string()], - Duration::from_secs(3600), - Some(fixed_id.clone()), - ); - - assert_eq!(id1, fixed_id); - assert_eq!(id2, fixed_id); - assert_ne!(secret1, secret2); - - let list = mgr.list_credentials(); - assert_eq!(list.len(), 1); - assert_eq!(list[0].credential_id, fixed_id); - assert_eq!(list[0].groups, vec!["fresh".to_string()]); - assert!(list[0].allow_relay); - assert_eq!(list[0].allowed_proxy_cidrs, vec!["10.0.0.0/24".to_string()]); - } - - #[test] - fn test_generate_non_reusable_credential() { - let mgr = CredentialManager::new(None); - let (_id, secret) = mgr.generate_credential_with_options( - vec!["single".to_string()], - false, - vec![], - Duration::from_secs(3600), - None, - false, - ); - - let privkey_bytes: [u8; 32] = BASE64_STANDARD.decode(&secret).unwrap().try_into().unwrap(); - let private = StaticSecret::from(privkey_bytes); - let pubkey_bytes = PublicKey::from(&private).as_bytes().to_vec(); - - let listed = mgr.list_credentials(); - assert_eq!(listed.len(), 1); - assert_eq!(listed[0].reusable, Some(false)); - assert!(mgr.is_pubkey_trusted(&pubkey_bytes)); - - let trusted = mgr.get_trusted_pubkeys("sec"); - assert_eq!(trusted.len(), 1); - assert_eq!( - trusted[0].credential.as_ref().unwrap().reusable, - Some(false) - ); - } - - #[test] - fn test_load_old_credentials_default_to_reusable() { - let dir = tempfile::tempdir().unwrap(); - let path = dir.path().join("legacy-creds.json"); - std::fs::write( - &path, - r#"{ - "legacy-id": { - "pubkey": "AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA=", - "secret": "BBBBBBBBBBBBBBBBBBBBBBBBBBBBBBBBBBBBBBBBBBB=", - "groups": ["legacy"], - "allow_relay": false, - "allowed_proxy_cidrs": [], - "expiry_unix": 4102444800, - "created_at_unix": 1700000000 - } -}"#, - ) - .unwrap(); - - let mgr = CredentialManager::new(Some(path)); - let list = mgr.list_credentials(); - assert_eq!(list.len(), 1); - assert_eq!(list[0].credential_id, "legacy-id"); - assert_eq!(list[0].reusable, Some(true)); - } -} diff --git a/easytier/src/peers/encrypt/mod.rs b/easytier/src/peers/encrypt/mod.rs deleted file mode 100644 index 8d8f2d38..00000000 --- a/easytier/src/peers/encrypt/mod.rs +++ /dev/null @@ -1,107 +0,0 @@ -use crate::{ - common::{config::EncryptionAlgorithm, log}, - tunnel::packet_def::ZCPacket, -}; -use std::sync::Arc; - -#[cfg(feature = "wireguard")] -pub mod ring; - -#[cfg(feature = "aes-gcm")] -pub mod aes_gcm; - -#[cfg(feature = "openssl-crypto")] -pub mod openssl; - -pub mod xor; - -#[derive(thiserror::Error, Debug)] -pub enum Error { - #[error("packet is too short. len: {0}")] - PacketTooShort(usize), - #[error("decryption failed")] - DecryptionFailed, - #[error("encryption failed")] - EncryptionFailed, - #[error("invalid tag. tag: {0:?}")] - InvalidTag(Vec), -} - -pub trait Encryptor: Send + Sync + 'static { - fn decrypt(&self, zc_packet: &mut ZCPacket) -> Result<(), Error>; - fn encrypt(&self, zc_packet: &mut ZCPacket) -> Result<(), Error>; - fn encrypt_with_nonce( - &self, - zc_packet: &mut ZCPacket, - _nonce: Option<&[u8]>, - ) -> Result<(), Error> { - self.encrypt(zc_packet) - } -} - -pub struct NullCipher; - -impl Encryptor for NullCipher { - fn decrypt(&self, zc_packet: &mut ZCPacket) -> Result<(), Error> { - let pm_header = zc_packet.peer_manager_header().unwrap(); - if pm_header.is_encrypted() { - Err(Error::DecryptionFailed) - } else { - Ok(()) - } - } - - fn encrypt(&self, _zc_packet: &mut ZCPacket) -> Result<(), Error> { - Ok(()) - } -} - -/// Create an encryptor based on the algorithm name -pub fn create_encryptor( - algorithm: &str, - key_128: [u8; 16], - #[allow(unused_variables)] key_256: [u8; 32], -) -> Arc { - let algorithm = match EncryptionAlgorithm::try_from(algorithm) { - Ok(algorithm) => algorithm, - Err(_) => { - let default = EncryptionAlgorithm::default(); - log::warn!( - "Unknown encryption algorithm: {}, falling back to default {}", - algorithm, - default - ); - default - } - }; - - match algorithm { - EncryptionAlgorithm::Xor => Arc::new(xor::XorCipher::new(&key_128)), - - #[cfg(any(feature = "aes-gcm", feature = "wireguard", feature = "openssl-crypto"))] - EncryptionAlgorithm::AesGcm => { - cfg_select! { - feature = "openssl-crypto" => Arc::new(openssl::OpenSslCipher::new_aes128_gcm(key_128)), - feature = "wireguard" => Arc::new(ring::RingCipher::new_aes128_gcm(key_128)), - feature = "aes-gcm" => Arc::new(aes_gcm::AesGcmCipher::new_128(key_128)), - } - } - - #[cfg(any(feature = "aes-gcm", feature = "wireguard", feature = "openssl-crypto"))] - EncryptionAlgorithm::Aes256Gcm => { - cfg_select! { - feature = "openssl-crypto" => Arc::new(openssl::OpenSslCipher::new_aes256_gcm(key_256)), - feature = "wireguard" => Arc::new(ring::RingCipher::new_aes256_gcm(key_256)), - feature = "aes-gcm" => Arc::new(aes_gcm::AesGcmCipher::new_256(key_256)), - } - } - - #[cfg(any(feature = "wireguard", feature = "openssl-crypto"))] - EncryptionAlgorithm::ChaCha20 => { - cfg_select! { - feature = "openssl-crypto" => Arc::new(openssl::OpenSslCipher::new_chacha20(key_256)), - feature = "wireguard" => Arc::new(ring::RingCipher::new_chacha20(key_256)), - } - } - } -} diff --git a/easytier/src/peers/encrypt/openssl.rs b/easytier/src/peers/encrypt/openssl.rs deleted file mode 100644 index f8bb8ec3..00000000 --- a/easytier/src/peers/encrypt/openssl.rs +++ /dev/null @@ -1,201 +0,0 @@ -use crate::tunnel::packet_def::{StandardAeadTail, ZCPacket}; -use openssl::symm::{Cipher, Crypter, Mode}; -use rand::RngCore; -use zerocopy::{AsBytes, FromBytes, FromZeroes}; - -use crate::peers::encrypt::{Encryptor, Error}; - -#[derive(Clone)] -pub struct OpenSslCipher { - pub(crate) cipher: OpenSslEnum, -} - -#[derive(Clone, Copy)] -pub enum OpenSslEnum { - Aes128Gcm([u8; 16]), - Aes256Gcm([u8; 32]), - ChaCha20([u8; 32]), -} - -impl OpenSslCipher { - pub fn new_aes128_gcm(key: [u8; 16]) -> Self { - Self { - cipher: OpenSslEnum::Aes128Gcm(key), - } - } - - pub fn new_aes256_gcm(key: [u8; 32]) -> Self { - Self { - cipher: OpenSslEnum::Aes256Gcm(key), - } - } - - pub fn new_chacha20(key: [u8; 32]) -> Self { - Self { - cipher: OpenSslEnum::ChaCha20(key), - } - } - - fn get_cipher_and_key(&self) -> (Cipher, &[u8]) { - match &self.cipher { - OpenSslEnum::Aes128Gcm(key) => (Cipher::aes_128_gcm(), key.as_slice()), - OpenSslEnum::Aes256Gcm(key) => (Cipher::aes_256_gcm(), key.as_slice()), - OpenSslEnum::ChaCha20(key) => (Cipher::chacha20_poly1305(), key.as_slice()), - } - } -} - -impl Encryptor for OpenSslCipher { - fn decrypt(&self, zc_packet: &mut ZCPacket) -> Result<(), Error> { - let pm_header = zc_packet.peer_manager_header().unwrap(); - if !pm_header.is_encrypted() { - return Ok(()); - } - - let payload = zc_packet.payload(); - let len = payload.len(); - if len < StandardAeadTail::SIZE { - return Err(Error::PacketTooShort(len)); - } - - let (cipher, key) = self.get_cipher_and_key(); - - // 提取 nonce/IV 和 tag - let tail = StandardAeadTail::ref_from_suffix(payload).unwrap(); - - let mut decrypter = Crypter::new(cipher, Mode::Decrypt, key, Some(&tail.nonce)) - .map_err(|_| Error::DecryptionFailed)?; - - decrypter - .set_tag(&tail.tag) - .map_err(|_| Error::DecryptionFailed)?; - - let text_len = len - StandardAeadTail::SIZE; - let mut output = vec![0u8; text_len + cipher.block_size()]; - let mut count = decrypter - .update(&payload[..text_len], &mut output) - .map_err(|_| Error::DecryptionFailed)?; - - count += decrypter - .finalize(&mut output[count..]) - .map_err(|_| Error::DecryptionFailed)?; - - // 更新数据包 - zc_packet.mut_payload()[..count].copy_from_slice(&output[..count]); - let pm_header = zc_packet.mut_peer_manager_header().unwrap(); - pm_header.set_encrypted(false); - - let len = zc_packet.buf_len() - (len - count); - zc_packet.mut_inner().truncate(len); - - Ok(()) - } - - fn encrypt(&self, zc_packet: &mut ZCPacket) -> Result<(), Error> { - self.encrypt_with_nonce(zc_packet, None) - } - - fn encrypt_with_nonce( - &self, - zc_packet: &mut ZCPacket, - nonce: Option<&[u8]>, - ) -> Result<(), Error> { - let pm_header = zc_packet.peer_manager_header().unwrap(); - if pm_header.is_encrypted() { - tracing::warn!(?zc_packet, "packet is already encrypted"); - return Ok(()); - } - - let (cipher, key) = self.get_cipher_and_key(); - - let mut tail = StandardAeadTail::new_zeroed(); - if let Some(nonce) = nonce { - if nonce.len() != StandardAeadTail::NONCE_SIZE { - return Err(Error::EncryptionFailed); - } - tail.nonce.copy_from_slice(nonce); - } else { - rand::thread_rng().fill_bytes(&mut tail.nonce); - } - - let mut encrypter = Crypter::new(cipher, Mode::Encrypt, key, Some(&tail.nonce)) - .map_err(|_| Error::EncryptionFailed)?; - - let payload_len = zc_packet.payload().len(); - let mut output = vec![0u8; payload_len + cipher.block_size()]; - - let mut count = encrypter - .update(zc_packet.payload(), &mut output) - .map_err(|_| Error::EncryptionFailed)?; - - count += encrypter - .finalize(&mut output[count..]) - .map_err(|_| Error::EncryptionFailed)?; - - // 更新数据包内容 - zc_packet.mut_payload()[..count].copy_from_slice(&output[..count]); - - encrypter - .get_tag(&mut tail.tag) - .map_err(|_| Error::EncryptionFailed)?; - - // 添加 nonce/IV & tag 的结构 - zc_packet.mut_inner().extend_from_slice(tail.as_bytes()); - - let pm_header = zc_packet.mut_peer_manager_header().unwrap(); - pm_header.set_encrypted(true); - - Ok(()) - } -} - -#[cfg(test)] -mod tests { - use super::*; - - fn run_cipher_test_with_nonce(cipher: OpenSslCipher) { - let text = b"Hello, World! This is a standardized test message."; - let nonce: [u8; 12] = [101, 102, 103, 104, 105, 106, 107, 108, 109, 110, 111, 112]; - - let mut packet = ZCPacket::new_with_payload(text); - packet.fill_peer_manager_hdr(0, 0, 0); - - cipher - .encrypt_with_nonce(&mut packet, Some(&nonce)) - .unwrap(); - - let payload = packet.payload(); - let len = payload.len(); - - assert!(len > text.len() + StandardAeadTail::SIZE - 1); - assert!(packet.peer_manager_header().unwrap().is_encrypted()); - - let tail = StandardAeadTail::ref_from_suffix(payload).unwrap().clone(); - assert_eq!(tail.nonce, nonce); - - cipher.decrypt(&mut packet).unwrap(); - assert_eq!(packet.payload(), text); - assert!(!packet.peer_manager_header().unwrap().is_encrypted()); - } - - #[test] - fn test_openssl_aes128_gcm() { - let key = [1u8; 16]; - let cipher = OpenSslCipher::new_aes128_gcm(key); - run_cipher_test_with_nonce(cipher); - } - - #[test] - fn test_openssl_aes256_gcm() { - let key = [2u8; 32]; - let cipher = OpenSslCipher::new_aes256_gcm(key); - run_cipher_test_with_nonce(cipher); - } - - #[test] - fn test_openssl_chacha20() { - let key = [3u8; 32]; - let cipher = OpenSslCipher::new_chacha20(key); - run_cipher_test_with_nonce(cipher); - } -} diff --git a/easytier/src/peers/encrypt/ring.rs b/easytier/src/peers/encrypt/ring.rs deleted file mode 100644 index b5177eb7..00000000 --- a/easytier/src/peers/encrypt/ring.rs +++ /dev/null @@ -1,252 +0,0 @@ -use rand::RngCore; -use ring::aead::{self}; -use ring::aead::{LessSafeKey, UnboundKey}; -use zerocopy::{AsBytes, FromBytes, FromZeroes}; - -use crate::tunnel::packet_def::{StandardAeadTail, ZCPacket}; - -use super::{Encryptor, Error}; - -#[derive(Clone)] -pub struct RingCipher { - pub(crate) cipher: RingEnum, -} - -pub enum RingEnum { - Aes128Gcm(LessSafeKey, [u8; 16]), - Aes256Gcm(LessSafeKey, [u8; 32]), - ChaCha20(LessSafeKey, [u8; 32]), -} - -impl RingEnum { - fn get_cipher(&self) -> &LessSafeKey { - match &self { - RingEnum::Aes128Gcm(cipher, _) => cipher, - RingEnum::Aes256Gcm(cipher, _) => cipher, - RingEnum::ChaCha20(cipher, _) => cipher, - } - } -} - -impl Clone for RingEnum { - fn clone(&self) -> Self { - match &self { - RingEnum::Aes128Gcm(_, key) => { - let c = - LessSafeKey::new(UnboundKey::new(&aead::AES_128_GCM, key.as_slice()).unwrap()); - RingEnum::Aes128Gcm(c, *key) - } - RingEnum::Aes256Gcm(_, key) => { - let c = - LessSafeKey::new(UnboundKey::new(&aead::AES_256_GCM, key.as_slice()).unwrap()); - RingEnum::Aes256Gcm(c, *key) - } - RingEnum::ChaCha20(_, key) => { - let c = LessSafeKey::new( - UnboundKey::new(&aead::CHACHA20_POLY1305, key.as_slice()).unwrap(), - ); - RingEnum::ChaCha20(c, *key) - } - } - } -} - -impl RingCipher { - pub fn new_aes128_gcm(key: [u8; 16]) -> Self { - let cipher = LessSafeKey::new(UnboundKey::new(&aead::AES_128_GCM, &key).unwrap()); - Self { - cipher: RingEnum::Aes128Gcm(cipher, key), - } - } - - pub fn new_aes256_gcm(key: [u8; 32]) -> Self { - let cipher = LessSafeKey::new(UnboundKey::new(&aead::AES_256_GCM, &key).unwrap()); - Self { - cipher: RingEnum::Aes256Gcm(cipher, key), - } - } - - pub fn new_chacha20(key: [u8; 32]) -> Self { - let unbound_key = UnboundKey::new(&aead::CHACHA20_POLY1305, &key).unwrap(); - let cipher = LessSafeKey::new(unbound_key); - Self { - cipher: RingEnum::ChaCha20(cipher, key), - } - } -} - -impl Encryptor for RingCipher { - fn decrypt(&self, zc_packet: &mut ZCPacket) -> Result<(), Error> { - let pm_header = zc_packet.peer_manager_header().unwrap(); - if !pm_header.is_encrypted() { - return Ok(()); - } - - let payload_len = zc_packet.payload().len(); - if payload_len < StandardAeadTail::SIZE { - return Err(Error::PacketTooShort(zc_packet.payload().len())); - } - - let text_and_tag_len = payload_len - StandardAeadTail::SIZE + StandardAeadTail::TAG_SIZE; - - let aes_tail = StandardAeadTail::ref_from_suffix(zc_packet.payload()).unwrap(); - let nonce = aead::Nonce::assume_unique_for_key(aes_tail.nonce); - - self.cipher - .get_cipher() - .open_in_place( - nonce, - aead::Aad::empty(), - &mut zc_packet.mut_payload()[..text_and_tag_len], - ) - .map_err(|_| Error::DecryptionFailed)?; - - let pm_header = zc_packet.mut_peer_manager_header().unwrap(); - pm_header.set_encrypted(false); - let old_len = zc_packet.buf_len(); - zc_packet - .mut_inner() - .truncate(old_len - StandardAeadTail::SIZE); - Ok(()) - } - - fn encrypt(&self, zc_packet: &mut ZCPacket) -> Result<(), Error> { - self.encrypt_with_nonce(zc_packet, None) - } - - fn encrypt_with_nonce( - &self, - zc_packet: &mut ZCPacket, - nonce: Option<&[u8]>, - ) -> Result<(), Error> { - let pm_header = zc_packet.peer_manager_header().unwrap(); - if pm_header.is_encrypted() { - tracing::warn!(?zc_packet, "packet is already encrypted"); - return Ok(()); - } - - let mut tail = StandardAeadTail::new_zeroed(); - - match nonce { - Some(n) => tail.nonce = n.try_into().map_err(|_| Error::EncryptionFailed)?, - None => rand::thread_rng().fill_bytes(&mut tail.nonce), - } - let nonce = aead::Nonce::assume_unique_for_key(tail.nonce); - - let tag = self - .cipher - .get_cipher() - .seal_in_place_separate_tag(nonce, aead::Aad::empty(), zc_packet.mut_payload()) - .map_err(|_| Error::EncryptionFailed)?; - - let tag = tag.as_ref(); - if tag.len() != StandardAeadTail::TAG_SIZE { - return Err(Error::InvalidTag(tag.to_vec())); - } - tail.tag.copy_from_slice(tag); - - let pm_header = zc_packet.mut_peer_manager_header().unwrap(); - pm_header.set_encrypted(true); - zc_packet.mut_inner().extend_from_slice(tail.as_bytes()); - Ok(()) - } -} - -#[cfg(test)] -mod tests { - use crate::{ - peers::encrypt::{Encryptor, ring::RingCipher}, - tunnel::packet_def::{StandardAeadTail, ZCPacket}, - }; - use zerocopy::FromBytes; - - #[test] - fn test_aes_gcm_cipher() { - let key = [0u8; 16]; - let cipher = RingCipher::new_aes128_gcm(key); - let text = b"1234567"; - let mut packet = ZCPacket::new_with_payload(text); - packet.fill_peer_manager_hdr(0, 0, 0); - cipher.encrypt(&mut packet).unwrap(); - assert_eq!(packet.payload().len(), text.len() + StandardAeadTail::SIZE); - assert!(packet.peer_manager_header().unwrap().is_encrypted()); - - cipher.decrypt(&mut packet).unwrap(); - assert_eq!(packet.payload(), text); - assert!(!packet.peer_manager_header().unwrap().is_encrypted()); - } - - #[test] - fn test_aes_gcm_cipher_with_nonce() { - let key = [7u8; 16]; - let cipher = RingCipher::new_aes128_gcm(key); - let text = b"Hello"; - let nonce = [3u8; 12]; - - let mut packet1 = ZCPacket::new_with_payload(text); - packet1.fill_peer_manager_hdr(0, 0, 0); - cipher - .encrypt_with_nonce(&mut packet1, Some(&nonce)) - .unwrap(); - - let mut packet2 = ZCPacket::new_with_payload(text); - packet2.fill_peer_manager_hdr(0, 0, 0); - cipher - .encrypt_with_nonce(&mut packet2, Some(&nonce)) - .unwrap(); - - assert_eq!(packet1.payload(), packet2.payload()); - - let tail = StandardAeadTail::ref_from_suffix(packet1.payload()).unwrap(); - assert_eq!(tail.nonce, nonce); - - cipher.decrypt(&mut packet1).unwrap(); - assert_eq!(packet1.payload(), text); - } - - #[test] - fn test_ring_chacha20_cipher() { - let key = [0u8; 32]; - let cipher = RingCipher::new_chacha20(key); - let text = b"Hello, World! This is a test message for Ring ChaCha20-Poly1305."; - let mut packet = ZCPacket::new_with_payload(text); - packet.fill_peer_manager_hdr(0, 0, 0); - - cipher.encrypt(&mut packet).unwrap(); - assert_eq!(packet.payload().len(), text.len() + StandardAeadTail::SIZE); - assert!(packet.peer_manager_header().unwrap().is_encrypted()); - - cipher.decrypt(&mut packet).unwrap(); - assert_eq!(packet.payload(), text); - assert!(!packet.peer_manager_header().unwrap().is_encrypted()); - } - - #[test] - fn test_ring_chacha20_cipher_with_nonce() { - let key = [9u8; 32]; - let cipher = RingCipher::new_chacha20(key); - let text = b"Hello"; - let nonce = [5u8; 12]; - - let mut packet1 = ZCPacket::new_with_payload(text); - packet1.fill_peer_manager_hdr(0, 0, 0); - cipher - .encrypt_with_nonce(&mut packet1, Some(&nonce)) - .unwrap(); - - let mut packet2 = ZCPacket::new_with_payload(text); - packet2.fill_peer_manager_hdr(0, 0, 0); - cipher - .encrypt_with_nonce(&mut packet2, Some(&nonce)) - .unwrap(); - - assert_eq!(packet1.payload(), packet2.payload()); - - let tail = StandardAeadTail::ref_from_suffix(packet1.payload()).unwrap(); - assert_eq!(tail.nonce, nonce); - - cipher.decrypt(&mut packet1).unwrap(); - assert_eq!(packet1.payload(), text); - assert!(!packet1.peer_manager_header().unwrap().is_encrypted()); - } -} diff --git a/easytier/src/peers/foreign_network_manager.rs b/easytier/src/peers/foreign_network_manager.rs deleted file mode 100644 index 9a9fb641..00000000 --- a/easytier/src/peers/foreign_network_manager.rs +++ /dev/null @@ -1,2597 +0,0 @@ -/* -foreign_network_manager is used to forward packets of other networks. currently -only forward packets of peers that directly connected to this node. - -in the future, with the help wo peer center we can forward packets of peers that -connected to any node in the local network. -*/ -use std::{ - sync::{ - Arc, Weak, - atomic::{AtomicBool, Ordering}, - }, - time::SystemTime, -}; - -use dashmap::{DashMap, DashSet}; -use guarden::{Guard, defer}; -use tokio::{ - sync::{ - Mutex, - mpsc::{self, UnboundedReceiver, UnboundedSender}, - }, - task::JoinSet, -}; - -use crate::{ - common::{ - PeerId, - config::{ConfigLoader, TomlConfigLoader}, - error::Error, - global_ctx::{ArcGlobalCtx, GlobalCtx, GlobalCtxEvent, NetworkIdentity, TrustedKeySource}, - join_joinset_background, shrink_dashmap, - stats_manager::{LabelSet, LabelType, MetricName, StatsManager}, - token_bucket::TokenBucket, - }, - peer_center::instance::{PeerCenterInstance, PeerMapWithPeerRpcManager}, - peers::route_trait::{Route, RouteInterface}, - proto::{ - api::instance::{ - ForeignNetworkEntryPb, ListForeignNetworkResponse, PeerInfo, TrustedKeyInfoPb, - TrustedKeySourcePb, - }, - common::LimiterConfig, - peer_rpc::{DirectConnectorRpcServer, PeerIdentityType}, - }, - tunnel::packet_def::{PacketType, ZCPacket}, - use_global_var, -}; - -use super::{ - PUBLIC_SERVER_HOSTNAME_PREFIX, PacketRecvChan, PacketRecvChanReceiver, create_packet_recv_chan, - peer_conn::PeerConn, - peer_map::PeerMap, - peer_ospf_route::PeerRoute, - peer_rpc::{PeerRpcManager, PeerRpcManagerTransport}, - peer_rpc_service::DirectConnectorManagerRpcServer, - peer_session::PeerSessionStore, - recv_packet_from_chan, - relay_peer_map::RelayPeerMap, - route_trait::NextHopPolicy, - traffic_metrics::{ - InstanceLabelKind, LogicalTrafficMetrics, TrafficKind, TrafficMetricRecorder, - is_relay_data_packet_type, route_peer_info_instance_id, traffic_kind, - }, -}; - -#[async_trait::async_trait] -#[auto_impl::auto_impl(&, Box, Arc)] -pub trait GlobalForeignNetworkAccessor: Send + Sync + 'static { - async fn list_global_foreign_peer(&self, network_identity: &NetworkIdentity) -> Vec; -} - -struct ForeignNetworkEntry { - my_peer_id: PeerId, - - // Node-global runtime flags, such as disable_relay_data, live on the parent - // context. The foreign context is scoped to the foreign network's OSPF view. - parent_global_ctx: ArcGlobalCtx, - global_ctx: ArcGlobalCtx, - network: NetworkIdentity, - peer_map: Arc, - relay_peer_map: Arc, - peer_session_store: Arc, - // Static per-network permission from the whitelist check. disable_relay_data - // is the node-wide runtime override layered on top of this value. - relay_data: bool, - pm_packet_sender: Mutex>, - - peer_rpc: Arc, - rpc_sender: UnboundedSender, - - packet_recv: Mutex>, - - bps_limiter: Option>, - - peer_center: Arc, - - stats_mgr: Arc, - traffic_metrics: Arc, - event_handler_started: AtomicBool, - - tasks: Mutex>, - - pub lock: Mutex<()>, -} - -impl ForeignNetworkEntry { - fn new( - network: NetworkIdentity, - // NOTICE: ospf route need my_peer_id be changed after restart. - my_peer_id: PeerId, - global_ctx: ArcGlobalCtx, - relay_data: bool, - peer_session_store: Arc, - pm_packet_sender: PacketRecvChan, - ) -> Self { - let stats_mgr = global_ctx.stats_manager().clone(); - let foreign_global_ctx = - Self::build_foreign_global_ctx(&network, global_ctx.clone(), relay_data); - let network_name = network.network_name.clone(); - - let (packet_sender, packet_recv) = create_packet_recv_chan(); - - let peer_map = Arc::new(PeerMap::new( - packet_sender, - foreign_global_ctx.clone(), - my_peer_id, - )); - let traffic_metrics = Arc::new(TrafficMetricRecorder::new( - my_peer_id, - Arc::new(LogicalTrafficMetrics::new( - stats_mgr.clone(), - network_name.clone(), - MetricName::TrafficBytesTx, - MetricName::TrafficPacketsTx, - MetricName::TrafficBytesTxByInstance, - MetricName::TrafficPacketsTxByInstance, - InstanceLabelKind::To, - )), - Arc::new(LogicalTrafficMetrics::new( - stats_mgr.clone(), - network_name.clone(), - MetricName::TrafficControlBytesTx, - MetricName::TrafficControlPacketsTx, - MetricName::TrafficControlBytesTxByInstance, - MetricName::TrafficControlPacketsTxByInstance, - InstanceLabelKind::To, - )), - Arc::new(LogicalTrafficMetrics::new( - stats_mgr.clone(), - network_name.clone(), - MetricName::TrafficBytesRx, - MetricName::TrafficPacketsRx, - MetricName::TrafficBytesRxByInstance, - MetricName::TrafficPacketsRxByInstance, - InstanceLabelKind::From, - )), - Arc::new(LogicalTrafficMetrics::new( - stats_mgr.clone(), - network_name.clone(), - MetricName::TrafficControlBytesRx, - MetricName::TrafficControlPacketsRx, - MetricName::TrafficControlBytesRxByInstance, - MetricName::TrafficControlPacketsRxByInstance, - InstanceLabelKind::From, - )), - { - let peer_map = Arc::downgrade(&peer_map); - move |peer_id| { - let peer_map = peer_map.clone(); - async move { - let peer_map = peer_map.upgrade()?; - peer_map - .get_route_peer_info(peer_id) - .await - .as_ref() - .and_then(route_peer_info_instance_id) - } - } - }, - )); - let relay_peer_map = RelayPeerMap::new( - peer_map.clone(), - None, - foreign_global_ctx.clone(), - my_peer_id, - peer_session_store.clone(), - ); - - let (peer_rpc, rpc_transport_sender) = Self::build_rpc_tspt(my_peer_id, peer_map.clone()); - - peer_rpc.rpc_server().registry().register( - DirectConnectorRpcServer::new(DirectConnectorManagerRpcServer::new( - foreign_global_ctx.clone(), - )), - &network.network_name, - ); - - let relay_bps_limit = global_ctx.config.get_flags().foreign_relay_bps_limit; - let bps_limiter = (relay_bps_limit != u64::MAX).then(|| { - let limiter_config = LimiterConfig { - burst_rate: None, - bps: Some(relay_bps_limit), - fill_duration_ms: None, - }; - global_ctx - .token_bucket_manager() - .get_or_create(&network.network_name, limiter_config.into()) - }); - - let peer_center = Arc::new(PeerCenterInstance::new(Arc::new( - PeerMapWithPeerRpcManager { - peer_map: peer_map.clone(), - rpc_mgr: peer_rpc.clone(), - }, - ))); - - Self { - my_peer_id, - - parent_global_ctx: global_ctx.clone(), - global_ctx: foreign_global_ctx, - network, - peer_map, - relay_peer_map, - peer_session_store, - relay_data, - pm_packet_sender: Mutex::new(Some(pm_packet_sender)), - - peer_rpc, - rpc_sender: rpc_transport_sender, - - packet_recv: Mutex::new(Some(packet_recv)), - - bps_limiter, - - stats_mgr, - traffic_metrics, - event_handler_started: AtomicBool::new(false), - - tasks: Mutex::new(JoinSet::new()), - - peer_center, - - lock: Mutex::new(()), - } - } - - fn desired_avoid_relay_data_feature_flag( - parent_global_ctx: &ArcGlobalCtx, - relay_data: bool, - ) -> bool { - !relay_data || parent_global_ctx.get_feature_flags().avoid_relay_data - } - - fn sync_parent_relay_data_feature_flag( - parent_global_ctx: &ArcGlobalCtx, - global_ctx: &ArcGlobalCtx, - relay_data: bool, - ) -> bool { - let avoid_relay_data = - Self::desired_avoid_relay_data_feature_flag(parent_global_ctx, relay_data); - if global_ctx.get_feature_flags().avoid_relay_data == avoid_relay_data { - return false; - } - - global_ctx.set_avoid_relay_data_preference(avoid_relay_data) - } - - fn build_foreign_global_ctx( - network: &NetworkIdentity, - global_ctx: ArcGlobalCtx, - relay_data: bool, - ) -> ArcGlobalCtx { - let config = TomlConfigLoader::default(); - config.set_network_identity(network.clone()); - config.set_hostname(Some(format!( - "{}{}", - PUBLIC_SERVER_HOSTNAME_PREFIX, - global_ctx.get_hostname() - ))); - config.set_secure_mode(global_ctx.config.get_secure_mode()); - - let mut flags = config.get_flags(); - flags.disable_relay_kcp = !global_ctx.get_flags().enable_relay_foreign_network_kcp; - flags.disable_relay_quic = !global_ctx.get_flags().enable_relay_foreign_network_quic; - // socket_mark is a host-wide socket option: propagate from parent so - // outbound sockets the foreign-network entry initiates inherit the same - // mark as the rest of the node. - flags.socket_mark = global_ctx.get_flags().socket_mark; - config.set_flags(flags); - - config.set_mapped_listeners(Some(global_ctx.config.get_mapped_listeners())); - - let foreign_global_ctx = Arc::new(GlobalCtx::new(config)); - foreign_global_ctx - .replace_stun_info_collector(Box::new(global_ctx.get_stun_info_collector().clone())); - - let mut feature_flag = global_ctx.get_feature_flags(); - feature_flag.is_public_server = true; - feature_flag.avoid_relay_data = - Self::desired_avoid_relay_data_feature_flag(&global_ctx, relay_data); - foreign_global_ctx.set_base_advertised_feature_flags(feature_flag); - - for u in global_ctx.get_running_listeners().into_iter() { - foreign_global_ctx.add_running_listener(u); - } - - foreign_global_ctx - } - - fn build_rpc_tspt( - my_peer_id: PeerId, - peer_map: Arc, - ) -> (Arc, UnboundedSender) { - struct RpcTransport { - my_peer_id: PeerId, - peer_map: Weak, - - packet_recv: Mutex>, - } - - #[async_trait::async_trait] - impl PeerRpcManagerTransport for RpcTransport { - fn my_peer_id(&self) -> PeerId { - self.my_peer_id - } - - async fn send(&self, msg: ZCPacket, dst_peer_id: PeerId) -> Result<(), Error> { - tracing::debug!( - "foreign network manager send rpc to peer: {:?}", - dst_peer_id - ); - let peer_map = self - .peer_map - .upgrade() - .ok_or(anyhow::anyhow!("peer map is gone"))?; - - // send to ourselves so we can handle it in forward logic. - peer_map.send_msg_directly(msg, self.my_peer_id).await - } - - async fn recv(&self) -> Result { - if let Some(o) = self.packet_recv.lock().await.recv().await { - tracing::trace!("recv rpc packet in foreign network manager rpc transport"); - Ok(o) - } else { - Err(Error::Unknown) - } - } - } - - impl Drop for RpcTransport { - fn drop(&mut self) { - tracing::debug!( - "drop rpc transport for foreign network manager, my_peer_id: {:?}", - self.my_peer_id - ); - } - } - - let (rpc_transport_sender, peer_rpc_tspt_recv) = mpsc::unbounded_channel(); - let tspt = RpcTransport { - my_peer_id, - peer_map: Arc::downgrade(&peer_map), - packet_recv: Mutex::new(peer_rpc_tspt_recv), - }; - - let peer_rpc = Arc::new(PeerRpcManager::new(tspt)); - (peer_rpc, rpc_transport_sender) - } - - async fn prepare_route(&self, accessor: Box) { - struct Interface { - my_peer_id: PeerId, - peer_map: Weak, - network_identity: NetworkIdentity, - accessor: Box, - } - - #[async_trait::async_trait] - impl RouteInterface for Interface { - async fn list_peers(&self) -> Vec { - let Some(peer_map) = self.peer_map.upgrade() else { - return vec![]; - }; - - let mut global = self - .accessor - .list_global_foreign_peer(&self.network_identity) - .await; - let local = peer_map.list_peers_with_conn().await; - global.extend(local.iter().cloned()); - global - .into_iter() - .filter(|x| *x != self.my_peer_id) - .collect() - } - - fn my_peer_id(&self) -> PeerId { - self.my_peer_id - } - - fn need_periodic_requery_peers(&self) -> bool { - true - } - - async fn get_peer_identity_type(&self, peer_id: PeerId) -> Option { - let peer_map = self.peer_map.upgrade()?; - peer_map.get_peer_identity_type(peer_id) - } - - async fn get_peer_public_key(&self, peer_id: PeerId) -> Option> { - let peer_map = self.peer_map.upgrade()?; - peer_map.get_peer_public_key(peer_id) - } - - async fn close_peer(&self, peer_id: PeerId) { - if let Some(peer_map) = self.peer_map.upgrade() { - let _ = peer_map.close_peer(peer_id).await; - } - } - } - - let route = PeerRoute::new( - self.my_peer_id, - self.global_ctx.clone(), - self.peer_rpc.clone(), - ); - route - .open(Box::new(Interface { - my_peer_id: self.my_peer_id, - network_identity: self.network.clone(), - peer_map: Arc::downgrade(&self.peer_map), - accessor, - })) - .await - .unwrap(); - - route - .set_route_cost_fn(self.peer_center.get_cost_calculator()) - .await; - - self.peer_map.add_route(Arc::new(Box::new(route))).await; - } - - async fn start_packet_recv(&self) { - let mut recv = self.packet_recv.lock().await.take().unwrap(); - let my_node_id = self.my_peer_id; - let rpc_sender = self.rpc_sender.clone(); - let peer_map = self.peer_map.clone(); - let relay_peer_map = self.relay_peer_map.clone(); - let traffic_metrics = self.traffic_metrics.clone(); - let parent_global_ctx = self.parent_global_ctx.clone(); - let relay_data = self.relay_data; - let pm_sender = self.pm_packet_sender.lock().await.take().unwrap(); - let network_name = self.network.network_name.clone(); - let bps_limiter = self.bps_limiter.clone(); - - let label_set = - LabelSet::new().with_label_type(LabelType::NetworkName(network_name.clone())); - let forward_data_bytes = self - .stats_mgr - .get_counter(MetricName::TrafficBytesForwarded, label_set.clone()); - let forward_data_packets = self - .stats_mgr - .get_counter(MetricName::TrafficPacketsForwarded, label_set.clone()); - let forward_control_bytes = self - .stats_mgr - .get_counter(MetricName::TrafficControlBytesForwarded, label_set.clone()); - let forward_control_packets = self.stats_mgr.get_counter( - MetricName::TrafficControlPacketsForwarded, - label_set.clone(), - ); - let rx_bytes = self - .stats_mgr - .get_counter(MetricName::TrafficBytesSelfRx, label_set.clone()); - let rx_packets = self - .stats_mgr - .get_counter(MetricName::TrafficPacketsRx, label_set.clone()); - - self.tasks.lock().await.spawn(async move { - while let Ok(mut zc_packet) = recv_packet_from_chan(&mut recv).await { - let buf_len = zc_packet.buf_len(); - let Some(hdr) = zc_packet.peer_manager_header() else { - tracing::warn!("invalid packet, skip"); - continue; - }; - tracing::trace!(?hdr, "recv packet in foreign network manager"); - let from_peer_id = hdr.from_peer_id.get(); - let packet_type = hdr.packet_type; - let len = hdr.len.get(); - let to_peer_id = hdr.to_peer_id.get(); - let is_local_delivery = to_peer_id == my_node_id; - let is_locally_originated = from_peer_id == my_node_id; - if is_local_delivery && !is_locally_originated { - traffic_metrics - .record_rx(from_peer_id, packet_type, buf_len as u64) - .await; - } - if is_local_delivery { - if packet_type == PacketType::RelayHandshake as u8 - || packet_type == PacketType::RelayHandshakeAck as u8 - { - let _ = relay_peer_map.handle_handshake_packet(zc_packet).await; - continue; - } - - if relay_peer_map.is_secure_mode_enabled() && hdr.is_encrypted() { - match relay_peer_map.decrypt_if_needed(&mut zc_packet).await { - Ok(true) => {} - Ok(false) => { - tracing::error!("secure session not found"); - continue; - } - Err(e) => { - tracing::error!(?e, "secure decrypt failed"); - continue; - } - } - } - - if packet_type == PacketType::TaRpc as u8 - || packet_type == PacketType::RpcReq as u8 - || packet_type == PacketType::RpcResp as u8 - { - rx_bytes.add(buf_len as u64); - rx_packets.inc(); - rpc_sender.send(zc_packet).unwrap(); - continue; - } - tracing::trace!( - ?packet_type, - ?len, - ?from_peer_id, - ?to_peer_id, - "ignore packet in foreign network" - ); - } else { - if is_relay_data_packet_type(packet_type) { - let disable_relay_data = parent_global_ctx.flags_arc().disable_relay_data; - if !relay_data || disable_relay_data { - tracing::debug!( - ?from_peer_id, - ?to_peer_id, - packet_type, - disable_relay_data, - "drop foreign network relay data" - ); - continue; - } - if let Some(bps_limiter) = bps_limiter.as_ref() - && !bps_limiter.try_consume(len.into()) - { - continue; - } - } - - match traffic_kind(packet_type) { - TrafficKind::Data => { - forward_data_bytes.add(buf_len as u64); - forward_data_packets.inc(); - } - TrafficKind::Control => { - forward_control_bytes.add(buf_len as u64); - forward_control_packets.inc(); - } - } - - let gateway_peer_id = peer_map - .get_gateway_peer_id(to_peer_id, NextHopPolicy::LeastHop) - .await; - - match gateway_peer_id { - Some(peer_id) if peer_map.has_peer(peer_id) => { - if peer_id != to_peer_id && hdr.from_peer_id.get() == my_node_id { - if let Err(e) = relay_peer_map - .send_msg(zc_packet, to_peer_id, NextHopPolicy::LeastHop) - .await - { - tracing::error!( - ?e, - "send packet to foreign peer inside relay peer map failed" - ); - } else if is_locally_originated { - traffic_metrics - .record_tx(to_peer_id, packet_type, buf_len as u64) - .await; - } - } else if let Err(e) = - peer_map.send_msg_directly(zc_packet, peer_id).await - { - tracing::error!( - ?e, - "send packet to foreign peer inside peer map failed" - ); - } else if is_locally_originated { - traffic_metrics - .record_tx(to_peer_id, packet_type, buf_len as u64) - .await; - } - } - _ => { - let mut foreign_packet = ZCPacket::new_for_foreign_network( - &network_name, - to_peer_id, - &zc_packet, - ); - let via_peer = gateway_peer_id.unwrap_or(to_peer_id); - foreign_packet.fill_peer_manager_hdr( - my_node_id, - via_peer, - PacketType::ForeignNetworkPacket as u8, - ); - if let Err(e) = pm_sender.send(foreign_packet).await { - tracing::error!("send packet to peer with pm failed: {:?}", e); - } else if is_locally_originated { - traffic_metrics - .record_tx(to_peer_id, packet_type, buf_len as u64) - .await; - } - } - }; - } - } - }); - } - - async fn run_relay_session_gc_routine(&self) { - let relay_peer_map = self.relay_peer_map.clone(); - self.tasks.lock().await.spawn(async move { - loop { - relay_peer_map.evict_idle_sessions(std::time::Duration::from_secs(60)); - tokio::time::sleep(std::time::Duration::from_secs(30)).await; - } - }); - } - - async fn run_parent_feature_flag_sync_routine(&self) { - let parent_global_ctx = self.parent_global_ctx.clone(); - let global_ctx = self.global_ctx.clone(); - let relay_data = self.relay_data; - self.tasks.lock().await.spawn(async move { - let mut parent_events = parent_global_ctx.subscribe(); - loop { - ForeignNetworkEntry::sync_parent_relay_data_feature_flag( - &parent_global_ctx, - &global_ctx, - relay_data, - ); - - if parent_events.recv().await.is_err() { - parent_events = parent_global_ctx.subscribe(); - } - } - }); - } - - async fn prepare(&self, accessor: Box) { - self.prepare_route(accessor).await; - self.start_packet_recv().await; - self.run_relay_session_gc_routine().await; - self.run_parent_feature_flag_sync_routine().await; - self.peer_rpc.run(); - self.peer_center.init().await; - } -} - -impl Drop for ForeignNetworkEntry { - fn drop(&mut self) { - self.peer_rpc - .rpc_server() - .registry() - .unregister_by_domain(&self.network.network_name); - self.global_ctx - .remove_trusted_keys(&self.network.network_name); - - tracing::debug!(self.my_peer_id, ?self.network, "drop foreign network entry"); - } -} - -struct ForeignNetworkManagerData { - network_peer_maps: DashMap>, - peer_network_map: DashMap>, - network_peer_last_update: DashMap, - accessor: Arc>, - lock: std::sync::Mutex<()>, - #[cfg(test)] - fail_next_add_peer_conn_after_entry_insert: AtomicBool, -} - -impl ForeignNetworkManagerData { - fn get_peer_network(&self, peer_id: PeerId) -> Option> { - self.peer_network_map.get(&peer_id).map(|v| v.clone()) - } - - fn get_network_entry(&self, network_name: &str) -> Option> { - self.network_peer_maps.get(network_name).map(|v| v.clone()) - } - - fn remove_peer(&self, peer_id: PeerId, network_name: &String) { - let _l = self.lock.lock().unwrap(); - self.peer_network_map.remove_if(&peer_id, |_, v| { - let _ = v.remove(network_name); - v.is_empty() - }); - if self - .network_peer_maps - .remove_if(network_name, |_, v| v.peer_map.is_empty()) - .is_some() - { - self.network_peer_last_update.remove(network_name); - } - shrink_dashmap(&self.peer_network_map, None); - shrink_dashmap(&self.network_peer_maps, None); - shrink_dashmap(&self.network_peer_last_update, None); - } - - async fn clear_no_conn_peer(&self, network_name: &String) { - let Some(peer_map) = self - .network_peer_maps - .get(network_name) - .map(|v| v.peer_map.clone()) - else { - return; - }; - peer_map.clean_peer_without_conn().await; - } - - fn remove_network(&self, network_name: &String) { - let _l = self.lock.lock().unwrap(); - if let Some(old) = self.network_peer_maps.remove(network_name) { - old.1.traffic_metrics.clear_peer_cache(); - let to_remove_peers = old.1.peer_map.list_peers(); - for p in to_remove_peers { - self.peer_network_map.remove_if(&p, |_, v| { - v.remove(network_name); - v.is_empty() - }); - } - } - self.network_peer_last_update.remove(network_name); - shrink_dashmap(&self.peer_network_map, None); - shrink_dashmap(&self.network_peer_maps, None); - shrink_dashmap(&self.network_peer_last_update, None); - } - - fn remove_network_if_current( - &self, - network_name: &String, - expected_entry: &Weak, - ) { - let _l = self.lock.lock().unwrap(); - let Some(expected_entry) = expected_entry.upgrade() else { - return; - }; - let old = self - .network_peer_maps - .remove_if(network_name, |_, entry| Arc::ptr_eq(entry, &expected_entry)); - let Some((_, old)) = old else { - return; - }; - - old.traffic_metrics.clear_peer_cache(); - let to_remove_peers = old.peer_map.list_peers(); - for p in to_remove_peers { - self.peer_network_map.remove_if(&p, |_, v| { - v.remove(network_name); - v.is_empty() - }); - } - self.network_peer_last_update.remove(network_name); - shrink_dashmap(&self.peer_network_map, None); - shrink_dashmap(&self.network_peer_maps, None); - shrink_dashmap(&self.network_peer_last_update, None); - } - - #[allow(clippy::too_many_arguments)] - async fn get_or_insert_entry( - &self, - network_identity: &NetworkIdentity, - my_peer_id: PeerId, - dst_peer_id: PeerId, - relay_data: bool, - global_ctx: &ArcGlobalCtx, - peer_session_store: Arc, - pm_packet_sender: &PacketRecvChan, - ) -> (Arc, bool) { - let mut new_added = false; - - let l = self.lock.lock().unwrap(); - let entry = self - .network_peer_maps - .entry(network_identity.network_name.clone()) - .or_insert_with(|| { - new_added = true; - Arc::new(ForeignNetworkEntry::new( - network_identity.clone(), - my_peer_id, - global_ctx.clone(), - relay_data, - peer_session_store, - pm_packet_sender.clone(), - )) - }) - .clone(); - - self.peer_network_map - .entry(dst_peer_id) - .or_default() - .insert(network_identity.network_name.clone()); - - self.network_peer_last_update - .insert(network_identity.network_name.clone(), SystemTime::now()); - - drop(l); - - if new_added { - entry.prepare(Box::new(self.accessor.clone())).await; - } - - (entry, new_added) - } -} - -pub const FOREIGN_NETWORK_SERVICE_ID: u32 = 1; - -pub struct ForeignNetworkManager { - my_peer_id: PeerId, - global_ctx: ArcGlobalCtx, - peer_session_store: Arc, - packet_sender_to_mgr: PacketRecvChan, - - data: Arc, - - tasks: Arc>>, -} - -impl ForeignNetworkManager { - fn network_secret_digest_is_empty(network: &NetworkIdentity) -> bool { - network - .network_secret_digest - .as_ref() - .is_none_or(|d| d.iter().all(|b| *b == 0)) - } - - fn should_reject_credential_trust_path(identity_type: PeerIdentityType) -> bool { - matches!(identity_type, PeerIdentityType::Admin) - } - - fn credential_pubkey_is_trusted( - global_ctx: &ArcGlobalCtx, - network_name: &str, - remote_static_pubkey: &[u8], - ) -> bool { - remote_static_pubkey.len() == 32 - && global_ctx.is_pubkey_trusted_with_source( - remote_static_pubkey, - network_name, - TrustedKeySource::OspfCredential, - ) - } - - fn is_credential_pubkey_trusted( - entry: &ForeignNetworkEntry, - remote_static_pubkey: &[u8], - ) -> bool { - Self::credential_pubkey_is_trusted( - &entry.global_ctx, - &entry.network.network_name, - remote_static_pubkey, - ) - } - - pub(crate) fn is_existing_credential_pubkey_trusted( - &self, - network_name: &str, - remote_static_pubkey: &[u8], - ) -> bool { - self.data - .get_network_entry(network_name) - .is_some_and(|entry| { - Self::credential_pubkey_is_trusted( - &entry.global_ctx, - &entry.network.network_name, - remote_static_pubkey, - ) - }) - } - - fn build_trusted_key_items(entry: &ForeignNetworkEntry) -> Vec { - entry - .global_ctx - .list_trusted_keys(&entry.network.network_name) - .into_iter() - .map(|(pubkey, metadata)| TrustedKeyInfoPb { - pubkey, - source: match metadata.source { - TrustedKeySource::OspfNode => TrustedKeySourcePb::OspfNode.into(), - TrustedKeySource::OspfCredential => TrustedKeySourcePb::OspfCredential.into(), - }, - expiry_unix: metadata.expiry_unix, - }) - .collect() - } - - pub fn new( - my_peer_id: PeerId, - global_ctx: ArcGlobalCtx, - peer_session_store: Arc, - packet_sender_to_mgr: PacketRecvChan, - accessor: Box, - ) -> Self { - let data = Arc::new(ForeignNetworkManagerData { - network_peer_maps: DashMap::new(), - peer_network_map: DashMap::new(), - network_peer_last_update: DashMap::new(), - accessor: Arc::new(accessor), - lock: std::sync::Mutex::new(()), - #[cfg(test)] - fail_next_add_peer_conn_after_entry_insert: AtomicBool::new(false), - }); - - let tasks = Arc::new(std::sync::Mutex::new(JoinSet::new())); - join_joinset_background(tasks.clone(), "ForeignNetworkManager".to_string()); - - Self { - my_peer_id, - global_ctx, - peer_session_store, - packet_sender_to_mgr, - - data, - - tasks, - } - } - - #[cfg(test)] - fn fail_next_add_peer_conn_after_entry_insert(&self) { - self.data - .fail_next_add_peer_conn_after_entry_insert - .store(true, Ordering::Release); - } - - pub fn get_network_peer_id(&self, network_name: &str) -> Option { - self.data - .network_peer_maps - .get(network_name) - .map(|v| v.my_peer_id) - } - - pub async fn add_peer_conn(&self, peer_conn: PeerConn) -> Result<(), Error> { - let conn_info = peer_conn.get_conn_info(); - let peer_network = peer_conn.get_network_identity(); - tracing::info!(peer_conn = ?conn_info, network = ?peer_network, "add new peer conn in foreign network manager"); - - let relay_peer_rpc = self.global_ctx.get_flags().relay_all_peer_rpc; - let ret = self - .global_ctx - .check_network_in_whitelist(&peer_network.network_name) - .map_err(Into::into); - if ret.is_err() && !relay_peer_rpc { - return ret; - } - - let peer_digest_empty = Self::network_secret_digest_is_empty(&peer_network); - if peer_digest_empty - && self - .data - .get_network_entry(&peer_network.network_name) - .is_none() - { - return Err(anyhow::anyhow!( - "foreign network {} is not established by a secret-verified peer yet", - peer_network.network_name - ) - .into()); - } - - let (entry, new_added) = self - .data - .get_or_insert_entry( - &peer_network, - peer_conn.get_my_peer_id(), - peer_conn.get_peer_id(), - ret.is_ok(), - &self.global_ctx, - self.peer_session_store.clone(), - &self.packet_sender_to_mgr, - ) - .await; - - defer!(rollback_new_entry => sync [ - data = self.data.clone(), - network_name = entry.network.network_name.clone(), - peer_id = peer_conn.get_peer_id(), - should_rollback = new_added - ] { - if should_rollback { - tracing::warn!( - %network_name, - "rollback newly added foreign network entry after add_peer_conn returned error" - ); - data.remove_peer(peer_id, &network_name); - } - }); - - #[cfg(test)] - if self - .data - .fail_next_add_peer_conn_after_entry_insert - .swap(false, Ordering::AcqRel) - { - return Err(anyhow::anyhow!( - "injected add_peer_conn failure after foreign network entry insert" - ) - .into()); - } - - self.ensure_event_handler_started(&entry); - - let same_identity = entry.network == peer_network; - let peer_identity_type = peer_conn.get_peer_identity_type(); - let credential_peer_trusted = peer_digest_empty - && Self::is_credential_pubkey_trusted(&entry, &conn_info.noise_remote_static_pubkey); - let credential_identity_mismatch = credential_peer_trusted - && Self::should_reject_credential_trust_path(peer_identity_type); - - let _g = entry.lock.lock().await; - - if (!(same_identity || credential_peer_trusted)) - || credential_identity_mismatch - || entry.my_peer_id != peer_conn.get_my_peer_id() - { - let err = if entry.my_peer_id != peer_conn.get_my_peer_id() { - anyhow::anyhow!( - "my peer id not match. exp: {:?} real: {:?}, need retry connect", - entry.my_peer_id, - peer_conn.get_my_peer_id() - ) - } else if credential_identity_mismatch { - anyhow::anyhow!( - "credential-trusted foreign peer has invalid identity type: {:?}", - peer_identity_type - ) - } else { - anyhow::anyhow!( - "foreign peer identity not trusted. exp: {:?} real: {:?}, remote_pubkey_len: {}, credential_trusted: {}", - entry.network, - peer_network, - conn_info.noise_remote_static_pubkey.len(), - credential_peer_trusted, - ) - }; - tracing::error!(?err, "foreign network entry not match, disconnect peer"); - return Err(err.into()); - } - - if !new_added && let Some(peer) = entry.peer_map.get_peer_by_id(peer_conn.get_peer_id()) { - let direct_conns_len = peer.get_directly_connections().len(); - let max_count = use_global_var!(MAX_DIRECT_CONNS_PER_PEER_IN_FOREIGN_NETWORK); - if direct_conns_len >= max_count as usize { - return Err(anyhow::anyhow!( - "too many direct conns, cur: {}, max: {}", - direct_conns_len, - max_count - ) - .into()); - } - } - - entry.peer_map.add_new_peer_conn(peer_conn).await?; - let _ = rollback_new_entry.defuse(); - Ok(()) - } - - fn ensure_event_handler_started(&self, entry: &Arc) { - if entry.event_handler_started.swap(true, Ordering::AcqRel) { - return; - } - - let data = self.data.clone(); - let network_name = entry.network.network_name.clone(); - let entry_for_cleanup = Arc::downgrade(entry); - let traffic_metrics = Arc::downgrade(&entry.traffic_metrics); - let mut s = entry.global_ctx.subscribe(); - self.tasks.lock().unwrap().spawn(async move { - while let Ok(e) = s.recv().await { - match &e { - GlobalCtxEvent::PeerRemoved(peer_id) => { - tracing::info!(?e, "remove peer from foreign network manager"); - if let Some(traffic_metrics) = traffic_metrics.upgrade() { - traffic_metrics.remove_peer(*peer_id); - } - data.network_peer_last_update - .insert(network_name.clone(), SystemTime::now()); - data.remove_peer(*peer_id, &network_name); - } - GlobalCtxEvent::PeerConnRemoved(..) => { - tracing::info!(?e, "clear no conn peer from foreign network manager"); - data.clear_no_conn_peer(&network_name).await; - } - GlobalCtxEvent::PeerAdded(_) => { - tracing::info!(?e, "add peer to foreign network manager"); - data.network_peer_last_update - .insert(network_name.clone(), SystemTime::now()); - } - _ => continue, - } - } - // if lagged or recv done just remove the network - tracing::error!("global event handler at foreign network manager exit"); - if let Some(traffic_metrics) = traffic_metrics.upgrade() { - traffic_metrics.clear_peer_cache(); - } - data.remove_network_if_current(&network_name, &entry_for_cleanup); - }); - } - - pub async fn list_foreign_networks(&self) -> ListForeignNetworkResponse { - self.list_foreign_networks_with_options(false).await - } - - pub async fn list_foreign_networks_with_options( - &self, - include_trusted_keys: bool, - ) -> ListForeignNetworkResponse { - let mut ret = ListForeignNetworkResponse::default(); - let networks = self - .data - .network_peer_maps - .iter() - .map(|v| v.key().clone()) - .collect::>(); - - for network_name in networks { - let Some(item) = self - .data - .network_peer_maps - .get(&network_name) - .map(|v| v.clone()) - else { - continue; - }; - - let mut entry = ForeignNetworkEntryPb { - network_secret_digest: item - .network - .network_secret_digest - .unwrap_or_default() - .to_vec(), - my_peer_id_for_this_network: item.my_peer_id, - peers: Default::default(), - trusted_keys: if include_trusted_keys { - Self::build_trusted_key_items(&item) - } else { - Default::default() - }, - }; - for peer in item.peer_map.list_peers() { - let peer_info = PeerInfo { - peer_id: peer, - conns: item.peer_map.list_peer_conns(peer).await.unwrap_or(vec![]), - ..Default::default() - }; - entry.peers.push(peer_info); - } - - ret.foreign_networks.insert(network_name, entry); - } - ret - } - - pub fn get_foreign_network_last_update(&self, network_name: &str) -> Option { - self.data - .network_peer_last_update - .get(network_name) - .map(|v| *v) - } - - pub async fn forward_foreign_network_packet( - &self, - network_name: &str, - dst_peer_id: PeerId, - msg: ZCPacket, - ) -> Result<(), Error> { - if let Some(entry) = self.data.get_network_entry(network_name) { - let packet_type = msg - .peer_manager_header() - .map(|hdr| hdr.packet_type) - .unwrap_or(0); - let msg_len = msg.buf_len() as u64; - let send_result = entry - .peer_map - .send_msg(msg, dst_peer_id, NextHopPolicy::LeastHop) - .await; - if send_result.is_ok() { - entry - .traffic_metrics - .record_tx(dst_peer_id, packet_type, msg_len) - .await; - } - send_result - } else { - Err(Error::RouteError(Some("network not found".to_string()))) - } - } - - pub async fn close_peer_conn( - &self, - peer_id: PeerId, - conn_id: &super::peer_conn::PeerConnId, - ) -> Result<(), Error> { - let network_names = self.data.get_peer_network(peer_id).unwrap_or_default(); - for network_name in network_names { - if let Some(entry) = self.data.get_network_entry(&network_name) { - let ret = entry.peer_map.close_peer_conn(peer_id, conn_id).await; - if ret.is_ok() || !matches!(ret.as_ref().unwrap_err(), Error::NotFound) { - return ret; - } - } - } - Err(Error::NotFound) - } -} - -impl Drop for ForeignNetworkManager { - fn drop(&mut self) { - self.data.peer_network_map.clear(); - self.data.network_peer_maps.clear(); - } -} - -#[cfg(test)] -pub mod tests { - use crate::{ - common::global_ctx::tests::get_mock_global_ctx_with_network, - common::stats_manager::{LabelSet, LabelType, MetricName}, - connector::udp_hole_punch::tests::{ - create_mock_peer_manager_with_mock_stun, replace_stun_info_collector, - }, - peers::{ - peer_conn::tests::set_secure_mode_cfg, - peer_manager::{PeerManager, RouteAlgoType}, - tests::{connect_peer_manager, wait_route_appear}, - }, - proto::common::NatType, - set_global_var, - tunnel::{ - common::tests::wait_for_condition, - packet_def::{PacketType, ZCPacket}, - }, - }; - use std::{collections::HashMap, time::Duration}; - - use super::*; - - fn metric_value(peer_mgr: &PeerManager, metric: MetricName, labels: LabelSet) -> u64 { - peer_mgr - .get_global_ctx() - .stats_manager() - .get_metric(metric, &labels) - .map(|metric| metric.value) - .unwrap_or(0) - } - - async fn create_mock_peer_manager_for_foreign_network_ext( - network: &str, - secret: &str, - ) -> Arc { - let (s, _r) = create_packet_recv_chan(); - let peer_mgr = Arc::new(PeerManager::new( - RouteAlgoType::Ospf, - get_mock_global_ctx_with_network(Some(NetworkIdentity::new( - network.to_string(), - secret.to_string(), - ))), - s, - )); - replace_stun_info_collector(peer_mgr.clone(), NatType::Unknown); - peer_mgr.run().await.unwrap(); - peer_mgr - } - - async fn create_mock_credential_peer_manager_for_foreign_network( - network: &str, - ) -> Arc { - let (s, _r) = create_packet_recv_chan(); - let global_ctx = get_mock_global_ctx_with_network(Some(NetworkIdentity::new_credential( - network.to_string(), - ))); - set_secure_mode_cfg(&global_ctx, true); - let peer_mgr = Arc::new(PeerManager::new(RouteAlgoType::Ospf, global_ctx, s)); - replace_stun_info_collector(peer_mgr.clone(), NatType::Unknown); - peer_mgr.run().await.unwrap(); - peer_mgr - } - - pub async fn create_mock_peer_manager_for_foreign_network(network: &str) -> Arc { - create_mock_peer_manager_for_foreign_network_ext(network, network).await - } - - pub async fn create_mock_peer_manager_for_secure_foreign_network( - network: &str, - ) -> Arc { - let (s, _r) = create_packet_recv_chan(); - let global_ctx = get_mock_global_ctx_with_network(Some(NetworkIdentity::new( - network.to_string(), - network.to_string(), - ))); - set_secure_mode_cfg(&global_ctx, true); - let peer_mgr = Arc::new(PeerManager::new(RouteAlgoType::Ospf, global_ctx, s)); - replace_stun_info_collector(peer_mgr.clone(), NatType::Unknown); - peer_mgr.run().await.unwrap(); - peer_mgr - } - - #[tokio::test] - async fn foreign_network_basic() { - let pm_center = create_mock_peer_manager_with_mock_stun(NatType::Unknown).await; - tracing::debug!("pm_center: {:?}", pm_center.my_peer_id()); - - let pma_net1 = create_mock_peer_manager_for_foreign_network("net1").await; - let pmb_net1 = create_mock_peer_manager_for_foreign_network("net1").await; - tracing::debug!( - "pma_net1: {:?}, pmb_net1: {:?}", - pma_net1.my_peer_id(), - pmb_net1.my_peer_id() - ); - connect_peer_manager(pma_net1.clone(), pm_center.clone()).await; - connect_peer_manager(pmb_net1.clone(), pm_center.clone()).await; - wait_route_appear(pma_net1.clone(), pmb_net1.clone()) - .await - .unwrap(); - assert_eq!(2, pma_net1.list_routes().await.len()); - assert_eq!(2, pmb_net1.list_routes().await.len()); - - println!("{:?}", pmb_net1.list_routes().await); - - let rpc_resp = pm_center - .get_foreign_network_manager() - .list_foreign_networks() - .await; - assert_eq!(1, rpc_resp.foreign_networks.len()); - assert_eq!(2, rpc_resp.foreign_networks["net1"].peers.len()); - } - - #[tokio::test] - async fn foreign_network_forwarding_records_traffic_metrics() { - let pm_center = create_mock_peer_manager_with_mock_stun(NatType::Unknown).await; - let pma_net1 = create_mock_peer_manager_for_foreign_network("net1").await; - let pmb_net1 = create_mock_peer_manager_for_foreign_network("net1").await; - - connect_peer_manager(pma_net1.clone(), pm_center.clone()).await; - connect_peer_manager(pmb_net1.clone(), pm_center.clone()).await; - wait_route_appear(pma_net1.clone(), pmb_net1.clone()) - .await - .unwrap(); - - let center_peer_id = pm_center - .get_foreign_network_manager() - .get_network_peer_id("net1") - .unwrap(); - - let mut rx_pkt = ZCPacket::new_with_payload(b"foreign-rx"); - rx_pkt.fill_peer_manager_hdr( - pma_net1.my_peer_id(), - center_peer_id, - PacketType::Data as u8, - ); - pma_net1 - .get_foreign_network_client() - .send_msg(rx_pkt, center_peer_id) - .await - .unwrap(); - - let mut tx_pkt = ZCPacket::new_with_payload(b"foreign-tx"); - tx_pkt.fill_peer_manager_hdr( - center_peer_id, - pmb_net1.my_peer_id(), - PacketType::Data as u8, - ); - pm_center - .get_foreign_network_manager() - .forward_foreign_network_packet("net1", pmb_net1.my_peer_id(), tx_pkt) - .await - .unwrap(); - - let network_labels = - LabelSet::new().with_label_type(LabelType::NetworkName("net1".to_string())); - let tx_instance_labels = network_labels - .clone() - .with_label_type(LabelType::ToInstanceId( - pmb_net1.get_global_ctx().get_id().to_string(), - )); - let rx_instance_labels = network_labels - .clone() - .with_label_type(LabelType::FromInstanceId( - pma_net1.get_global_ctx().get_id().to_string(), - )); - - wait_for_condition( - || { - let pm_center = pm_center.clone(); - let network_labels = network_labels.clone(); - let tx_instance_labels = tx_instance_labels.clone(); - let rx_instance_labels = rx_instance_labels.clone(); - async move { - metric_value( - &pm_center, - MetricName::TrafficBytesTx, - network_labels.clone(), - ) > 0 - && metric_value( - &pm_center, - MetricName::TrafficBytesRx, - network_labels.clone(), - ) > 0 - && metric_value( - &pm_center, - MetricName::TrafficBytesTxByInstance, - tx_instance_labels.clone(), - ) > 0 - && metric_value( - &pm_center, - MetricName::TrafficBytesRxByInstance, - rx_instance_labels.clone(), - ) > 0 - } - }, - Duration::from_secs(5), - ) - .await; - } - - #[tokio::test] - async fn foreign_network_transit_forwarding_only_records_forwarded_metrics() { - let pm_center = create_mock_peer_manager_with_mock_stun(NatType::Unknown).await; - let pma_net1 = create_mock_peer_manager_for_foreign_network("net1").await; - let pmb_net1 = create_mock_peer_manager_for_foreign_network("net1").await; - - connect_peer_manager(pma_net1.clone(), pm_center.clone()).await; - connect_peer_manager(pmb_net1.clone(), pm_center.clone()).await; - wait_route_appear(pma_net1.clone(), pmb_net1.clone()) - .await - .unwrap(); - - let center_peer_id = pm_center - .get_foreign_network_manager() - .get_network_peer_id("net1") - .unwrap(); - let network_labels = - LabelSet::new().with_label_type(LabelType::NetworkName("net1".to_string())); - let forwarded_bytes_before = metric_value( - &pm_center, - MetricName::TrafficBytesForwarded, - network_labels.clone(), - ); - let forwarded_packets_before = metric_value( - &pm_center, - MetricName::TrafficPacketsForwarded, - network_labels.clone(), - ); - let rx_bytes_before = metric_value( - &pm_center, - MetricName::TrafficBytesRx, - network_labels.clone(), - ); - let rx_packets_before = metric_value( - &pm_center, - MetricName::TrafficPacketsRx, - network_labels.clone(), - ); - let tx_bytes_before = metric_value( - &pm_center, - MetricName::TrafficBytesTx, - network_labels.clone(), - ); - let tx_packets_before = metric_value( - &pm_center, - MetricName::TrafficPacketsTx, - network_labels.clone(), - ); - - let mut transit_pkt = ZCPacket::new_with_payload(b"foreign-transit"); - transit_pkt.fill_peer_manager_hdr( - pma_net1.my_peer_id(), - pmb_net1.my_peer_id(), - PacketType::Data as u8, - ); - let transit_pkt_len = transit_pkt.buf_len() as u64; - pma_net1 - .get_foreign_network_client() - .send_msg(transit_pkt, center_peer_id) - .await - .unwrap(); - wait_for_condition( - || { - let pm_center = pm_center.clone(); - let network_labels = network_labels.clone(); - async move { - metric_value( - &pm_center, - MetricName::TrafficBytesForwarded, - network_labels.clone(), - ) >= forwarded_bytes_before + transit_pkt_len - && metric_value( - &pm_center, - MetricName::TrafficPacketsForwarded, - network_labels.clone(), - ) > forwarded_packets_before - } - }, - Duration::from_secs(5), - ) - .await; - - assert_eq!( - metric_value( - &pm_center, - MetricName::TrafficBytesRx, - network_labels.clone() - ), - rx_bytes_before - ); - assert_eq!( - metric_value( - &pm_center, - MetricName::TrafficPacketsRx, - network_labels.clone() - ), - rx_packets_before - ); - assert_eq!( - metric_value( - &pm_center, - MetricName::TrafficBytesTx, - network_labels.clone() - ), - tx_bytes_before - ); - assert_eq!( - metric_value(&pm_center, MetricName::TrafficPacketsTx, network_labels), - tx_packets_before - ); - } - - #[tokio::test] - async fn disable_relay_data_blocks_foreign_network_transit_data() { - let pm_center = create_mock_peer_manager_with_mock_stun(NatType::Unknown).await; - let pma_net1 = create_mock_peer_manager_for_foreign_network("net1").await; - let pmb_net1 = create_mock_peer_manager_for_foreign_network("net1").await; - - connect_peer_manager(pma_net1.clone(), pm_center.clone()).await; - connect_peer_manager(pmb_net1.clone(), pm_center.clone()).await; - wait_route_appear(pma_net1.clone(), pmb_net1.clone()) - .await - .unwrap(); - - let mut flags = pm_center.get_global_ctx().get_flags(); - flags.disable_relay_data = true; - pm_center.get_global_ctx().set_flags(flags); - pm_center - .get_global_ctx() - .issue_event(GlobalCtxEvent::ConfigPatched(Default::default())); - - let center_peer_id = pm_center - .get_foreign_network_manager() - .get_network_peer_id("net1") - .unwrap(); - wait_for_condition( - || { - let pma_net1 = pma_net1.clone(); - async move { - pma_net1.list_routes().await.iter().any(|route| { - route.peer_id == center_peer_id - && route - .feature_flag - .as_ref() - .map(|flag| flag.avoid_relay_data) - .unwrap_or(false) - }) - } - }, - Duration::from_secs(5), - ) - .await; - - let network_labels = - LabelSet::new().with_label_type(LabelType::NetworkName("net1".to_string())); - let forwarded_bytes_before = metric_value( - &pm_center, - MetricName::TrafficBytesForwarded, - network_labels.clone(), - ); - let forwarded_packets_before = metric_value( - &pm_center, - MetricName::TrafficPacketsForwarded, - network_labels.clone(), - ); - - let mut transit_pkt = ZCPacket::new_with_payload(b"foreign-transit-disabled"); - transit_pkt.fill_peer_manager_hdr( - pma_net1.my_peer_id(), - pmb_net1.my_peer_id(), - PacketType::Data as u8, - ); - pma_net1 - .get_foreign_network_client() - .send_msg(transit_pkt, center_peer_id) - .await - .unwrap(); - - tokio::time::sleep(Duration::from_millis(300)).await; - - assert_eq!( - metric_value( - &pm_center, - MetricName::TrafficBytesForwarded, - network_labels.clone() - ), - forwarded_bytes_before - ); - assert_eq!( - metric_value( - &pm_center, - MetricName::TrafficPacketsForwarded, - network_labels - ), - forwarded_packets_before - ); - } - - #[tokio::test] - async fn foreign_network_transit_control_forwarding_records_control_forwarded_metrics() { - let pm_center = create_mock_peer_manager_with_mock_stun(NatType::Unknown).await; - let pma_net1 = create_mock_peer_manager_for_foreign_network("net1").await; - let pmb_net1 = create_mock_peer_manager_for_foreign_network("net1").await; - - connect_peer_manager(pma_net1.clone(), pm_center.clone()).await; - connect_peer_manager(pmb_net1.clone(), pm_center.clone()).await; - wait_route_appear(pma_net1.clone(), pmb_net1.clone()) - .await - .unwrap(); - - let mut flags = pm_center.get_global_ctx().get_flags(); - flags.disable_relay_data = true; - pm_center.get_global_ctx().set_flags(flags); - - let center_peer_id = pm_center - .get_foreign_network_manager() - .get_network_peer_id("net1") - .unwrap(); - let network_labels = - LabelSet::new().with_label_type(LabelType::NetworkName("net1".to_string())); - let forwarded_bytes_before = metric_value( - &pm_center, - MetricName::TrafficControlBytesForwarded, - network_labels.clone(), - ); - let forwarded_packets_before = metric_value( - &pm_center, - MetricName::TrafficControlPacketsForwarded, - network_labels.clone(), - ); - - let mut transit_pkt = ZCPacket::new_with_payload(b"foreign-control-transit"); - transit_pkt.fill_peer_manager_hdr( - pma_net1.my_peer_id(), - pmb_net1.my_peer_id(), - PacketType::RpcReq as u8, - ); - let transit_pkt_len = transit_pkt.buf_len() as u64; - pma_net1 - .get_foreign_network_client() - .send_msg(transit_pkt, center_peer_id) - .await - .unwrap(); - - wait_for_condition( - || { - let pm_center = pm_center.clone(); - let network_labels = network_labels.clone(); - async move { - metric_value( - &pm_center, - MetricName::TrafficControlBytesForwarded, - network_labels.clone(), - ) >= forwarded_bytes_before + transit_pkt_len - && metric_value( - &pm_center, - MetricName::TrafficControlPacketsForwarded, - network_labels.clone(), - ) > forwarded_packets_before - } - }, - Duration::from_secs(5), - ) - .await; - } - - #[tokio::test] - async fn failed_new_foreign_peer_conn_rolls_back_entry_maps() { - let pm_center = create_mock_peer_manager_with_mock_stun(NatType::Unknown).await; - let pma_net1 = create_mock_peer_manager_for_foreign_network("net1").await; - let foreign_mgr = pm_center.get_foreign_network_manager(); - - foreign_mgr.fail_next_add_peer_conn_after_entry_insert(); - - let (a_ring, b_ring) = crate::tunnel::ring::create_ring_tunnel_pair(); - let (client_ret, server_ret) = tokio::time::timeout(Duration::from_secs(5), async { - tokio::join!( - pma_net1.add_client_tunnel(a_ring, false), - pm_center.add_tunnel_as_server(b_ring, true) - ) - }) - .await - .unwrap(); - - assert!(client_ret.is_ok()); - assert!(server_ret.is_err()); - assert!(foreign_mgr.data.get_network_entry("net1").is_none()); - assert!( - foreign_mgr - .data - .get_peer_network(pma_net1.my_peer_id()) - .is_none() - ); - } - - #[tokio::test] - async fn foreign_network_peer_removed_clears_traffic_metric_peer_cache() { - let pm_center = create_mock_peer_manager_with_mock_stun(NatType::Unknown).await; - let pma_net1 = create_mock_peer_manager_for_foreign_network("net1").await; - - connect_peer_manager(pma_net1.clone(), pm_center.clone()).await; - wait_for_condition( - || { - let pm_center = pm_center.clone(); - async move { - pm_center - .get_foreign_network_manager() - .get_network_peer_id("net1") - .is_some() - } - }, - Duration::from_secs(5), - ) - .await; - - let entry = pm_center - .get_foreign_network_manager() - .data - .get_network_entry("net1") - .unwrap(); - - entry - .traffic_metrics - .record_rx(pma_net1.my_peer_id(), PacketType::Data as u8, 128) - .await; - - assert!( - entry - .traffic_metrics - .contains_peer_cache(pma_net1.my_peer_id()) - ); - - entry - .global_ctx - .issue_event(GlobalCtxEvent::PeerRemoved(pma_net1.my_peer_id())); - - wait_for_condition( - || { - let entry = entry.clone(); - let peer_id = pma_net1.my_peer_id(); - async move { !entry.traffic_metrics.contains_peer_cache(peer_id) } - }, - Duration::from_secs(5), - ) - .await; - } - - #[tokio::test] - async fn foreign_network_encapsulated_forwarding_records_tx_metrics() { - set_global_var!(OSPF_UPDATE_MY_GLOBAL_FOREIGN_NETWORK_INTERVAL_SEC, 1); - - let pm_center1 = create_mock_peer_manager_with_mock_stun(NatType::Unknown).await; - let pm_center2 = create_mock_peer_manager_with_mock_stun(NatType::Unknown).await; - - connect_peer_manager(pm_center1.clone(), pm_center2.clone()).await; - - let pma_net1 = create_mock_peer_manager_for_foreign_network("net1").await; - let pmb_net1 = create_mock_peer_manager_for_foreign_network("net1").await; - connect_peer_manager(pma_net1.clone(), pm_center1.clone()).await; - connect_peer_manager(pmb_net1.clone(), pm_center2.clone()).await; - wait_route_appear(pma_net1.clone(), pmb_net1.clone()) - .await - .unwrap(); - - let center_peer_id = pm_center1 - .get_foreign_network_manager() - .get_network_peer_id("net1") - .unwrap(); - - let mut encapsulated_tx_pkt = ZCPacket::new_with_payload(b"foreign-encap-tx"); - encapsulated_tx_pkt.fill_peer_manager_hdr( - center_peer_id, - pmb_net1.my_peer_id(), - PacketType::Data as u8, - ); - pma_net1 - .get_foreign_network_client() - .send_msg(encapsulated_tx_pkt, center_peer_id) - .await - .unwrap(); - - let network_labels = - LabelSet::new().with_label_type(LabelType::NetworkName("net1".to_string())); - let tx_instance_labels = network_labels - .clone() - .with_label_type(LabelType::ToInstanceId( - pmb_net1.get_global_ctx().get_id().to_string(), - )); - - wait_for_condition( - || { - let pm_center1 = pm_center1.clone(); - let network_labels = network_labels.clone(); - let tx_instance_labels = tx_instance_labels.clone(); - async move { - metric_value( - &pm_center1, - MetricName::TrafficBytesTx, - network_labels.clone(), - ) > 0 - && metric_value( - &pm_center1, - MetricName::TrafficBytesTxByInstance, - tx_instance_labels.clone(), - ) > 0 - } - }, - Duration::from_secs(5), - ) - .await; - } - - #[tokio::test] - async fn foreign_network_list_can_include_trusted_keys() { - let pm_center = create_mock_peer_manager_with_mock_stun(NatType::Unknown).await; - set_secure_mode_cfg(&pm_center.get_global_ctx(), true); - - let pma_net1 = create_mock_peer_manager_for_secure_foreign_network("net1").await; - let pmb_net1 = create_mock_peer_manager_for_secure_foreign_network("net1").await; - connect_peer_manager(pma_net1.clone(), pm_center.clone()).await; - connect_peer_manager(pmb_net1.clone(), pm_center.clone()).await; - wait_route_appear(pma_net1.clone(), pmb_net1.clone()) - .await - .unwrap(); - - let without_trusted_keys = pm_center - .get_foreign_network_manager() - .list_foreign_networks() - .await; - assert!( - without_trusted_keys.foreign_networks["net1"] - .trusted_keys - .is_empty() - ); - - let foreign_mgr = pm_center.get_foreign_network_manager(); - wait_for_condition( - || { - let foreign_mgr = foreign_mgr.clone(); - async move { - foreign_mgr - .list_foreign_networks_with_options(true) - .await - .foreign_networks - .get("net1") - .map(|entry| !entry.trusted_keys.is_empty()) - .unwrap_or(false) - } - }, - Duration::from_secs(5), - ) - .await; - - let with_trusted_keys = foreign_mgr.list_foreign_networks_with_options(true).await; - assert!( - !with_trusted_keys.foreign_networks["net1"] - .trusted_keys - .is_empty() - ); - } - - #[tokio::test] - async fn secure_center_can_serve_legacy_and_secure_foreign_networks() { - let pm_center = create_mock_peer_manager_with_mock_stun(NatType::Unknown).await; - set_secure_mode_cfg(&pm_center.get_global_ctx(), true); - - let legacy_a = create_mock_peer_manager_for_foreign_network("legacy-net").await; - let legacy_b = create_mock_peer_manager_for_foreign_network("legacy-net").await; - connect_peer_manager(legacy_a.clone(), pm_center.clone()).await; - connect_peer_manager(legacy_b.clone(), pm_center.clone()).await; - wait_route_appear(legacy_a.clone(), legacy_b.clone()) - .await - .unwrap(); - - let secure_a = create_mock_peer_manager_for_secure_foreign_network("secure-net").await; - let secure_b = create_mock_peer_manager_for_secure_foreign_network("secure-net").await; - connect_peer_manager(secure_a.clone(), pm_center.clone()).await; - connect_peer_manager(secure_b.clone(), pm_center.clone()).await; - wait_route_appear(secure_a.clone(), secure_b.clone()) - .await - .unwrap(); - - assert_eq!(2, legacy_a.list_routes().await.len()); - assert_eq!(2, legacy_b.list_routes().await.len()); - assert_eq!(2, secure_a.list_routes().await.len()); - assert_eq!(2, secure_b.list_routes().await.len()); - - let rpc_resp = pm_center - .get_foreign_network_manager() - .list_foreign_networks() - .await; - assert_eq!(2, rpc_resp.foreign_networks.len()); - assert_eq!(2, rpc_resp.foreign_networks["legacy-net"].peers.len()); - assert_eq!(2, rpc_resp.foreign_networks["secure-net"].peers.len()); - } - - #[tokio::test] - async fn credential_pubkey_trust_requires_ospf_credential_source() { - let global_ctx = get_mock_global_ctx_with_network(Some(NetworkIdentity::new( - "__access__".to_string(), - "access_secret".to_string(), - ))); - let foreign_network = NetworkIdentity::new("net1".to_string(), "net1_secret".to_string()); - let (pm_packet_sender, _pm_packet_recv) = create_packet_recv_chan(); - let entry = ForeignNetworkEntry::new( - foreign_network.clone(), - 1, - global_ctx.clone(), - false, - Arc::new(PeerSessionStore::new()), - pm_packet_sender, - ); - let pubkey = vec![7; 32]; - - entry.global_ctx.update_trusted_keys( - HashMap::from([( - pubkey.clone(), - crate::common::global_ctx::TrustedKeyMetadata { - source: TrustedKeySource::OspfNode, - expiry_unix: None, - }, - )]), - &foreign_network.network_name, - ); - assert!(!ForeignNetworkManager::is_credential_pubkey_trusted( - &entry, &pubkey - )); - - entry.global_ctx.update_trusted_keys( - HashMap::from([( - pubkey.clone(), - crate::common::global_ctx::TrustedKeyMetadata { - source: TrustedKeySource::OspfCredential, - expiry_unix: None, - }, - )]), - &foreign_network.network_name, - ); - assert!(ForeignNetworkManager::is_credential_pubkey_trusted( - &entry, &pubkey - )); - } - - #[tokio::test] - async fn foreign_entry_feature_flag_tracks_parent_disable_relay_data_toggle() { - let global_ctx = get_mock_global_ctx_with_network(Some(NetworkIdentity::new( - "__access__".to_string(), - "access_secret".to_string(), - ))); - let foreign_network = NetworkIdentity::new("net1".to_string(), "net1_secret".to_string()); - let (pm_packet_sender, _pm_packet_recv) = create_packet_recv_chan(); - let entry = ForeignNetworkEntry::new( - foreign_network, - 1, - global_ctx.clone(), - true, - Arc::new(PeerSessionStore::new()), - pm_packet_sender, - ); - assert!(!entry.global_ctx.get_feature_flags().avoid_relay_data); - - entry.run_parent_feature_flag_sync_routine().await; - - let mut flags = global_ctx.get_flags(); - flags.disable_relay_data = true; - global_ctx.set_flags(flags); - global_ctx.issue_event(GlobalCtxEvent::ConfigPatched(Default::default())); - - wait_for_condition( - || async { entry.global_ctx.get_feature_flags().avoid_relay_data }, - Duration::from_secs(2), - ) - .await; - - let mut flags = global_ctx.get_flags(); - flags.disable_relay_data = false; - global_ctx.set_flags(flags); - global_ctx.issue_event(GlobalCtxEvent::ConfigPatched(Default::default())); - - wait_for_condition( - || async { !entry.global_ctx.get_feature_flags().avoid_relay_data }, - Duration::from_secs(2), - ) - .await; - } - - #[tokio::test] - async fn foreign_entry_without_relay_data_keeps_avoid_feature_flag() { - let global_ctx = get_mock_global_ctx_with_network(Some(NetworkIdentity::new( - "__access__".to_string(), - "access_secret".to_string(), - ))); - let foreign_network = NetworkIdentity::new("net1".to_string(), "net1_secret".to_string()); - let (pm_packet_sender, _pm_packet_recv) = create_packet_recv_chan(); - let entry = ForeignNetworkEntry::new( - foreign_network, - 1, - global_ctx.clone(), - false, - Arc::new(PeerSessionStore::new()), - pm_packet_sender, - ); - - assert!(entry.global_ctx.get_feature_flags().avoid_relay_data); - - let mut flags = global_ctx.get_flags(); - flags.disable_relay_data = false; - global_ctx.set_flags(flags); - - ForeignNetworkEntry::sync_parent_relay_data_feature_flag( - &global_ctx, - &entry.global_ctx, - entry.relay_data, - ); - - assert!(entry.global_ctx.get_feature_flags().avoid_relay_data); - } - - #[test] - fn credential_trust_path_rejects_admin_identity() { - assert!(ForeignNetworkManager::should_reject_credential_trust_path( - PeerIdentityType::Admin - )); - assert!(!ForeignNetworkManager::should_reject_credential_trust_path( - PeerIdentityType::Credential - )); - assert!(!ForeignNetworkManager::should_reject_credential_trust_path( - PeerIdentityType::SharedNode - )); - } - - #[tokio::test] - async fn zero_digest_peer_cannot_bootstrap_foreign_network() { - let pm_center = create_mock_peer_manager_with_mock_stun(NatType::Unknown).await; - set_secure_mode_cfg(&pm_center.get_global_ctx(), true); - - let pma_net1 = create_mock_credential_peer_manager_for_foreign_network("net1").await; - - let (a_ring, b_ring) = crate::tunnel::ring::create_ring_tunnel_pair(); - let a_mgr_copy = pma_net1.clone(); - let client = tokio::spawn(async move { a_mgr_copy.add_client_tunnel(a_ring, false).await }); - let b_mgr_copy = pm_center.clone(); - let server = - tokio::spawn(async move { b_mgr_copy.add_tunnel_as_server(b_ring, true).await }); - - assert!(client.await.unwrap().is_ok()); - assert!(server.await.unwrap().is_err()); - assert!( - pm_center - .get_foreign_network_manager() - .list_foreign_networks() - .await - .foreign_networks - .is_empty() - ); - } - - async fn foreign_network_whitelist_helper(name: String) { - let pm_center = create_mock_peer_manager_with_mock_stun(NatType::Unknown).await; - tracing::debug!("pm_center: {:?}", pm_center.my_peer_id()); - let mut flag = pm_center.get_global_ctx().get_flags(); - flag.relay_network_whitelist = ["net1".to_string(), "net2*".to_string()].join(" "); - pm_center.get_global_ctx().set_flags(flag); - - let pma_net1 = create_mock_peer_manager_for_foreign_network(name.as_str()).await; - - let (a_ring, b_ring) = crate::tunnel::ring::create_ring_tunnel_pair(); - let b_mgr_copy = pm_center.clone(); - let s_ret = - tokio::spawn(async move { b_mgr_copy.add_tunnel_as_server(b_ring, true).await }); - - pma_net1.add_client_tunnel(a_ring, false).await.unwrap(); - - s_ret.await.unwrap().unwrap(); - } - - #[tokio::test] - async fn foreign_network_whitelist() { - foreign_network_whitelist_helper("net1".to_string()).await; - foreign_network_whitelist_helper("net2".to_string()).await; - foreign_network_whitelist_helper("net2abc".to_string()).await; - } - - #[tokio::test] - async fn only_relay_peer_rpc() { - let pm_center = create_mock_peer_manager_with_mock_stun(NatType::Unknown).await; - let mut flag = pm_center.get_global_ctx().get_flags(); - flag.relay_network_whitelist = "".to_string(); - flag.relay_all_peer_rpc = true; - pm_center.get_global_ctx().set_flags(flag); - tracing::debug!("pm_center: {:?}", pm_center.my_peer_id()); - - let pma_net1 = create_mock_peer_manager_for_foreign_network("net1").await; - let pmb_net1 = create_mock_peer_manager_for_foreign_network("net1").await; - tracing::debug!( - "pma_net1: {:?}, pmb_net1: {:?}", - pma_net1.my_peer_id(), - pmb_net1.my_peer_id() - ); - connect_peer_manager(pma_net1.clone(), pm_center.clone()).await; - connect_peer_manager(pmb_net1.clone(), pm_center.clone()).await; - wait_route_appear(pma_net1.clone(), pmb_net1.clone()) - .await - .unwrap(); - assert_eq!(2, pma_net1.list_routes().await.len()); - assert_eq!(2, pmb_net1.list_routes().await.len()); - } - - #[tokio::test] - #[should_panic] - async fn foreign_network_whitelist_fail() { - foreign_network_whitelist_helper("net3".to_string()).await; - } - - #[tokio::test] - async fn test_foreign_network_manager() { - let pm_center = create_mock_peer_manager_with_mock_stun(NatType::Unknown).await; - let pm_center2 = create_mock_peer_manager_with_mock_stun(NatType::Unknown).await; - connect_peer_manager(pm_center.clone(), pm_center2.clone()).await; - - tracing::debug!( - "pm_center: {:?}, pm_center2: {:?}", - pm_center.my_peer_id(), - pm_center2.my_peer_id() - ); - - let pma_net1 = create_mock_peer_manager_for_foreign_network("net1").await; - let pmb_net1 = create_mock_peer_manager_for_foreign_network("net1").await; - connect_peer_manager(pma_net1.clone(), pm_center.clone()).await; - connect_peer_manager(pmb_net1.clone(), pm_center.clone()).await; - - tracing::debug!( - "pma_net1: {:?}, pmb_net1: {:?}", - pma_net1.my_peer_id(), - pmb_net1.my_peer_id() - ); - - wait_route_appear(pma_net1.clone(), pmb_net1.clone()) - .await - .unwrap(); - - assert_eq!( - vec![ - pm_center - .get_foreign_network_manager() - .get_network_peer_id("net1") - .unwrap() - ], - pma_net1 - .get_foreign_network_client() - .get_peer_map() - .list_peers() - ); - assert_eq!( - vec![ - pm_center - .get_foreign_network_manager() - .get_network_peer_id("net1") - .unwrap() - ], - pmb_net1 - .get_foreign_network_client() - .get_peer_map() - .list_peers() - ); - - assert_eq!(2, pma_net1.list_routes().await.len()); - assert_eq!(2, pmb_net1.list_routes().await.len()); - - let pmc_net1 = create_mock_peer_manager_for_foreign_network("net1").await; - connect_peer_manager(pmc_net1.clone(), pm_center.clone()).await; - wait_route_appear(pma_net1.clone(), pmc_net1.clone()) - .await - .unwrap(); - wait_route_appear(pmb_net1.clone(), pmc_net1.clone()) - .await - .unwrap(); - assert_eq!(3, pmc_net1.list_routes().await.len()); - - tracing::debug!("pmc_net1: {:?}", pmc_net1.my_peer_id()); - - let pma_net2 = create_mock_peer_manager_for_foreign_network("net2").await; - let pmb_net2 = create_mock_peer_manager_for_foreign_network("net2").await; - tracing::debug!( - "pma_net2: {:?}, pmb_net2: {:?}", - pma_net2.my_peer_id(), - pmb_net2.my_peer_id() - ); - connect_peer_manager(pma_net2.clone(), pm_center.clone()).await; - connect_peer_manager(pmb_net2.clone(), pm_center.clone()).await; - wait_route_appear(pma_net2.clone(), pmb_net2.clone()) - .await - .unwrap(); - assert_eq!(2, pma_net2.list_routes().await.len()); - assert_eq!(2, pmb_net2.list_routes().await.len()); - - assert_eq!( - 5, - pm_center - .get_foreign_network_manager() - .data - .peer_network_map - .len() - ); - - assert_eq!( - 2, - pm_center - .get_foreign_network_manager() - .data - .network_peer_maps - .len() - ); - - let rpc_resp = pm_center - .get_foreign_network_manager() - .list_foreign_networks() - .await; - assert_eq!(2, rpc_resp.foreign_networks.len()); - assert_eq!(3, rpc_resp.foreign_networks["net1"].peers.len()); - assert_eq!(2, rpc_resp.foreign_networks["net2"].peers.len()); - - drop(pmb_net2); - tokio::time::sleep(std::time::Duration::from_secs(1)).await; - assert_eq!( - 4, - pm_center - .get_foreign_network_manager() - .data - .peer_network_map - .len() - ); - drop(pma_net2); - tokio::time::sleep(std::time::Duration::from_secs(1)).await; - assert_eq!( - 3, - pm_center - .get_foreign_network_manager() - .data - .peer_network_map - .len() - ); - assert_eq!( - 1, - pm_center - .get_foreign_network_manager() - .data - .network_peer_maps - .len() - ); - } - - #[tokio::test] - async fn test_disconnect_foreign_network() { - let pm_center = create_mock_peer_manager_with_mock_stun(NatType::Unknown).await; - tracing::debug!("pm_center: {:?}", pm_center.my_peer_id()); - let pma_net1 = create_mock_peer_manager_for_foreign_network("net1").await; - tracing::debug!("pma_net1: {:?}", pma_net1.my_peer_id(),); - - connect_peer_manager(pma_net1.clone(), pm_center.clone()).await; - - wait_for_condition( - || async { pma_net1.list_routes().await.len() == 1 }, - Duration::from_secs(5), - ) - .await; - - drop(pm_center); - wait_for_condition( - || async { pma_net1.list_routes().await.is_empty() }, - Duration::from_secs(5), - ) - .await; - } - - #[tokio::test] - async fn test_foreign_network_manager_cluster_simple() { - set_global_var!(OSPF_UPDATE_MY_GLOBAL_FOREIGN_NETWORK_INTERVAL_SEC, 1); - - let pm_center1 = create_mock_peer_manager_with_mock_stun(NatType::Unknown).await; - let pm_center2 = create_mock_peer_manager_with_mock_stun(NatType::Unknown).await; - - connect_peer_manager(pm_center1.clone(), pm_center2.clone()).await; - - let pma_net1 = create_mock_peer_manager_for_foreign_network("net1").await; - let pmb_net1 = create_mock_peer_manager_for_foreign_network("net1").await; - connect_peer_manager(pma_net1.clone(), pm_center1.clone()).await; - connect_peer_manager(pmb_net1.clone(), pm_center2.clone()).await; - - wait_route_appear(pma_net1.clone(), pmb_net1.clone()) - .await - .unwrap(); - - let pma_net2 = create_mock_peer_manager_for_foreign_network("net2").await; - let pmb_net2 = create_mock_peer_manager_for_foreign_network("net2").await; - connect_peer_manager(pma_net2.clone(), pm_center1.clone()).await; - connect_peer_manager(pmb_net2.clone(), pm_center2.clone()).await; - - wait_route_appear(pma_net2.clone(), pmb_net2.clone()) - .await - .unwrap(); - } - - #[tokio::test] - async fn test_foreign_network_manager_cluster_multiple_hops() { - set_global_var!(OSPF_UPDATE_MY_GLOBAL_FOREIGN_NETWORK_INTERVAL_SEC, 1); - - let pm_center1 = create_mock_peer_manager_with_mock_stun(NatType::Unknown).await; - let pm_center2 = create_mock_peer_manager_with_mock_stun(NatType::Unknown).await; - let pm_center3 = create_mock_peer_manager_with_mock_stun(NatType::Unknown).await; - let pm_center4 = create_mock_peer_manager_with_mock_stun(NatType::Unknown).await; - - connect_peer_manager(pm_center1.clone(), pm_center2.clone()).await; - connect_peer_manager(pm_center2.clone(), pm_center3.clone()).await; - connect_peer_manager(pm_center3.clone(), pm_center4.clone()).await; - - let pma_net1 = create_mock_peer_manager_for_foreign_network("net1").await; - let pmb_net1 = create_mock_peer_manager_for_foreign_network("net1").await; - connect_peer_manager(pma_net1.clone(), pm_center1.clone()).await; - connect_peer_manager(pmb_net1.clone(), pm_center3.clone()).await; - wait_route_appear(pma_net1.clone(), pmb_net1.clone()) - .await - .unwrap(); - let pmc_net1 = create_mock_peer_manager_for_foreign_network("net1").await; - connect_peer_manager(pmc_net1.clone(), pm_center4.clone()).await; - wait_route_appear(pma_net1.clone(), pmc_net1.clone()) - .await - .unwrap(); - - let pma_net2 = create_mock_peer_manager_for_foreign_network("net2").await; - let pmb_net2 = create_mock_peer_manager_for_foreign_network("net2").await; - connect_peer_manager(pma_net2.clone(), pm_center1.clone()).await; - connect_peer_manager(pmb_net2.clone(), pm_center4.clone()).await; - wait_route_appear(pma_net2.clone(), pmb_net2.clone()) - .await - .unwrap(); - drop(pmb_net2); - wait_for_condition( - || async { pma_net2.list_routes().await.len() == 1 }, - Duration::from_secs(5), - ) - .await; - } - - #[tokio::test] - async fn test_foreign_network_manager_cluster() { - set_global_var!(OSPF_UPDATE_MY_GLOBAL_FOREIGN_NETWORK_INTERVAL_SEC, 1); - - let pm_center1 = create_mock_peer_manager_with_mock_stun(NatType::Unknown).await; - let pm_center2 = create_mock_peer_manager_with_mock_stun(NatType::Unknown).await; - let pm_center3 = create_mock_peer_manager_with_mock_stun(NatType::Unknown).await; - - connect_peer_manager(pm_center1.clone(), pm_center2.clone()).await; - connect_peer_manager(pm_center2.clone(), pm_center3.clone()).await; - - tracing::debug!( - "pm_center: {:?}, pm_center2: {:?}", - pm_center1.my_peer_id(), - pm_center2.my_peer_id() - ); - - let pma_net1 = create_mock_peer_manager_for_foreign_network("net1").await; - let pmb_net1 = create_mock_peer_manager_for_foreign_network("net1").await; - connect_peer_manager(pma_net1.clone(), pm_center1.clone()).await; - connect_peer_manager(pmb_net1.clone(), pm_center2.clone()).await; - - tracing::debug!( - "pma_net1: {:?}, pmb_net1: {:?}", - pma_net1.my_peer_id(), - pmb_net1.my_peer_id() - ); - - wait_route_appear(pma_net1.clone(), pmb_net1.clone()) - .await - .unwrap(); - - assert_eq!(3, pma_net1.list_routes().await.len(),); - - let pmc_net1 = create_mock_peer_manager_for_foreign_network("net1").await; - connect_peer_manager(pmc_net1.clone(), pm_center3.clone()).await; - wait_route_appear(pma_net1.clone(), pmc_net1.clone()) - .await - .unwrap(); - assert_eq!(5, pma_net1.list_routes().await.len(),); - - println!( - "pm_center1: {:?}, pm_center2: {:?}, pm_center3: {:?}", - pm_center1.my_peer_id(), - pm_center2.my_peer_id(), - pm_center3.my_peer_id() - ); - println!( - "pma_net1: {:?}, pmb_net1: {:?}, pmc_net1: {:?}", - pma_net1.my_peer_id(), - pmb_net1.my_peer_id(), - pmc_net1.my_peer_id() - ); - - println!("drop pmc_net1, id: {:?}", pmc_net1.my_peer_id()); - - // foreign network node disconnect - drop(pmc_net1); - wait_for_condition( - || async { pma_net1.list_routes().await.len() == 3 }, - Duration::from_secs(15), - ) - .await; - - println!("drop pm_center1, id: {:?}", pm_center1.my_peer_id()); - drop(pm_center1); - wait_for_condition( - || async { pma_net1.list_routes().await.is_empty() }, - Duration::from_secs(5), - ) - .await; - wait_for_condition( - || async { - let n = pmb_net1 - .get_route() - .get_next_hop(pma_net1.my_peer_id()) - .await; - n.is_none() - }, - Duration::from_secs(5), - ) - .await; - wait_for_condition( - || async { - // only remain pmb center - pmb_net1.list_routes().await.len() == 1 - }, - Duration::from_secs(15), - ) - .await; - } - - #[tokio::test] - async fn test_foreign_network_manager_cluster_multi_net() { - set_global_var!(OSPF_UPDATE_MY_GLOBAL_FOREIGN_NETWORK_INTERVAL_SEC, 1); - - let pm_center1 = create_mock_peer_manager_with_mock_stun(NatType::Unknown).await; - let pm_center2 = create_mock_peer_manager_with_mock_stun(NatType::Unknown).await; - let pm_center3 = create_mock_peer_manager_with_mock_stun(NatType::Unknown).await; - - connect_peer_manager(pm_center1.clone(), pm_center2.clone()).await; - connect_peer_manager(pm_center2.clone(), pm_center3.clone()).await; - - let pma_net1 = create_mock_peer_manager_for_foreign_network("net1").await; - let pmb_net1 = create_mock_peer_manager_for_foreign_network("net1").await; - connect_peer_manager(pma_net1.clone(), pm_center1.clone()).await; - connect_peer_manager(pmb_net1.clone(), pm_center2.clone()).await; - - let pma_net2 = create_mock_peer_manager_for_foreign_network("net2").await; - let pmb_net2 = create_mock_peer_manager_for_foreign_network("net2").await; - connect_peer_manager(pma_net2.clone(), pm_center2.clone()).await; - connect_peer_manager(pmb_net2.clone(), pm_center3.clone()).await; - - let pma_net3 = create_mock_peer_manager_for_foreign_network("net3").await; - let pmb_net3 = create_mock_peer_manager_for_foreign_network("net3").await; - connect_peer_manager(pma_net3.clone(), pm_center1.clone()).await; - connect_peer_manager(pmb_net3.clone(), pm_center3.clone()).await; - - let pma_net4 = create_mock_peer_manager_for_foreign_network("net4").await; - let pmb_net4 = create_mock_peer_manager_for_foreign_network("net4").await; - let pmc_net4 = create_mock_peer_manager_for_foreign_network("net4").await; - connect_peer_manager(pma_net4.clone(), pm_center1.clone()).await; - connect_peer_manager(pmb_net4.clone(), pm_center2.clone()).await; - connect_peer_manager(pmc_net4.clone(), pm_center3.clone()).await; - - tokio::time::sleep(Duration::from_secs(5)).await; - - wait_route_appear(pma_net1.clone(), pmb_net1.clone()) - .await - .unwrap(); - wait_route_appear(pma_net2.clone(), pmb_net2.clone()) - .await - .unwrap(); - wait_route_appear(pma_net3.clone(), pmb_net3.clone()) - .await - .unwrap(); - wait_route_appear(pma_net4.clone(), pmb_net4.clone()) - .await - .unwrap(); - wait_route_appear(pma_net4.clone(), pmc_net4.clone()) - .await - .unwrap(); - wait_route_appear(pmb_net4.clone(), pmc_net4.clone()) - .await - .unwrap(); - - assert_eq!(3, pma_net1.list_routes().await.len()); - assert_eq!(3, pmb_net1.list_routes().await.len()); - - assert_eq!(3, pma_net2.list_routes().await.len()); - assert_eq!(3, pmb_net2.list_routes().await.len()); - - assert_eq!(3, pma_net3.list_routes().await.len()); - assert_eq!(3, pmb_net3.list_routes().await.len()); - - assert_eq!(5, pma_net4.list_routes().await.len()); - assert_eq!(5, pmb_net4.list_routes().await.len()); - assert_eq!(5, pmc_net4.list_routes().await.len()); - - drop(pm_center3); - tokio::time::sleep(Duration::from_secs(5)).await; - assert_eq!(1, pma_net2.list_routes().await.len()); - assert_eq!(1, pma_net3.list_routes().await.len()); - assert_eq!(3, pma_net4.list_routes().await.len()); - } - - #[tokio::test] - async fn test_foreign_network_manager_cluster_secret_mismatch() { - set_global_var!(OSPF_UPDATE_MY_GLOBAL_FOREIGN_NETWORK_INTERVAL_SEC, 1); - - let pm_center1 = create_mock_peer_manager_with_mock_stun(NatType::Unknown).await; - let pm_center2 = create_mock_peer_manager_with_mock_stun(NatType::Unknown).await; - let pm_center3 = create_mock_peer_manager_with_mock_stun(NatType::Unknown).await; - - connect_peer_manager(pm_center1.clone(), pm_center2.clone()).await; - connect_peer_manager(pm_center2.clone(), pm_center3.clone()).await; - - let pma_net4 = create_mock_peer_manager_for_foreign_network_ext("net4", "1").await; - let pmb_net4 = create_mock_peer_manager_for_foreign_network_ext("net4", "2").await; - let pmc_net4 = create_mock_peer_manager_for_foreign_network_ext("net4", "3").await; - connect_peer_manager(pma_net4.clone(), pm_center1.clone()).await; - connect_peer_manager(pmb_net4.clone(), pm_center2.clone()).await; - connect_peer_manager(pmc_net4.clone(), pm_center3.clone()).await; - - tokio::time::sleep(Duration::from_secs(5)).await; - assert_eq!(1, pma_net4.list_routes().await.len()); - assert_eq!(1, pmb_net4.list_routes().await.len()); - assert_eq!(1, pmc_net4.list_routes().await.len()); - } - - #[tokio::test] - async fn test_foreign_network_manager_cluster_max_direct_conns() { - set_global_var!(MAX_DIRECT_CONNS_PER_PEER_IN_FOREIGN_NETWORK, 1); - - let pm_center1 = create_mock_peer_manager_with_mock_stun(NatType::Unknown).await; - - let pma_net1 = create_mock_peer_manager_for_foreign_network("net1").await; - - connect_peer_manager(pma_net1.clone(), pm_center1.clone()).await; - wait_for_condition( - || async { pma_net1.list_routes().await.len() == 1 }, - Duration::from_secs(5), - ) - .await; - - println!("routes: {:?}", pma_net1.list_routes().await); - - let (a_ring, b_ring) = crate::tunnel::ring::create_ring_tunnel_pair(); - let a_mgr_copy = pma_net1.clone(); - tokio::spawn(async move { - a_mgr_copy.add_client_tunnel(a_ring, false).await.unwrap(); - }); - let b_mgr_copy = pm_center1.clone(); - - assert!(b_mgr_copy.add_tunnel_as_server(b_ring, true).await.is_err()); - } -} diff --git a/easytier/src/peers/mod.rs b/easytier/src/peers/mod.rs deleted file mode 100644 index c94a65de..00000000 --- a/easytier/src/peers/mod.rs +++ /dev/null @@ -1,73 +0,0 @@ -mod graph_algo; - -pub mod acl_filter; -pub mod credential_manager; -pub mod peer; -pub mod peer_conn; -pub mod peer_conn_ping; -pub mod peer_manager; -pub mod peer_map; -pub mod peer_ospf_route; -pub mod peer_rpc; -pub mod peer_rpc_service; -pub mod peer_session; -pub(crate) mod public_ipv6; -pub mod relay_peer_map; -pub mod route_trait; -pub mod rpc_service; -mod traffic_metrics; - -pub mod foreign_network_client; -pub mod foreign_network_manager; - -pub mod encrypt; -pub(crate) mod secure_datagram; - -pub mod peer_task; - -#[cfg(test)] -pub mod tests; - -use crate::tunnel::packet_def::ZCPacket; - -#[async_trait::async_trait] -#[auto_impl::auto_impl(Arc)] -pub trait PeerPacketFilter { - async fn try_process_packet_from_peer(&self, _zc_packet: ZCPacket) -> Option { - Some(_zc_packet) - } -} - -#[async_trait::async_trait] -#[auto_impl::auto_impl(Arc)] -pub trait NicPacketFilter { - async fn try_process_packet_from_nic(&self, data: &mut ZCPacket) -> bool; - - fn id(&self) -> String { - format!("{:p}", self) - } -} - -type BoxPeerPacketFilter = Box; -type BoxNicPacketFilter = Box; - -// pub type PacketRecvChan = tachyonix::Sender; -// pub type PacketRecvChanReceiver = tachyonix::Receiver; -// pub fn create_packet_recv_chan() -> (PacketRecvChan, PacketRecvChanReceiver) { -// tachyonix::channel(128) -// } -pub type PacketRecvChan = tokio::sync::mpsc::Sender; -pub type PacketRecvChanReceiver = tokio::sync::mpsc::Receiver; -pub fn create_packet_recv_chan() -> (PacketRecvChan, PacketRecvChanReceiver) { - tokio::sync::mpsc::channel(128) -} -pub async fn recv_packet_from_chan( - packet_recv_chan_receiver: &mut PacketRecvChanReceiver, -) -> Result { - packet_recv_chan_receiver - .recv() - .await - .ok_or(anyhow::anyhow!("recv_packet_from_chan failed")) -} - -pub const PUBLIC_SERVER_HOSTNAME_PREFIX: &str = "PublicServer_"; diff --git a/easytier/src/peers/peer.rs b/easytier/src/peers/peer.rs deleted file mode 100644 index 63af93f8..00000000 --- a/easytier/src/peers/peer.rs +++ /dev/null @@ -1,566 +0,0 @@ -use std::sync::Arc; - -use crossbeam::atomic::AtomicCell; -use dashmap::{DashMap, DashSet}; -use parking_lot::RwLock; - -use tokio::{select, sync::mpsc}; - -use tracing::Instrument; - -use super::{ - PacketRecvChan, - peer_conn::{PeerConn, PeerConnId}, -}; -use crate::{common::shrink_dashmap, proto::api::instance::PeerConnInfo}; -use crate::{ - common::{ - PeerId, - error::Error, - global_ctx::{ArcGlobalCtx, GlobalCtxEvent}, - }, - proto::peer_rpc::PeerIdentityType, - tunnel::packet_def::ZCPacket, -}; -use tokio_util::task::AbortOnDropHandle; - -type ArcPeerConn = Arc; -type ConnMap = Arc>; - -pub struct Peer { - pub peer_node_id: PeerId, - conns: ConnMap, - global_ctx: ArcGlobalCtx, - - packet_recv_chan: PacketRecvChan, - - close_event_sender: mpsc::Sender, - close_event_listener: AbortOnDropHandle<()>, - - shutdown_notifier: Arc, - - default_conn_id: Arc>, - peer_identity_type: Arc>>, - peer_public_key: Arc>>>, - default_conn_id_clear_task: AbortOnDropHandle<()>, -} - -impl Peer { - pub fn new( - peer_node_id: PeerId, - packet_recv_chan: PacketRecvChan, - global_ctx: ArcGlobalCtx, - ) -> Self { - let conns: ConnMap = Arc::new(DashMap::new()); - let (close_event_sender, mut close_event_receiver) = mpsc::channel(10); - let shutdown_notifier = Arc::new(tokio::sync::Notify::new()); - let peer_identity_type = Arc::new(AtomicCell::new(None)); - let peer_identity_type_copy = peer_identity_type.clone(); - let peer_public_key = Arc::new(RwLock::new(None)); - let peer_public_key_copy = peer_public_key.clone(); - - let conns_copy = conns.clone(); - let shutdown_notifier_copy = shutdown_notifier.clone(); - let global_ctx_copy = global_ctx.clone(); - let close_event_listener = AbortOnDropHandle::new(tokio::spawn( - async move { - loop { - select! { - ret = close_event_receiver.recv() => { - if ret.is_none() { - break; - } - let ret = ret.unwrap(); - tracing::warn!( - ?peer_node_id, - ?ret, - "notified that peer conn is closed", - ); - - if let Some((_, conn)) = conns_copy.remove(&ret) { - global_ctx_copy.issue_event(GlobalCtxEvent::PeerConnRemoved( - conn.get_conn_info(), - )); - shrink_dashmap(&conns_copy, Some(4)); - if conns_copy.is_empty() { - peer_identity_type_copy.store(None); - *peer_public_key_copy.write() = None; - } - } - } - - _ = shutdown_notifier_copy.notified() => { - close_event_receiver.close(); - tracing::warn!(?peer_node_id, "peer close event listener notified"); - } - } - } - tracing::info!("peer {} close event listener exit", peer_node_id); - } - .instrument(tracing::info_span!( - "peer_close_event_listener", - ?peer_node_id, - )), - )); - - let default_conn_id = Arc::new(AtomicCell::new(PeerConnId::default())); - - let conns_copy = conns.clone(); - let default_conn_id_copy = default_conn_id.clone(); - let default_conn_id_clear_task = AbortOnDropHandle::new(tokio::spawn(async move { - loop { - tokio::time::sleep(std::time::Duration::from_secs(5)).await; - if conns_copy.len() > 1 { - default_conn_id_copy.store(PeerConnId::default()); - } - } - })); - - Peer { - peer_node_id, - conns, - packet_recv_chan, - global_ctx, - - close_event_sender, - close_event_listener, - - shutdown_notifier, - default_conn_id, - peer_identity_type, - peer_public_key, - default_conn_id_clear_task, - } - } - - pub async fn add_peer_conn(&self, mut conn: PeerConn) -> Result<(), Error> { - let conn_identity_type = conn.get_peer_identity_type(); - let peer_identity_type = self.peer_identity_type.load(); - if let Some(peer_identity_type) = peer_identity_type { - if peer_identity_type != conn_identity_type { - return Err(Error::SecretKeyError(format!( - "peer identity type mismatch. peer: {:?}, conn: {:?}", - peer_identity_type, conn_identity_type - ))); - } - } else { - self.peer_identity_type.store(Some(conn_identity_type)); - } - - let close_notifier = conn.get_close_notifier(); - let conn_info = conn.get_conn_info(); - let conn_pubkey = conn_info.noise_remote_static_pubkey.clone(); - { - let mut peer_pubkey = self.peer_public_key.write(); - if let Some(existing_pubkey) = peer_pubkey.as_ref() { - if existing_pubkey != &conn_pubkey { - return Err(Error::SecretKeyError(format!( - "peer public key mismatch. peer_id: {}, existing_len: {}, new_len: {}", - self.peer_node_id, - existing_pubkey.len(), - conn_pubkey.len() - ))); - } - } else { - *peer_pubkey = Some(conn_pubkey); - } - } - - conn.start_recv_loop(self.packet_recv_chan.clone()).await; - conn.start_pingpong(); - self.conns.insert(conn.get_conn_id(), Arc::new(conn)); - - let close_event_sender = self.close_event_sender.clone(); - tokio::spawn(async move { - let conn_id = close_notifier.get_conn_id(); - if let Some(mut waiter) = close_notifier.get_waiter().await { - let _ = waiter.recv().await; - } - if let Err(e) = close_event_sender.send(conn_id).await { - tracing::warn!(?conn_id, "failed to send close event: {}", e); - } - }); - - self.global_ctx - .issue_event(GlobalCtxEvent::PeerConnAdded(conn_info)); - Ok(()) - } - - async fn select_conn(&self) -> Option { - let default_conn_id = self.default_conn_id.load(); - if let Some(conn) = self.conns.get(&default_conn_id) { - return Some(conn.clone()); - } - - // find a conn with the smallest latency - let mut min_latency = u64::MAX; - for conn in self.conns.iter() { - let latency = conn.value().get_stats().latency_us; - if latency < min_latency { - min_latency = latency; - self.default_conn_id.store(conn.get_conn_id()); - } - } - - self.conns - .get(&self.default_conn_id.load()) - .map(|conn| conn.clone()) - } - - pub async fn send_msg(&self, msg: ZCPacket) -> Result<(), Error> { - let Some(conn) = self.select_conn().await else { - return Err(Error::PeerNoConnectionError(self.peer_node_id)); - }; - conn.send_msg(msg).await?; - - Ok(()) - } - - pub async fn close_peer_conn(&self, conn_id: &PeerConnId) -> Result<(), Error> { - let has_key = self.conns.contains_key(conn_id); - if !has_key { - return Err(Error::NotFound); - } - self.close_event_sender.send(*conn_id).await.unwrap(); - Ok(()) - } - - pub async fn list_peer_conns(&self) -> Vec { - let mut conns = vec![]; - for conn in self.conns.iter() { - // do not lock here, otherwise it will cause dashmap deadlock - conns.push(conn.clone()); - } - - let mut ret = Vec::new(); - for conn in conns { - let info = conn.get_conn_info(); - if !info.is_closed { - ret.push(info); - } else { - let conn_id = info.conn_id.parse().unwrap(); - let _ = self.close_peer_conn(&conn_id).await; - } - } - ret - } - - pub fn has_live_conns(&self) -> bool { - self.conns.iter().any(|entry| !entry.value().is_closed()) - } - - pub fn has_directly_connected_conn(&self) -> bool { - self.conns - .iter() - .any(|entry| !entry.value().is_closed() && !entry.value().is_hole_punched()) - } - - pub fn get_directly_connections(&self) -> DashSet { - self.conns - .iter() - .filter(|entry| !(entry.value()).is_hole_punched()) - .map(|entry| (entry.value()).get_conn_id()) - .collect() - } - - pub fn get_default_conn_id(&self) -> PeerConnId { - self.default_conn_id.load() - } - - pub fn get_peer_identity_type(&self) -> Option { - self.peer_identity_type.load() - } - - pub fn get_peer_public_key(&self) -> Option> { - self.peer_public_key.read().clone() - } -} - -// pritn on drop -impl Drop for Peer { - fn drop(&mut self) { - self.conns.retain(|_, conn| { - self.global_ctx - .issue_event(GlobalCtxEvent::PeerConnRemoved(conn.get_conn_info())); - false - }); - self.shutdown_notifier.notify_one(); - tracing::info!("peer {} drop", self.peer_node_id); - } -} - -#[cfg(test)] -mod tests { - use base64::prelude::{BASE64_STANDARD, Engine as _}; - use rand::rngs::OsRng; - use std::sync::Arc; - use tokio::time::timeout; - - use crate::{ - common::{ - config::{NetworkIdentity, PeerConfig}, - global_ctx::{GlobalCtx, tests::get_mock_global_ctx}, - new_peer_id, - }, - peers::{create_packet_recv_chan, peer_conn::PeerConn, peer_session::PeerSessionStore}, - proto::common::SecureModeConfig, - tunnel::ring::create_ring_tunnel_pair, - }; - - use super::Peer; - - fn set_secure_mode_cfg(global_ctx: &GlobalCtx, enabled: bool) { - if !enabled { - global_ctx.config.set_secure_mode(None); - } else { - let private = x25519_dalek::StaticSecret::random_from_rng(OsRng); - let public = x25519_dalek::PublicKey::from(&private); - global_ctx.config.set_secure_mode(Some(SecureModeConfig { - enabled: true, - local_private_key: Some(BASE64_STANDARD.encode(private.as_bytes())), - local_public_key: Some(BASE64_STANDARD.encode(public.as_bytes())), - })); - } - } - - #[tokio::test] - async fn close_peer() { - let (local_packet_send, _local_packet_recv) = create_packet_recv_chan(); - let (remote_packet_send, _remote_packet_recv) = create_packet_recv_chan(); - let global_ctx = get_mock_global_ctx(); - let local_peer = Peer::new(new_peer_id(), local_packet_send, global_ctx.clone()); - let remote_peer = Peer::new(new_peer_id(), remote_packet_send, global_ctx.clone()); - - let ps = Arc::new(PeerSessionStore::new()); - let (local_tunnel, remote_tunnel) = create_ring_tunnel_pair(); - let mut local_peer_conn = PeerConn::new( - local_peer.peer_node_id, - global_ctx.clone(), - local_tunnel, - ps.clone(), - ); - let mut remote_peer_conn = PeerConn::new( - remote_peer.peer_node_id, - global_ctx.clone(), - remote_tunnel, - ps.clone(), - ); - - assert!(!local_peer_conn.handshake_done()); - assert!(!remote_peer_conn.handshake_done()); - - let (a, b) = tokio::join!( - local_peer_conn.do_handshake_as_client(), - remote_peer_conn.do_handshake_as_server() - ); - a.unwrap(); - b.unwrap(); - - let local_conn_id = local_peer_conn.get_conn_id(); - - local_peer.add_peer_conn(local_peer_conn).await.unwrap(); - remote_peer.add_peer_conn(remote_peer_conn).await.unwrap(); - - assert_eq!(local_peer.list_peer_conns().await.len(), 1); - assert_eq!(remote_peer.list_peer_conns().await.len(), 1); - - let close_handler = - tokio::spawn(async move { local_peer.close_peer_conn(&local_conn_id).await }); - - // wait for remote peer conn close - timeout(std::time::Duration::from_secs(5), async { - while !remote_peer.list_peer_conns().await.is_empty() { - tokio::time::sleep(std::time::Duration::from_millis(100)).await; - } - }) - .await - .unwrap(); - - println!("wait for close handler"); - close_handler.await.unwrap().unwrap(); - } - - #[tokio::test] - async fn reject_peer_conn_with_mismatched_identity_type() { - let (packet_send, _packet_recv) = create_packet_recv_chan(); - let global_ctx = get_mock_global_ctx(); - let local_peer_id = new_peer_id(); - let remote_peer_id = new_peer_id(); - let peer = Peer::new(remote_peer_id, packet_send, global_ctx); - - let ps = Arc::new(PeerSessionStore::new()); - - let (shared_client_tunnel, shared_server_tunnel) = create_ring_tunnel_pair(); - let shared_client_ctx = get_mock_global_ctx(); - let shared_server_ctx = get_mock_global_ctx(); - shared_client_ctx - .config - .set_network_identity(NetworkIdentity::new("net1".to_string(), "sec2".to_string())); - shared_server_ctx - .config - .set_network_identity(NetworkIdentity { - network_name: "net2".to_string(), - network_secret: None, - network_secret_digest: None, - }); - set_secure_mode_cfg(&shared_client_ctx, true); - set_secure_mode_cfg(&shared_server_ctx, true); - let remote_url: url::Url = shared_client_tunnel - .info() - .unwrap() - .remote_addr - .unwrap() - .url - .parse() - .unwrap(); - shared_client_ctx.config.set_peers(vec![PeerConfig { - uri: remote_url, - peer_public_key: Some( - shared_server_ctx - .config - .get_secure_mode() - .unwrap() - .local_public_key - .unwrap(), - ), - }]); - let mut shared_client_conn = PeerConn::new( - local_peer_id, - shared_client_ctx, - Box::new(shared_client_tunnel), - ps.clone(), - ); - let mut shared_server_conn = PeerConn::new( - remote_peer_id, - shared_server_ctx, - Box::new(shared_server_tunnel), - ps.clone(), - ); - let (c1, s1) = tokio::join!( - shared_client_conn.do_handshake_as_client(), - shared_server_conn.do_handshake_as_server() - ); - c1.unwrap(); - s1.unwrap(); - assert_eq!( - shared_client_conn.get_peer_identity_type(), - crate::proto::peer_rpc::PeerIdentityType::SharedNode - ); - - let (admin_client_tunnel, admin_server_tunnel) = create_ring_tunnel_pair(); - let admin_client_ctx = get_mock_global_ctx(); - let admin_server_ctx = get_mock_global_ctx(); - admin_client_ctx - .config - .set_network_identity(NetworkIdentity::new("net1".to_string(), "sec2".to_string())); - admin_server_ctx - .config - .set_network_identity(NetworkIdentity::new("net1".to_string(), "sec2".to_string())); - set_secure_mode_cfg(&admin_client_ctx, true); - set_secure_mode_cfg(&admin_server_ctx, true); - let mut admin_client_conn = PeerConn::new( - local_peer_id, - admin_client_ctx, - Box::new(admin_client_tunnel), - Arc::new(PeerSessionStore::new()), - ); - let mut admin_server_conn = PeerConn::new( - remote_peer_id, - admin_server_ctx, - Box::new(admin_server_tunnel), - Arc::new(PeerSessionStore::new()), - ); - let (c2, s2) = tokio::join!( - admin_client_conn.do_handshake_as_client(), - admin_server_conn.do_handshake_as_server() - ); - c2.unwrap(); - s2.unwrap(); - assert_eq!( - admin_client_conn.get_peer_identity_type(), - crate::proto::peer_rpc::PeerIdentityType::Admin - ); - - peer.add_peer_conn(shared_client_conn).await.unwrap(); - let ret = peer.add_peer_conn(admin_client_conn).await; - assert!(ret.is_err()); - } - - #[tokio::test] - async fn reject_peer_conn_with_mismatched_public_key() { - let (packet_send, _packet_recv) = create_packet_recv_chan(); - let local_peer_id = new_peer_id(); - let remote_peer_id = new_peer_id(); - let peer = Peer::new(remote_peer_id, packet_send, get_mock_global_ctx()); - let ps = Arc::new(PeerSessionStore::new()); - - let (client_tunnel_1, server_tunnel_1) = create_ring_tunnel_pair(); - let client_ctx_1 = get_mock_global_ctx(); - let server_ctx_1 = get_mock_global_ctx(); - client_ctx_1 - .config - .set_network_identity(NetworkIdentity::new("net1".to_string(), "sec1".to_string())); - server_ctx_1 - .config - .set_network_identity(NetworkIdentity::new("net1".to_string(), "sec1".to_string())); - set_secure_mode_cfg(&client_ctx_1, true); - set_secure_mode_cfg(&server_ctx_1, true); - let mut client_conn_1 = PeerConn::new( - local_peer_id, - client_ctx_1, - Box::new(client_tunnel_1), - ps.clone(), - ); - let mut server_conn_1 = PeerConn::new( - remote_peer_id, - server_ctx_1, - Box::new(server_tunnel_1), - ps.clone(), - ); - let (c1, s1) = tokio::join!( - client_conn_1.do_handshake_as_client(), - server_conn_1.do_handshake_as_server() - ); - c1.unwrap(); - s1.unwrap(); - - let (client_tunnel_2, server_tunnel_2) = create_ring_tunnel_pair(); - let client_ctx_2 = get_mock_global_ctx(); - let server_ctx_2 = get_mock_global_ctx(); - client_ctx_2 - .config - .set_network_identity(NetworkIdentity::new("net1".to_string(), "sec1".to_string())); - server_ctx_2 - .config - .set_network_identity(NetworkIdentity::new("net1".to_string(), "sec1".to_string())); - set_secure_mode_cfg(&client_ctx_2, true); - set_secure_mode_cfg(&server_ctx_2, true); - let mut client_conn_2 = PeerConn::new( - local_peer_id, - client_ctx_2, - Box::new(client_tunnel_2), - Arc::new(PeerSessionStore::new()), - ); - let mut server_conn_2 = PeerConn::new( - remote_peer_id, - server_ctx_2, - Box::new(server_tunnel_2), - Arc::new(PeerSessionStore::new()), - ); - let (c2, s2) = tokio::join!( - client_conn_2.do_handshake_as_client(), - server_conn_2.do_handshake_as_server() - ); - c2.unwrap(); - s2.unwrap(); - - let pubkey_1 = client_conn_1.get_conn_info().noise_remote_static_pubkey; - let pubkey_2 = client_conn_2.get_conn_info().noise_remote_static_pubkey; - assert_ne!(pubkey_1, pubkey_2); - - peer.add_peer_conn(client_conn_1).await.unwrap(); - assert_eq!(peer.get_peer_public_key(), Some(pubkey_1)); - let ret = peer.add_peer_conn(client_conn_2).await; - assert!(ret.is_err()); - } -} diff --git a/easytier/src/peers/peer_manager.rs b/easytier/src/peers/peer_manager.rs deleted file mode 100644 index 42ff5396..00000000 --- a/easytier/src/peers/peer_manager.rs +++ /dev/null @@ -1,3756 +0,0 @@ -use anyhow::Context; -use async_trait::async_trait; -use cidr::{Ipv4Cidr, Ipv6Cidr}; -use dashmap::DashMap; -use quanta::Instant; -use std::collections::BTreeSet; -use std::{ - fmt::Debug, - net::{IpAddr, Ipv4Addr, Ipv6Addr}, - sync::{Arc, Weak, atomic::AtomicBool}, - time::{Duration, SystemTime}, -}; - -use tokio::sync::{Mutex, RwLock}; -use tokio::{ - sync::mpsc::{self, UnboundedReceiver, UnboundedSender}, - task::JoinSet, -}; - -use crate::{ - common::{ - PeerId, - compressor::{Compressor as _, DefaultCompressor}, - constants::EASYTIER_VERSION, - error::Error, - global_ctx::{ArcGlobalCtx, GlobalCtxEvent, NetworkIdentity}, - shrink_dashmap, - stats_manager::{CounterHandle, LabelSet, LabelType, MetricName}, - stun::StunInfoCollectorTrait, - }, - peers::{ - PeerPacketFilter, - peer_conn::PeerConn, - peer_rpc::PeerRpcManagerTransport, - peer_session::PeerSessionStore, - recv_packet_from_chan, - route_trait::{ForeignNetworkRouteInfoMap, MockRoute, NextHopPolicy, RouteInterface}, - traffic_metrics::{ - InstanceLabelKind, LogicalTrafficMetrics, TrafficKind, TrafficMetricRecorder, - is_relay_data_packet_type, route_peer_info_instance_id, traffic_kind, - }, - }, - proto::{ - api::instance::{ - self, ListGlobalForeignNetworkResponse, - list_global_foreign_network_response::OneForeignNetwork, - }, - peer_rpc::{ - ForeignNetworkRouteInfoEntry, ForeignNetworkRouteInfoKey, PeerIdentityType, - RouteForeignNetworkSummary, - }, - }, - tunnel::{ - self, Tunnel, TunnelConnector, - packet_def::{CompressorAlgo, PacketType, ZCPacket}, - }, -}; - -use super::{ - BoxNicPacketFilter, BoxPeerPacketFilter, PacketRecvChan, PacketRecvChanReceiver, - create_packet_recv_chan, - encrypt::{Encryptor, NullCipher}, - foreign_network_client::ForeignNetworkClient, - foreign_network_manager::{ForeignNetworkManager, GlobalForeignNetworkAccessor}, - peer_conn::PeerConnId, - peer_map::PeerMap, - peer_ospf_route::PeerRoute, - peer_rpc::PeerRpcManager, - peer_task::ExternalTaskSignal, - relay_peer_map::RelayPeerMap, - route_trait::{ArcRoute, Route}, -}; - -struct RpcTransport { - my_peer_id: PeerId, - peers: Weak, - // TODO: this seems can be removed - foreign_peers: Mutex>>, - - packet_recv: Mutex>, - peer_rpc_tspt_sender: UnboundedSender, - - encryptor: Arc, - is_secure_mode_enabled: bool, -} - -#[async_trait::async_trait] -impl PeerRpcManagerTransport for RpcTransport { - fn my_peer_id(&self) -> PeerId { - self.my_peer_id - } - - async fn send(&self, mut msg: ZCPacket, dst_peer_id: PeerId) -> Result<(), Error> { - let peers = self.peers.upgrade().ok_or(Error::Unknown)?; - // NOTE: if route info is not exchanged, this will return None. treat it as public server. - let is_dst_peer_public_server = peers - .get_route_peer_info(dst_peer_id) - .await - .and_then(|x| x.feature_flag.map(|x| x.is_public_server)) - // if dst is directly connected, it's must not public server - .unwrap_or(!peers.has_peer(dst_peer_id)); - if !is_dst_peer_public_server && !self.is_secure_mode_enabled { - self.encryptor - .encrypt(&mut msg) - .with_context(|| "encrypt failed")?; - } - // send to self and this packet will be forwarded in peer_recv loop - peers.send_msg_directly(msg, self.my_peer_id).await - } - - async fn recv(&self) -> Result { - if let Some(o) = self.packet_recv.lock().await.recv().await { - Ok(o) - } else { - Err(Error::Unknown) - } - } -} - -pub enum RouteAlgoType { - Ospf, - None, -} - -enum RouteAlgoInst { - Ospf(Arc), - None, -} - -impl Clone for RouteAlgoInst { - fn clone(&self) -> Self { - match self { - RouteAlgoInst::Ospf(route) => RouteAlgoInst::Ospf(route.clone()), - RouteAlgoInst::None => RouteAlgoInst::None, - } - } -} - -struct SelfTxCounters { - self_tx_packets: CounterHandle, - self_tx_bytes: CounterHandle, - compress_tx_bytes_before: CounterHandle, - compress_tx_bytes_after: CounterHandle, -} - -pub struct PeerManager { - my_peer_id: PeerId, - - global_ctx: ArcGlobalCtx, - nic_channel: PacketRecvChan, - - tasks: Mutex>, - - packet_recv: Arc>>, - - peers: Arc, - - peer_rpc_mgr: Arc, - peer_rpc_tspt: Arc, - - peer_packet_process_pipeline: Arc>>, - nic_packet_process_pipeline: Arc>>, - - route_algo_inst: RouteAlgoInst, - - foreign_network_manager: Arc, - foreign_network_client: Arc, - relay_peer_map: Arc, - - encryptor: Arc, - data_compress_algo: CompressorAlgo, - - exit_nodes: RwLock>, - - reserved_my_peer_id_map: DashMap, - recent_have_traffic: Arc>, - p2p_demand_notify: Arc, - - allow_loopback_tunnel: AtomicBool, - - self_tx_counters: SelfTxCounters, - traffic_metrics: Arc, - - peer_session_store: Arc, - is_secure_mode_enabled: bool, -} - -impl Debug for PeerManager { - fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { - f.debug_struct("PeerManager") - .field("my_peer_id", &self.my_peer_id()) - .field("instance_name", &self.global_ctx.inst_name) - .field("net_ns", &self.global_ctx.net_ns.name()) - .finish() - } -} - -impl PeerManager { - // Keep lazy-p2p demand alive across the 5s task rescan interval and a full on-demand - // connect attempt, without retaining extra per-task state in the hot path. - const RECENT_HAVE_TRAFFIC_TTL: Duration = Duration::from_secs(30); - - fn should_mark_recent_traffic_for_fanout(total_dst_peers: usize) -> bool { - total_dst_peers <= 1 - } - - fn gc_recent_traffic_entries( - recent_have_traffic: &DashMap, - now: Instant, - mut has_directly_connected_conn: F, - ) where - F: FnMut(PeerId) -> bool, - { - let mut to_remove = Vec::new(); - for entry in recent_have_traffic.iter() { - let peer_id = *entry.key(); - let expired = - now.saturating_duration_since(*entry.value()) > Self::RECENT_HAVE_TRAFFIC_TTL; - if expired || has_directly_connected_conn(peer_id) { - to_remove.push(peer_id); - } - } - - if !to_remove.is_empty() { - for peer_id in to_remove { - recent_have_traffic.remove(&peer_id); - } - shrink_dashmap(recent_have_traffic, None); - } - } - - pub fn new( - route_algo: RouteAlgoType, - global_ctx: ArcGlobalCtx, - nic_channel: PacketRecvChan, - ) -> Self { - let my_peer_id = rand::random(); - - let (packet_send, packet_recv) = create_packet_recv_chan(); - let peers = Arc::new(PeerMap::new( - packet_send.clone(), - global_ctx.clone(), - my_peer_id, - )); - let peer_session_store = Arc::new(PeerSessionStore::new()); - - let encryptor = if global_ctx.get_flags().enable_encryption { - // 只有在启用加密时才使用工厂函数选择算法 - let algorithm = &global_ctx.get_flags().encryption_algorithm; - super::encrypt::create_encryptor( - algorithm, - global_ctx.get_128_key(), - global_ctx.get_256_key(), - ) - } else { - // disable_encryption = true 时使用 NullCipher - Arc::new(NullCipher) - }; - - if global_ctx - .check_network_in_whitelist(&global_ctx.get_network_name()) - .is_err() - { - // if local network is not in whitelist, avoid relay data when exist any other route path - global_ctx.set_avoid_relay_data_preference(true); - } - - let is_secure_mode_enabled = global_ctx - .config - .get_secure_mode() - .map(|cfg| cfg.enabled) - .unwrap_or(false); - - // TODO: remove these because we have impl pipeline processor. - let (peer_rpc_tspt_sender, peer_rpc_tspt_recv) = mpsc::unbounded_channel(); - let rpc_tspt = Arc::new(RpcTransport { - my_peer_id, - peers: Arc::downgrade(&peers), - foreign_peers: Mutex::new(None), - packet_recv: Mutex::new(peer_rpc_tspt_recv), - peer_rpc_tspt_sender, - encryptor: encryptor.clone(), - is_secure_mode_enabled, - }); - let peer_rpc_mgr = Arc::new(PeerRpcManager::new_with_stats_manager( - rpc_tspt.clone(), - global_ctx.stats_manager().clone(), - )); - - let route_algo_inst = match route_algo { - RouteAlgoType::Ospf => RouteAlgoInst::Ospf(PeerRoute::new( - my_peer_id, - global_ctx.clone(), - peer_rpc_mgr.clone(), - )), - RouteAlgoType::None => RouteAlgoInst::None, - }; - - let foreign_network_manager = Arc::new(ForeignNetworkManager::new( - my_peer_id, - global_ctx.clone(), - peer_session_store.clone(), - packet_send.clone(), - Self::build_foreign_network_manager_accessor(&peers), - )); - let foreign_network_client = Arc::new(ForeignNetworkClient::new( - global_ctx.clone(), - packet_send, - peer_rpc_mgr.clone(), - my_peer_id, - )); - - let data_compress_algo = global_ctx - .get_flags() - .data_compress_algo() - .try_into() - .expect("invalid data compress algo, maybe some features not enabled"); - - let exit_nodes = global_ctx.config.get_exit_nodes(); - - let stats_manager = global_ctx.stats_manager(); - let network_name = global_ctx.get_network_name(); - let traffic_tx_metrics = Arc::new(LogicalTrafficMetrics::new( - stats_manager.clone(), - network_name.clone(), - MetricName::TrafficBytesTx, - MetricName::TrafficPacketsTx, - MetricName::TrafficBytesTxByInstance, - MetricName::TrafficPacketsTxByInstance, - InstanceLabelKind::To, - )); - let traffic_control_tx_metrics = Arc::new(LogicalTrafficMetrics::new( - stats_manager.clone(), - network_name.clone(), - MetricName::TrafficControlBytesTx, - MetricName::TrafficControlPacketsTx, - MetricName::TrafficControlBytesTxByInstance, - MetricName::TrafficControlPacketsTxByInstance, - InstanceLabelKind::To, - )); - let relay_peer_map = RelayPeerMap::new( - peers.clone(), - Some(foreign_network_client.clone()), - global_ctx.clone(), - my_peer_id, - peer_session_store.clone(), - ); - let self_tx_counters = SelfTxCounters { - self_tx_packets: stats_manager.get_counter( - MetricName::TrafficPacketsSelfTx, - LabelSet::new().with_label_type(LabelType::NetworkName(network_name.clone())), - ), - self_tx_bytes: stats_manager.get_counter( - MetricName::TrafficBytesSelfTx, - LabelSet::new().with_label_type(LabelType::NetworkName(network_name.clone())), - ), - compress_tx_bytes_before: stats_manager.get_counter( - MetricName::CompressionBytesTxBefore, - LabelSet::new().with_label_type(LabelType::NetworkName(network_name.clone())), - ), - compress_tx_bytes_after: stats_manager.get_counter( - MetricName::CompressionBytesTxAfter, - LabelSet::new().with_label_type(LabelType::NetworkName(network_name.clone())), - ), - }; - let traffic_rx_metrics = Arc::new(LogicalTrafficMetrics::new( - stats_manager.clone(), - network_name, - MetricName::TrafficBytesRx, - MetricName::TrafficPacketsRx, - MetricName::TrafficBytesRxByInstance, - MetricName::TrafficPacketsRxByInstance, - InstanceLabelKind::From, - )); - let traffic_control_rx_metrics = Arc::new(LogicalTrafficMetrics::new( - stats_manager.clone(), - global_ctx.get_network_name(), - MetricName::TrafficControlBytesRx, - MetricName::TrafficControlPacketsRx, - MetricName::TrafficControlBytesRxByInstance, - MetricName::TrafficControlPacketsRxByInstance, - InstanceLabelKind::From, - )); - let route_algo_inst_for_metrics = route_algo_inst.clone(); - let traffic_metrics = Arc::new(TrafficMetricRecorder::new( - my_peer_id, - traffic_tx_metrics, - traffic_control_tx_metrics, - traffic_rx_metrics, - traffic_control_rx_metrics, - move |peer_id| { - let route_algo_inst = route_algo_inst_for_metrics.clone(); - async move { - match &route_algo_inst { - RouteAlgoInst::Ospf(route) => route - .get_peer_info(peer_id) - .await - .as_ref() - .and_then(route_peer_info_instance_id), - RouteAlgoInst::None => None, - } - } - }, - )); - - PeerManager { - my_peer_id, - - global_ctx, - nic_channel, - - tasks: Mutex::new(JoinSet::new()), - - packet_recv: Arc::new(Mutex::new(Some(packet_recv))), - - peers, - - peer_rpc_mgr, - peer_rpc_tspt: rpc_tspt, - - peer_packet_process_pipeline: Arc::new(RwLock::new(Vec::new())), - nic_packet_process_pipeline: Arc::new(RwLock::new(Vec::new())), - - route_algo_inst, - - foreign_network_manager, - foreign_network_client, - relay_peer_map, - - encryptor, - data_compress_algo, - - exit_nodes: RwLock::new(exit_nodes), - - reserved_my_peer_id_map: DashMap::new(), - recent_have_traffic: Arc::new(DashMap::new()), - p2p_demand_notify: Arc::new(ExternalTaskSignal::new()), - - allow_loopback_tunnel: AtomicBool::new(true), - - self_tx_counters, - traffic_metrics, - - peer_session_store, - is_secure_mode_enabled, - } - } - - pub fn set_allow_loopback_tunnel(&self, allow_loopback_tunnel: bool) { - self.allow_loopback_tunnel - .store(allow_loopback_tunnel, std::sync::atomic::Ordering::Relaxed); - } - - pub fn mark_recent_traffic(&self, dst_peer_id: PeerId) { - if dst_peer_id == self.my_peer_id { - return; - } - - let flags = self.global_ctx.flags_arc(); - if flags.disable_p2p || !flags.lazy_p2p || self.has_directly_connected_conn(dst_peer_id) { - return; - } - - let now = Instant::now(); - if let Some(mut last_seen) = self.recent_have_traffic.get_mut(&dst_peer_id) { - let should_notify = - now.saturating_duration_since(*last_seen) > Self::RECENT_HAVE_TRAFFIC_TTL; - *last_seen = now; - if !should_notify { - return; - } - } else { - self.recent_have_traffic.insert(dst_peer_id, now); - } - self.p2p_demand_notify.notify(); - } - - pub fn has_recent_traffic(&self, peer_id: PeerId, now: Instant) -> bool { - if self.has_directly_connected_conn(peer_id) { - return false; - } - - self.recent_have_traffic - .get(&peer_id) - .map(|last_seen| { - now.saturating_duration_since(*last_seen) <= Self::RECENT_HAVE_TRAFFIC_TTL - }) - .unwrap_or(false) - } - - pub fn clear_recent_traffic(&self, peer_id: PeerId) { - self.recent_have_traffic.remove(&peer_id); - } - - pub fn p2p_demand_notify(&self) -> Arc { - self.p2p_demand_notify.clone() - } - - fn gc_recent_traffic(&self) { - Self::gc_recent_traffic_entries(&self.recent_have_traffic, Instant::now(), |peer_id| { - self.has_directly_connected_conn(peer_id) - }); - } - - async fn close_untrusted_credential_peers(peer_map: &Arc, global_ctx: &ArcGlobalCtx) { - let network_name = global_ctx.get_network_name(); - for peer_id in peer_map.list_peers() { - if !matches!( - peer_map.get_peer_identity_type(peer_id), - Some(PeerIdentityType::Credential) - ) { - continue; - } - let Some(peer) = peer_map.get_peer_by_id(peer_id) else { - continue; - }; - let Some(pubkey) = peer.get_peer_public_key() else { - continue; - }; - - if global_ctx.is_pubkey_trusted(&pubkey, &network_name) { - continue; - } - - tracing::warn!(?peer_id, "closing untrusted credential peer"); - if let Err(e) = peer_map.close_peer(peer_id).await { - tracing::warn!(?e, ?peer_id, "failed to close untrusted credential peer"); - } - } - } - - fn build_foreign_network_manager_accessor( - peer_map: &Arc, - ) -> Box { - struct T { - peer_map: Weak, - } - - #[async_trait::async_trait] - impl GlobalForeignNetworkAccessor for T { - async fn list_global_foreign_peer( - &self, - network_identity: &NetworkIdentity, - ) -> Vec { - let Some(peer_map) = self.peer_map.upgrade() else { - return vec![]; - }; - - peer_map - .list_peers_own_foreign_network(network_identity) - .await - } - } - - Box::new(T { - peer_map: Arc::downgrade(peer_map), - }) - } - - async fn add_new_peer_conn(&self, peer_conn: PeerConn) -> Result<(), Error> { - let my_identity = self.global_ctx.get_network_identity(); - let peer_identity = peer_conn.get_network_identity(); - let conn_info = peer_conn.get_conn_info(); - let local_secure_mode = self - .global_ctx - .config - .get_secure_mode() - .as_ref() - .map(|cfg| cfg.enabled) - .unwrap_or(false); - let peer_secure_mode = !conn_info.noise_remote_static_pubkey.is_empty(); - - if local_secure_mode != peer_secure_mode { - return Err(Error::SecretKeyError( - "same-network peers must use the same secure mode".to_string(), - )); - } - - // For credential nodes, network_secret_digest is either None or all-zeros - // (all-zeros when received over the wire via handshake). - // In this case, only compare network_name. - let my_digest_empty = my_identity - .network_secret_digest - .as_ref() - .is_none_or(|d| d.iter().all(|b| *b == 0)); - let peer_digest_empty = peer_identity - .network_secret_digest - .as_ref() - .is_none_or(|d| d.iter().all(|b| *b == 0)); - - let identity_ok = if my_digest_empty || peer_digest_empty { - // Credential node: only check network_name - my_identity.network_name == peer_identity.network_name - } else { - my_identity == peer_identity - }; - - if !identity_ok { - return Err(Error::SecretKeyError( - "network identity not match".to_string(), - )); - } - let peer_id = peer_conn.get_peer_id(); - self.peers.add_new_peer_conn(peer_conn).await?; - self.clear_recent_traffic(peer_id); - Ok(()) - } - - pub async fn add_client_tunnel( - &self, - tunnel: Box, - is_directly_connected: bool, - ) -> Result<(PeerId, PeerConnId), Error> { - self.add_client_tunnel_with_peer_id_hint(tunnel, is_directly_connected, None) - .await - } - - pub async fn add_client_tunnel_with_peer_id_hint( - &self, - tunnel: Box, - is_directly_connected: bool, - peer_id_hint: Option, - ) -> Result<(PeerId, PeerConnId), Error> { - let mut peer = PeerConn::new_with_peer_id_hint( - self.my_peer_id, - self.global_ctx.clone(), - tunnel, - peer_id_hint, - self.peer_session_store.clone(), - ); - peer.set_is_hole_punched(!is_directly_connected); - peer.do_handshake_as_client().await?; - let conn_id = peer.get_conn_id(); - let peer_id = peer.get_peer_id(); - if peer.get_network_identity().network_name - == self.global_ctx.get_network_identity().network_name - { - self.add_new_peer_conn(peer).await?; - } else { - self.foreign_network_client.add_new_peer_conn(peer).await?; - } - Ok((peer_id, conn_id)) - } - - pub fn has_directly_connected_conn(&self, peer_id: PeerId) -> bool { - if let Some(peer) = self.peers.get_peer_by_id(peer_id) { - peer.has_directly_connected_conn() - } else { - self.foreign_network_client.get_peer_map().has_peer(peer_id) - } - } - - #[tracing::instrument] - pub async fn try_direct_connect(&self, connector: C) -> Result<(PeerId, PeerConnId), Error> - where - C: TunnelConnector + Debug, - { - self.try_direct_connect_with_peer_id_hint(connector, None) - .await - } - - #[tracing::instrument] - pub async fn try_direct_connect_with_peer_id_hint( - &self, - connector: C, - peer_id_hint: Option, - ) -> Result<(PeerId, PeerConnId), Error> - where - C: TunnelConnector + Debug, - { - let t = self.connect_tunnel(connector).await?; - self.add_client_tunnel_with_peer_id_hint(t, true, peer_id_hint) - .await - } - - pub(crate) async fn connect_tunnel(&self, mut connector: C) -> Result, Error> - where - C: TunnelConnector + Debug, - { - let ns = self.global_ctx.net_ns.clone(); - Ok(ns - .run_async(|| async move { connector.connect().await }) - .await?) - } - - // avoid loop back to virtual network - fn check_remote_addr_not_from_virtual_network( - &self, - tunnel: &dyn Tunnel, - ) -> Result<(), anyhow::Error> { - tracing::info!("check remote addr not from virtual network"); - let Some(tunnel_info) = tunnel.info() else { - anyhow::bail!("tunnel info is not set"); - }; - let Some(src) = tunnel_info.remote_addr.map(url::Url::from) else { - anyhow::bail!("tunnel info remote addr is not set"); - }; - if src.scheme() == "ring" { - return Ok(()); - } - let Ok(Some(addr)) = src.socket_addrs(|| Some(1)).map(|x| x.first().cloned()) else { - // if the tunnel is not rely on ip address, skip check - return Ok(()); - }; - - // if no-tun is enabled, the src ip of packet in virtual network is converted to loopback address - // we already filter out the connection in tcp/quic/kcp proxy so no need check here. - if addr.ip().is_loopback() { - // allow other loopback address, good for conn from cdn/l4 connection - return Ok(()); - } - - if self.global_ctx.is_ip_in_same_network(&addr.ip()) { - anyhow::bail!( - "tunnel src {} is from the same network (ignore this error please)", - addr - ); - } - - Ok(()) - } - - fn release_reserved_peer_id(&self, network_name: &str) { - self.reserved_my_peer_id_map.remove(network_name); - shrink_dashmap(&self.reserved_my_peer_id_map, None); - } - - #[tracing::instrument(ret)] - pub async fn add_tunnel_as_server( - &self, - tunnel: Box, - is_directly_connected: bool, - ) -> Result<(), Error> { - tracing::info!("add tunnel as server start"); - self.check_remote_addr_not_from_virtual_network(&tunnel)?; - - let mut conn = PeerConn::new( - self.my_peer_id, - self.global_ctx.clone(), - tunnel, - self.peer_session_store.clone(), - ); - let mut reserved_peer_id_network_name = None; - let handshake_ret = conn.do_handshake_as_server_ext(|peer, network_name:&str| { - if network_name - == self.global_ctx.get_network_identity().network_name - { - return Ok(()); - } - - let mut peer_id = self - .foreign_network_manager - .get_network_peer_id(network_name); - if peer_id.is_none() { - reserved_peer_id_network_name = Some(network_name.to_string()); - peer_id = Some(*self.reserved_my_peer_id_map.entry(network_name.to_string()).or_insert_with(|| { - rand::random::() - }).value()); - } - peer.set_peer_id(peer_id.unwrap()); - - tracing::info!( - ?peer_id, - ?network_name, - "handshake as server with foreign network, new peer id: {}, peer id in foreign manager: {:?}", - peer.get_my_peer_id(), peer_id - ); - - Ok(()) - }) - .await; - - if let Err(err) = handshake_ret { - if let Some(network_name) = reserved_peer_id_network_name { - self.release_reserved_peer_id(&network_name); - } - return Err(err); - } - - let peer_identity = conn.get_network_identity(); - let peer_network_name = peer_identity.network_name.clone(); - let my_identity = self.global_ctx.get_network_identity(); - let is_local_network = peer_network_name == my_identity.network_name; - let trusted_foreign_credential = - matches!(conn.get_peer_identity_type(), PeerIdentityType::Credential) - && self - .foreign_network_manager - .is_existing_credential_pubkey_trusted( - &peer_network_name, - &conn.get_conn_info().noise_remote_static_pubkey, - ); - let foreign_network_allowed = - conn.matches_local_network_secret() || trusted_foreign_credential; - - if !is_local_network && self.global_ctx.get_flags().private_mode && !foreign_network_allowed - { - self.release_reserved_peer_id(&peer_network_name); - return Err(Error::SecretKeyError( - "private mode is turned on, foreign network secret mismatch".to_string(), - )); - } - - conn.set_is_hole_punched(!is_directly_connected); - - let add_peer_ret = if is_local_network { - self.add_new_peer_conn(conn).await - } else { - self.foreign_network_manager.add_peer_conn(conn).await - }; - - if let Err(err) = add_peer_ret { - self.release_reserved_peer_id(&peer_network_name); - return Err(err); - } - - self.release_reserved_peer_id(&peer_network_name); - - tracing::info!("add tunnel as server done"); - Ok(()) - } - - async fn try_handle_foreign_network_packet( - mut packet: ZCPacket, - my_peer_id: PeerId, - peer_map: &PeerMap, - foreign_network_mgr: &ForeignNetworkManager, - disable_relay_data: bool, - ) -> Result<(), ZCPacket> { - let pm_header = packet.peer_manager_header().unwrap(); - if pm_header.packet_type != PacketType::ForeignNetworkPacket as u8 { - return Err(packet); - } - - let from_peer_id = pm_header.from_peer_id.get(); - let to_peer_id = pm_header.to_peer_id.get(); - - if disable_relay_data && Self::is_relay_data_zc_packet(&packet) { - tracing::debug!( - ?from_peer_id, - ?to_peer_id, - inner_packet_type = ?packet.foreign_network_inner_packet_type(), - "drop foreign network relay data while relay data is disabled" - ); - return Ok(()); - } - - let foreign_hdr = packet.foreign_network_hdr().unwrap(); - let foreign_network_name = foreign_hdr.get_network_name(packet.payload()); - let foreign_peer_id = foreign_hdr.get_dst_peer_id(); - - let foreign_network_my_peer_id = - foreign_network_mgr.get_network_peer_id(&foreign_network_name); - - let buf_len = packet.buf_len(); - let stats_manager = peer_map.get_global_ctx().stats_manager().clone(); - let label_set = - LabelSet::new().with_label_type(LabelType::NetworkName(foreign_network_name.clone())); - let add_counter = move |bytes_metric, packets_metric| { - stats_manager - .get_counter(bytes_metric, label_set.clone()) - .add(buf_len as u64); - stats_manager.get_counter(packets_metric, label_set).inc(); - }; - - // NOTICE: the to peer id is modified by the src from foreign network my peer id to the origin my peer id - if to_peer_id == my_peer_id { - // packet sent from other peer to me, extract the inner packet and forward it - add_counter( - MetricName::TrafficBytesForeignForwardRx, - MetricName::TrafficPacketsForeignForwardRx, - ); - if let Err(e) = foreign_network_mgr - .forward_foreign_network_packet( - &foreign_network_name, - foreign_peer_id, - packet.foreign_network_packet(), - ) - .await - { - tracing::debug!( - ?e, - ?foreign_network_name, - ?foreign_peer_id, - "foreign network mgr send_msg_to_peer failed" - ); - } - Ok(()) - } else if Some(from_peer_id) == foreign_network_my_peer_id { - // to_peer_id is my peer id for the foreign network, need to convert to the origin my_peer_id of dst - let Some(to_peer_id) = peer_map - .get_origin_my_peer_id(&foreign_network_name, to_peer_id) - .await - else { - tracing::debug!( - ?foreign_network_name, - ?to_peer_id, - "cannot find origin my peer id for foreign network." - ); - return Err(packet); - }; - - add_counter( - MetricName::TrafficBytesForeignForwardTx, - MetricName::TrafficPacketsForeignForwardTx, - ); - - // modify the to_peer id from foreign network my peer id to the origin my peer id - packet - .mut_peer_manager_header() - .unwrap() - .to_peer_id - .set(to_peer_id); - - // packet is generated from foreign network mgr and should be forward to other peer - if let Err(e) = peer_map - .send_msg(packet, to_peer_id, NextHopPolicy::LeastHop) - .await - { - tracing::debug!( - ?e, - ?to_peer_id, - "send_msg_directly failed when forward local generated foreign network packet" - ); - } - Ok(()) - } else { - // target is not me, forward it. try get origin peer id - add_counter( - MetricName::TrafficBytesForeignForwardForwarded, - MetricName::TrafficPacketsForeignForwardForwarded, - ); - Err(packet) - } - } - - fn is_relay_data_packet(packet_type: u8) -> bool { - is_relay_data_packet_type(packet_type) - } - - fn is_relay_data_zc_packet(packet: &ZCPacket) -> bool { - let Some(hdr) = packet.peer_manager_header() else { - return false; - }; - - if hdr.packet_type == PacketType::ForeignNetworkPacket as u8 { - let inner_packet_type = packet.foreign_network_inner_packet_type(); - if inner_packet_type.is_none() { - tracing::warn!( - ?hdr, - "foreign network packet has unparseable inner peer manager header" - ); - } - return inner_packet_type.is_none_or(Self::is_relay_data_packet); - } - - Self::is_relay_data_packet(hdr.packet_type) - } - - async fn start_peer_recv(&self) { - let mut recv = self.packet_recv.lock().await.take().unwrap(); - let my_peer_id = self.my_peer_id; - let peers = self.peers.clone(); - let pipe_line = self.peer_packet_process_pipeline.clone(); - let foreign_client = self.foreign_network_client.clone(); - let relay_peer_map = self.relay_peer_map.clone(); - let foreign_mgr = self.foreign_network_manager.clone(); - let encryptor = self.encryptor.clone(); - let compress_algo = self.data_compress_algo; - let acl_filter = self.global_ctx.get_acl_filter().clone(); - let global_ctx = self.global_ctx.clone(); - let secure_mode_enabled = self.is_secure_mode_enabled; - let stats_mgr = self.global_ctx.stats_manager().clone(); - let route = self.get_route(); - let is_credential_node = self - .global_ctx - .get_network_identity() - .network_secret - .is_none() - && secure_mode_enabled; - - let label_set = - LabelSet::new().with_label_type(LabelType::NetworkName(global_ctx.get_network_name())); - - let self_tx_bytes = self.self_tx_counters.self_tx_bytes.clone(); - let self_tx_packets = self.self_tx_counters.self_tx_packets.clone(); - let self_rx_bytes = - stats_mgr.get_counter(MetricName::TrafficBytesSelfRx, label_set.clone()); - let self_rx_packets = - stats_mgr.get_counter(MetricName::TrafficPacketsSelfRx, label_set.clone()); - let forward_data_tx_bytes = - stats_mgr.get_counter(MetricName::TrafficBytesForwarded, label_set.clone()); - let forward_data_tx_packets = - stats_mgr.get_counter(MetricName::TrafficPacketsForwarded, label_set.clone()); - let forward_control_tx_bytes = - stats_mgr.get_counter(MetricName::TrafficControlBytesForwarded, label_set.clone()); - let forward_control_tx_packets = stats_mgr.get_counter( - MetricName::TrafficControlPacketsForwarded, - label_set.clone(), - ); - - let compress_tx_bytes_before = self.self_tx_counters.compress_tx_bytes_before.clone(); - let compress_tx_bytes_after = self.self_tx_counters.compress_tx_bytes_after.clone(); - let compress_rx_bytes_before = - stats_mgr.get_counter(MetricName::CompressionBytesRxBefore, label_set.clone()); - let compress_rx_bytes_after = - stats_mgr.get_counter(MetricName::CompressionBytesRxAfter, label_set.clone()); - let traffic_metrics = self.traffic_metrics.clone(); - - self.tasks.lock().await.spawn(async move { - tracing::trace!("start_peer_recv"); - while let Ok(ret) = recv_packet_from_chan(&mut recv).await { - let disable_relay_data = global_ctx.flags_arc().disable_relay_data; - let Err(mut ret) = Self::try_handle_foreign_network_packet( - ret, - my_peer_id, - &peers, - &foreign_mgr, - disable_relay_data, - ) - .await - else { - continue; - }; - - let buf_len = ret.buf_len(); - let is_relay_data_packet = Self::is_relay_data_zc_packet(&ret); - let Some(hdr) = ret.mut_peer_manager_header() else { - tracing::warn!(?ret, "invalid packet, skip"); - continue; - }; - - tracing::trace!(?hdr, "peer recv a packet..."); - let from_peer_id = hdr.from_peer_id.get(); - let to_peer_id = hdr.to_peer_id.get(); - let packet_type = hdr.packet_type; - let is_encrypted = hdr.is_encrypted(); - if to_peer_id != my_peer_id { - if disable_relay_data && is_relay_data_packet { - tracing::debug!( - ?from_peer_id, - ?to_peer_id, - packet_type, - "drop forwarded relay data while relay data is disabled" - ); - continue; - } - - if hdr.forward_counter > 7 { - tracing::warn!(?hdr, "forward counter exceed, drop packet"); - continue; - } - - // Step 10b: credential nodes don't forward handshake packets - if is_credential_node - && (packet_type == PacketType::HandShake as u8 - || packet_type == PacketType::NoiseHandshakeMsg1 as u8 - || packet_type == PacketType::NoiseHandshakeMsg2 as u8 - || packet_type == PacketType::NoiseHandshakeMsg3 as u8) - { - tracing::debug!("credential node dropping forwarded handshake packet"); - continue; - } - - if hdr.forward_counter > 2 && hdr.is_latency_first() { - tracing::trace!(?hdr, "set_latency_first false because too many hop"); - hdr.set_latency_first(false); - } - - hdr.forward_counter += 1; - - if from_peer_id == my_peer_id { - compress_tx_bytes_before.add(buf_len as u64); - - if packet_type == PacketType::Data as u8 - || packet_type == PacketType::KcpSrc as u8 - || packet_type == PacketType::KcpDst as u8 - { - let _ = Self::try_compress_and_encrypt( - compress_algo, - &encryptor, - &mut ret, - secure_mode_enabled, - ) - .await; - } - - compress_tx_bytes_after.add(ret.buf_len() as u64); - self_tx_bytes.add(ret.buf_len() as u64); - self_tx_packets.inc(); - } else { - match traffic_kind(packet_type) { - TrafficKind::Data => { - forward_data_tx_bytes.add(buf_len as u64); - forward_data_tx_packets.inc(); - } - TrafficKind::Control => { - forward_control_tx_bytes.add(buf_len as u64); - forward_control_tx_packets.inc(); - } - } - } - - tracing::trace!(?to_peer_id, ?my_peer_id, "need forward"); - let tx_metrics = if from_peer_id == my_peer_id { - Some(&traffic_metrics) - } else { - None - }; - let ret = Self::send_msg_internal( - &peers, - &foreign_client, - &relay_peer_map, - tx_metrics, - ret, - to_peer_id, - ) - .await; - if ret.is_err() { - tracing::error!(?ret, ?to_peer_id, ?from_peer_id, "forward packet error"); - } - } else { - if packet_type == PacketType::RelayHandshake as u8 - || packet_type == PacketType::RelayHandshakeAck as u8 - { - let _ = relay_peer_map.handle_handshake_packet(ret).await; - continue; - } - if !secure_mode_enabled { - if let Err(e) = encryptor.decrypt(&mut ret) { - tracing::error!(?e, "decrypt failed"); - continue; - } - } else if is_encrypted { - match relay_peer_map.decrypt_if_needed(&mut ret).await { - Ok(true) => {} - Ok(false) => { - tracing::error!("secure session not found"); - continue; - } - Err(e) => { - tracing::error!(?e, "secure decrypt failed"); - continue; - } - } - } - - self_rx_bytes.add(buf_len as u64); - self_rx_packets.inc(); - traffic_metrics - .record_rx(from_peer_id, packet_type, buf_len as u64) - .await; - compress_rx_bytes_before.add(buf_len as u64); - - let compressor = DefaultCompressor {}; - if let Err(e) = compressor.decompress(&mut ret).await { - tracing::error!(?e, "decompress failed"); - continue; - } - - compress_rx_bytes_after.add(ret.buf_len() as u64); - - if !acl_filter.process_packet_with_acl( - &ret, - true, - global_ctx.get_ipv4().map(|x| x.address()), - |dst| global_ctx.is_ip_local_ipv6(&dst), - &route, - ) { - continue; - } - - let mut processed = false; - let mut zc_packet = Some(ret); - tracing::trace!(?zc_packet, "try_process_packet_from_peer"); - for pipeline in pipe_line.read().await.iter().rev() { - zc_packet = pipeline - .try_process_packet_from_peer(zc_packet.unwrap()) - .await; - if zc_packet.is_none() { - processed = true; - break; - } - } - if !processed { - tracing::error!(?zc_packet, "unhandled packet"); - } - } - } - panic!("done_peer_recv"); - }); - } - - pub async fn add_packet_process_pipeline(&self, pipeline: BoxPeerPacketFilter) { - // newest pipeline will be executed first - self.peer_packet_process_pipeline - .write() - .await - .push(pipeline); - } - - pub async fn add_nic_packet_process_pipeline(&self, pipeline: BoxNicPacketFilter) { - // newest pipeline will be executed first - self.nic_packet_process_pipeline - .write() - .await - .push(pipeline); - } - - async fn init_packet_process_pipeline(&self) { - // for tun/tap ip/eth packet. - struct NicPacketProcessor { - nic_channel: PacketRecvChan, - } - #[async_trait::async_trait] - impl PeerPacketFilter for NicPacketProcessor { - async fn try_process_packet_from_peer(&self, packet: ZCPacket) -> Option { - let hdr = packet.peer_manager_header().unwrap(); - if hdr.packet_type == PacketType::Data as u8 && !hdr.is_not_send_to_tun() { - if hdr.is_encrypted() || hdr.is_compressed() { - tracing::warn!( - from_peer_id = hdr.from_peer_id.get(), - to_peer_id = hdr.to_peer_id.get(), - encrypted = hdr.is_encrypted(), - compressed = hdr.is_compressed(), - "dropping packet before nic because it is not fully decoded" - ); - return None; - } - tracing::trace!(?packet, "send packet to nic channel"); - // TODO: use a function to get the body ref directly for zero copy - let _ = self.nic_channel.send(packet).await; - None - } else { - Some(packet) - } - } - } - self.add_packet_process_pipeline(Box::new(NicPacketProcessor { - nic_channel: self.nic_channel.clone(), - })) - .await; - - // for peer rpc packet - struct PeerRpcPacketProcessor { - peer_rpc_tspt_sender: UnboundedSender, - } - - #[async_trait::async_trait] - impl PeerPacketFilter for PeerRpcPacketProcessor { - async fn try_process_packet_from_peer(&self, packet: ZCPacket) -> Option { - let hdr = packet.peer_manager_header().unwrap(); - if hdr.packet_type == PacketType::TaRpc as u8 - || hdr.packet_type == PacketType::RpcReq as u8 - || hdr.packet_type == PacketType::RpcResp as u8 - { - self.peer_rpc_tspt_sender.send(packet).unwrap(); - None - } else { - Some(packet) - } - } - } - self.add_packet_process_pipeline(Box::new(PeerRpcPacketProcessor { - peer_rpc_tspt_sender: self.peer_rpc_tspt.peer_rpc_tspt_sender.clone(), - })) - .await; - } - - pub async fn add_route(&self, route: T) - where - T: Route + PeerPacketFilter + Send + Sync + Clone + 'static, - { - // for route - self.add_packet_process_pipeline(Box::new(route.clone())) - .await; - - struct Interface { - my_peer_id: PeerId, - peers: Weak, - foreign_network_client: Weak, - foreign_network_manager: Weak, - } - - #[async_trait] - impl RouteInterface for Interface { - async fn list_peers(&self) -> Vec { - let Some(foreign_client) = self.foreign_network_client.upgrade() else { - return vec![]; - }; - - let Some(peer_map) = self.peers.upgrade() else { - return vec![]; - }; - - let mut peers = foreign_client.list_public_peers().await; - peers.extend(peer_map.list_peers_with_conn().await); - peers - } - - fn my_peer_id(&self) -> PeerId { - self.my_peer_id - } - - async fn close_peer(&self, peer_id: PeerId) { - if let Some(peer_map) = self.peers.upgrade() { - let _ = peer_map.close_peer(peer_id).await; - } - - if let Some(foreign_client) = self.foreign_network_client.upgrade() { - let _ = foreign_client.get_peer_map().close_peer(peer_id).await; - } - } - - async fn get_peer_public_key(&self, peer_id: PeerId) -> Option> { - let peer_map = self.peers.upgrade()?; - peer_map.get_peer_public_key(peer_id) - } - - async fn get_peer_identity_type(&self, peer_id: PeerId) -> Option { - let peer_map = self.peers.upgrade()?; - peer_map.get_peer_identity_type(peer_id) - } - - async fn list_foreign_networks(&self) -> ForeignNetworkRouteInfoMap { - let ret = DashMap::new(); - let Some(foreign_mgr) = self.foreign_network_manager.upgrade() else { - return ret; - }; - - let networks = foreign_mgr.list_foreign_networks().await; - for (network_name, info) in networks.foreign_networks.iter() { - if info.peers.is_empty() { - continue; - } - - let last_update = foreign_mgr - .get_foreign_network_last_update(network_name) - .unwrap_or(SystemTime::now()); - ret.insert( - ForeignNetworkRouteInfoKey { - peer_id: self.my_peer_id, - network_name: network_name.clone(), - }, - ForeignNetworkRouteInfoEntry { - foreign_peer_ids: info.peers.iter().map(|x| x.peer_id).collect(), - last_update: Some(last_update.into()), - version: 0, - network_secret_digest: info.network_secret_digest.clone(), - my_peer_id_for_this_network: info.my_peer_id_for_this_network, - }, - ); - } - ret - } - } - - let my_peer_id = self.my_peer_id; - let _route_id = route - .open(Box::new(Interface { - my_peer_id, - peers: Arc::downgrade(&self.peers), - foreign_network_client: Arc::downgrade(&self.foreign_network_client), - foreign_network_manager: Arc::downgrade(&self.foreign_network_manager), - })) - .await - .unwrap(); - - let arc_route: ArcRoute = Arc::new(Box::new(route)); - self.peers.add_route(arc_route).await; - } - - pub fn get_route(&self) -> Box { - match &self.route_algo_inst { - RouteAlgoInst::Ospf(route) => Box::new(route.clone()), - RouteAlgoInst::None => Box::new(MockRoute {}), - } - } - - pub async fn list_routes(&self) -> Vec { - self.get_route().list_routes().await - } - - pub async fn get_route_peer_info_last_update_time(&self) -> Instant { - self.get_route().get_peer_info_last_update_time().await - } - - pub async fn list_proxy_cidrs(&self) -> BTreeSet { - self.get_route().list_proxy_cidrs().await - } - - pub async fn list_proxy_cidrs_v6(&self) -> BTreeSet { - self.get_route().list_proxy_cidrs_v6().await - } - - pub async fn list_public_ipv6_routes(&self) -> BTreeSet { - self.get_route().list_public_ipv6_routes().await - } - - pub async fn get_my_public_ipv6_addr(&self) -> Option { - self.get_route().get_my_public_ipv6_addr().await - } - - pub async fn get_local_public_ipv6_info(&self) -> instance::ListPublicIpv6InfoResponse { - self.get_route().get_local_public_ipv6_info().await - } - - pub async fn dump_route(&self) -> String { - self.get_route().dump().await - } - - pub async fn list_global_foreign_network(&self) -> ListGlobalForeignNetworkResponse { - let mut resp = ListGlobalForeignNetworkResponse::default(); - let ret = self.get_route().list_foreign_network_info().await; - for info in ret.infos.iter() { - let entry = resp - .foreign_networks - .entry(info.key.as_ref().unwrap().peer_id) - .or_insert_with(Default::default); - let Some(route_info) = info.value.as_ref() else { - continue; - }; - - let f = OneForeignNetwork { - network_name: info.key.as_ref().unwrap().network_name.clone(), - peer_ids: route_info.foreign_peer_ids.clone(), - last_updated: serde_json::to_string(&route_info.last_update.unwrap()).unwrap(), - version: route_info.version, - }; - - entry.foreign_networks.push(f); - } - - resp - } - - pub async fn get_foreign_network_summary(&self) -> RouteForeignNetworkSummary { - self.get_route().get_foreign_network_summary().await - } - - async fn run_nic_packet_process_pipeline(&self, data: &mut ZCPacket) -> bool { - // Enforce ACL for outbound (NIC-originated) packets. If ACL denies, stop processing. - if !self.global_ctx.get_acl_filter().process_packet_with_acl( - data, - false, - None, - |_| false, - &self.get_route(), - ) { - return false; - } - - for pipeline in self.nic_packet_process_pipeline.read().await.iter().rev() { - let _ = pipeline.try_process_packet_from_nic(data).await; - } - - true - } - - pub async fn remove_nic_packet_process_pipeline(&self, id: String) -> Result<(), Error> { - let mut pipelines = self.nic_packet_process_pipeline.write().await; - if let Some(pos) = pipelines.iter().position(|x| x.id() == id) { - pipelines.remove(pos); - Ok(()) - } else { - Err(Error::NotFound) - } - } - - fn get_next_hop_policy(is_first_latency: bool) -> NextHopPolicy { - if is_first_latency { - NextHopPolicy::LeastCost - } else { - NextHopPolicy::LeastHop - } - } - - fn check_p2p_only_before_send(&self, dst_peer_id: PeerId) -> Result<(), Error> { - if self.global_ctx.p2p_only() && !self.peers.has_peer(dst_peer_id) { - return Err(Error::RouteError(None)); - } - Ok(()) - } - - pub async fn send_msg_for_proxy( - &self, - mut msg: ZCPacket, - dst_peer_id: PeerId, - ) -> Result<(), Error> { - self.mark_recent_traffic(dst_peer_id); - self.check_p2p_only_before_send(dst_peer_id)?; - - self.self_tx_counters - .compress_tx_bytes_before - .add(msg.buf_len() as u64); - - Self::try_compress_and_encrypt( - self.data_compress_algo, - &self.encryptor, - &mut msg, - self.is_secure_mode_enabled, - ) - .await?; - - self.self_tx_counters - .compress_tx_bytes_after - .add(msg.buf_len() as u64); - - let msg_len = msg.buf_len() as u64; - let result = Self::send_msg_internal( - &self.peers, - &self.foreign_network_client, - &self.relay_peer_map, - Some(&self.traffic_metrics), - msg, - dst_peer_id, - ) - .await; - if result.is_ok() { - self.self_tx_counters.self_tx_bytes.add(msg_len); - self.self_tx_counters.self_tx_packets.inc(); - } - result - } - - async fn send_msg_internal( - peers: &Arc, - foreign_network_client: &Arc, - relay_peer_map: &Arc, - direct_tx_metrics: Option<&Arc>, - msg: ZCPacket, - dst_peer_id: PeerId, - ) -> Result<(), Error> { - let policy = - Self::get_next_hop_policy(msg.peer_manager_header().unwrap().is_latency_first()); - let is_latency_first = msg.peer_manager_header().unwrap().is_latency_first(); - let packet_type = msg.peer_manager_header().unwrap().packet_type; - let msg_len = msg.buf_len() as u64; - let latency_first_gateway = if is_latency_first { - peers - .get_gateway_peer_id(dst_peer_id, policy.clone()) - .await - .filter(|gateway| *gateway != dst_peer_id) - } else { - None - }; - let send_result = if let Some(gateway) = latency_first_gateway - && (peers.has_peer(gateway) || foreign_network_client.has_next_hop(gateway)) - { - relay_peer_map.send_msg(msg, dst_peer_id, policy).await - } else if peers.has_peer(dst_peer_id) { - peers.send_msg_directly(msg, dst_peer_id).await - } else if foreign_network_client.has_next_hop(dst_peer_id) { - foreign_network_client.send_msg(msg, dst_peer_id).await - } else if let Some(gateway) = peers.get_gateway_peer_id(dst_peer_id, policy.clone()).await { - if peers.has_peer(gateway) || foreign_network_client.has_next_hop(gateway) { - relay_peer_map.send_msg(msg, dst_peer_id, policy).await - } else { - tracing::warn!( - ?gateway, - ?dst_peer_id, - "cannot send msg to peer through gateway" - ); - Err(Error::RouteError(None)) - } - } else if foreign_network_client.has_next_hop(dst_peer_id) { - // check foreign network again. so in happy path we can avoid extra check - foreign_network_client.send_msg(msg, dst_peer_id).await - } else { - tracing::debug!(?dst_peer_id, "no gateway for peer"); - Err(Error::RouteError(None)) - }; - - if send_result.is_ok() - && let Some(metrics) = direct_tx_metrics - { - metrics.record_tx(dst_peer_id, packet_type, msg_len).await; - } - - send_result - } - - pub async fn get_msg_dst_peer(&self, addr: &IpAddr) -> (Vec, bool) { - match addr { - IpAddr::V4(ipv4_addr) => self.get_msg_dst_peer_ipv4(ipv4_addr).await, - IpAddr::V6(ipv6_addr) => self.get_msg_dst_peer_ipv6(ipv6_addr).await, - } - } - - fn is_all_peers_broadcast_ipv4(&self, ipv4_addr: &Ipv4Addr) -> bool { - let network_length = self - .global_ctx - .get_ipv4() - .map(|x| x.network_length()) - .unwrap_or(24); - let ipv4_inet = cidr::Ipv4Inet::new(*ipv4_addr, network_length).unwrap(); - ipv4_addr.is_broadcast() - || ipv4_addr.is_multicast() - || *ipv4_addr == ipv4_inet.last_address() - } - - fn is_all_peers_broadcast_ipv6(&self, ipv6_addr: &Ipv6Addr) -> bool { - let network_length = self - .global_ctx - .get_ipv6() - .map(|x| x.network_length()) - .unwrap_or(64); - let ipv6_inet = cidr::Ipv6Inet::new(*ipv6_addr, network_length).unwrap(); - ipv6_addr.is_multicast() || *ipv6_addr == ipv6_inet.last_address() - } - - fn select_ipv4_broadcast_peers<'a>( - routes: impl IntoIterator, - my_peer_id: PeerId, - ) -> Vec { - routes - .into_iter() - .filter_map(|route| { - (route.peer_id != my_peer_id && route.ipv4_addr.is_some()).then_some(route.peer_id) - }) - .collect() - } - - pub async fn get_msg_dst_peer_ipv4(&self, ipv4_addr: &Ipv4Addr) -> (Vec, bool) { - let mut is_exit_node = false; - let mut dst_peers = vec![]; - if self.is_all_peers_broadcast_ipv4(ipv4_addr) { - dst_peers.extend(Self::select_ipv4_broadcast_peers( - &self.peers.list_route_infos().await, - self.my_peer_id, - )); - } else if let Some(peer_id) = self.peers.get_peer_id_by_ipv4(ipv4_addr).await { - dst_peers.push(peer_id); - } else if !self - .global_ctx - .is_ip_in_same_network(&std::net::IpAddr::V4(*ipv4_addr)) - { - for exit_node in self.exit_nodes.read().await.iter() { - let IpAddr::V4(exit_node) = exit_node else { - continue; - }; - if let Some(peer_id) = self.peers.get_peer_id_by_ipv4(exit_node).await { - dst_peers.push(peer_id); - is_exit_node = true; - break; - } - } - } - #[cfg(target_env = "ohos")] - { - if dst_peers.is_empty() - && !self - .global_ctx - .is_ip_in_same_network(&std::net::IpAddr::V4(*ipv4_addr)) - { - tracing::trace!("no peer id for ipv4: {}, set exit_node for ohos", ipv4_addr); - dst_peers.push(self.my_peer_id.clone()); - is_exit_node = true; - } - } - (dst_peers, is_exit_node) - } - - pub async fn get_msg_dst_peer_ipv6(&self, ipv6_addr: &Ipv6Addr) -> (Vec, bool) { - let mut is_exit_node = false; - let mut dst_peers = vec![]; - if self.is_all_peers_broadcast_ipv6(ipv6_addr) { - dst_peers.extend(self.peers.list_routes().await.iter().map(|x| *x.key())); - } else if let Some(peer_id) = self.peers.get_peer_id_by_ipv6(ipv6_addr).await { - dst_peers.push(peer_id); - } else if !ipv6_addr.is_unicast_link_local() - && let Some(peer_id) = self.get_route().get_public_ipv6_gateway_peer_id().await - { - dst_peers.push(peer_id); - } else if !ipv6_addr.is_unicast_link_local() { - // NOTE: never route link local address to exit node. - for exit_node in self.exit_nodes.read().await.iter() { - let IpAddr::V6(exit_node) = exit_node else { - continue; - }; - if let Some(peer_id) = self.peers.get_peer_id_by_ipv6(exit_node).await { - dst_peers.push(peer_id); - is_exit_node = true; - break; - } - } - } - - (dst_peers, is_exit_node) - } - - pub async fn try_compress_and_encrypt( - compress_algo: CompressorAlgo, - encryptor: &Arc, - msg: &mut ZCPacket, - secure_mode_enabled: bool, - ) -> Result<(), Error> { - let compressor = DefaultCompressor {}; - compressor - .compress(msg, compress_algo) - .await - .with_context(|| "compress failed")?; - if !secure_mode_enabled { - encryptor.encrypt(msg).with_context(|| "encrypt failed")?; - } - Ok(()) - } - - pub async fn send_msg_by_ip( - &self, - mut msg: ZCPacket, - ip_addr: IpAddr, - not_send_to_self: bool, - ) -> Result<(), Error> { - tracing::trace!( - "do send_msg in peer manager, msg: {:?}, ip_addr: {}", - msg, - ip_addr - ); - - msg.fill_peer_manager_hdr( - self.my_peer_id, - 0, - tunnel::packet_def::PacketType::Data as u8, - ); - if !self.run_nic_packet_process_pipeline(&mut msg).await { - return Ok(()); - } - let cur_to_peer_id = msg.peer_manager_header().unwrap().to_peer_id.into(); - if cur_to_peer_id != 0 { - self.mark_recent_traffic(cur_to_peer_id); - return Self::send_msg_internal( - &self.peers, - &self.foreign_network_client, - &self.relay_peer_map, - Some(&self.traffic_metrics), - msg, - cur_to_peer_id, - ) - .await; - } - - let (dst_peers, is_exit_node) = match ip_addr { - IpAddr::V4(ipv4_addr) => self.get_msg_dst_peer_ipv4(&ipv4_addr).await, - IpAddr::V6(ipv6_addr) => self.get_msg_dst_peer_ipv6(&ipv6_addr).await, - }; - - if dst_peers.is_empty() { - tracing::info!("no peer id for ip: {}", ip_addr); - return Ok(()); - } - - self.self_tx_counters - .compress_tx_bytes_before - .add(msg.buf_len() as u64); - - Self::try_compress_and_encrypt( - self.data_compress_algo, - &self.encryptor, - &mut msg, - self.is_secure_mode_enabled, - ) - .await?; - - self.self_tx_counters - .compress_tx_bytes_after - .add(msg.buf_len() as u64); - - let is_latency_first = self.global_ctx.latency_first(); - msg.mut_peer_manager_header() - .unwrap() - .set_latency_first(is_latency_first) - .set_exit_node(is_exit_node); - - let mut errs: Vec = vec![]; - let mut msg = Some(msg); - let total_dst_peers = dst_peers.len(); - let should_mark_recent_traffic = - Self::should_mark_recent_traffic_for_fanout(total_dst_peers); - for (i, peer_id) in dst_peers.iter().enumerate() { - if should_mark_recent_traffic { - self.mark_recent_traffic(*peer_id); - } - if let Err(e) = self.check_p2p_only_before_send(*peer_id) { - errs.push(e); - continue; - } - - let mut msg = if i == total_dst_peers - 1 { - msg.take().unwrap() - } else { - msg.clone().unwrap() - }; - - let hdr = msg.mut_peer_manager_header().unwrap(); - hdr.to_peer_id.set(*peer_id); - - #[cfg(not(target_env = "ohos"))] - { - if not_send_to_self - && *peer_id == self.my_peer_id - && !self.global_ctx.is_ip_local_virtual_ip(&ip_addr) - { - // Keep the loop-prevention flags for proxy-induced self-delivery where - // the destination is not this node's own EasyTier-managed IP. - hdr.set_not_send_to_tun(true); - hdr.set_no_proxy(true); - } - } - - self.self_tx_counters - .self_tx_bytes - .add(msg.buf_len() as u64); - self.self_tx_counters.self_tx_packets.inc(); - - if let Err(e) = Self::send_msg_internal( - &self.peers, - &self.foreign_network_client, - &self.relay_peer_map, - Some(&self.traffic_metrics), - msg, - *peer_id, - ) - .await - { - errs.push(e); - } - } - - tracing::trace!(?dst_peers, "do send_msg in peer manager done"); - - if errs.is_empty() { - Ok(()) - } else { - tracing::error!(?errs, "send_msg has error"); - Err(anyhow::anyhow!("send_msg has error: {:?}", errs).into()) - } - } - - async fn run_clean_peer_without_conn_routine(&self) { - let peer_map = self.peers.clone(); - self.tasks.lock().await.spawn(async move { - loop { - peer_map.clean_peer_without_conn().await; - tokio::time::sleep(std::time::Duration::from_secs(3)).await; - } - }); - } - - async fn run_relay_session_gc_routine(&self) { - let relay_peer_map = self.relay_peer_map.clone(); - self.tasks.lock().await.spawn(async move { - loop { - relay_peer_map.evict_idle_sessions(std::time::Duration::from_secs(60)); - tokio::time::sleep(std::time::Duration::from_secs(30)).await; - } - }); - } - - async fn run_recent_traffic_gc_routine(&self) { - let recent_have_traffic = self.recent_have_traffic.clone(); - let peers = self.peers.clone(); - let foreign_network_client = self.foreign_network_client.clone(); - self.tasks.lock().await.spawn(async move { - loop { - PeerManager::gc_recent_traffic_entries( - recent_have_traffic.as_ref(), - Instant::now(), - |peer_id| { - if let Some(peer) = peers.get_peer_by_id(peer_id) { - peer.has_directly_connected_conn() - } else { - foreign_network_client.get_peer_map().has_peer(peer_id) - } - }, - ); - tokio::time::sleep(std::time::Duration::from_secs(30)).await; - } - }); - } - - async fn run_peer_session_gc_routine(&self) { - let peer_session_store = self.peer_session_store.clone(); - self.tasks.lock().await.spawn(async move { - loop { - tokio::time::sleep(std::time::Duration::from_secs(60)).await; - peer_session_store.evict_unused_sessions(); - } - }); - } - - async fn run_credential_gc_routine(&self) { - let global_ctx = self.global_ctx.clone(); - let peer_map = self.peers.clone(); - self.tasks.lock().await.spawn(async move { - loop { - if global_ctx.get_network_identity().network_secret.is_some() { - if global_ctx - .get_credential_manager() - .remove_expired_credentials() - { - global_ctx.issue_event(GlobalCtxEvent::CredentialChanged); - } - - Self::close_untrusted_credential_peers(&peer_map, &global_ctx).await; - } - tokio::time::sleep(std::time::Duration::from_secs(1)).await; - } - }); - } - - async fn run_traffic_metrics_gc_routine(&self) { - let mut event_receiver = self.global_ctx.subscribe(); - let traffic_metrics = self.traffic_metrics.clone(); - self.tasks.lock().await.spawn(async move { - loop { - match event_receiver.recv().await { - Ok(GlobalCtxEvent::PeerRemoved(peer_id)) => { - traffic_metrics.remove_peer(peer_id); - } - Ok(_) => {} - Err(tokio::sync::broadcast::error::RecvError::Lagged(skipped)) => { - tracing::warn!( - skipped, - "traffic metrics GC receiver lagged; clearing peer cache to avoid stale metric attribution" - ); - traffic_metrics.clear_peer_cache(); - event_receiver = event_receiver.resubscribe(); - } - Err(tokio::sync::broadcast::error::RecvError::Closed) => break, - } - } - }); - } - - async fn run_foriegn_network(&self) { - self.peer_rpc_tspt - .foreign_peers - .lock() - .await - .replace(Arc::downgrade(&self.foreign_network_client)); - - self.foreign_network_client.run().await; - } - - pub async fn run(&self) -> Result<(), Error> { - match &self.route_algo_inst { - RouteAlgoInst::Ospf(route) => self.add_route(route.clone()).await, - RouteAlgoInst::None => {} - }; - - self.init_packet_process_pipeline().await; - self.peer_rpc_mgr.run(); - - self.start_peer_recv().await; - self.run_clean_peer_without_conn_routine().await; - self.run_relay_session_gc_routine().await; - self.run_recent_traffic_gc_routine().await; - self.run_peer_session_gc_routine().await; - self.run_credential_gc_routine().await; - self.run_traffic_metrics_gc_routine().await; - - self.run_foriegn_network().await; - - Ok(()) - } - - pub fn get_peer_map(&self) -> Arc { - self.peers.clone() - } - - pub fn get_relay_peer_map(&self) -> Arc { - self.relay_peer_map.clone() - } - - pub fn get_peer_rpc_mgr(&self) -> Arc { - self.peer_rpc_mgr.clone() - } - - pub fn get_peer_session_store(&self) -> Arc { - self.peer_session_store.clone() - } - - pub fn my_node_id(&self) -> uuid::Uuid { - self.global_ctx.get_id() - } - - pub fn my_peer_id(&self) -> PeerId { - self.my_peer_id - } - - pub fn get_global_ctx(&self) -> ArcGlobalCtx { - self.global_ctx.clone() - } - - pub fn get_global_ctx_ref(&self) -> &ArcGlobalCtx { - &self.global_ctx - } - - pub fn get_nic_channel(&self) -> PacketRecvChan { - self.nic_channel.clone() - } - - pub fn get_foreign_network_manager(&self) -> Arc { - self.foreign_network_manager.clone() - } - - pub fn get_foreign_network_client(&self) -> Arc { - self.foreign_network_client.clone() - } - - pub async fn get_my_info(&self) -> instance::NodeInfo { - instance::NodeInfo { - peer_id: self.my_peer_id, - ipv4_addr: self - .global_ctx - .get_ipv4() - .map(|x| x.to_string()) - .unwrap_or_default(), - proxy_cidrs: self - .global_ctx - .config - .get_proxy_cidrs() - .into_iter() - .map(|x| match x.mapped_cidr { - None => x.cidr.to_string(), - Some(mapped) => format!("{}->{}", x.cidr, mapped), - }) - .collect(), - hostname: self.global_ctx.get_hostname(), - stun_info: Some(self.global_ctx.get_stun_info_collector().get_stun_info()), - inst_id: self.global_ctx.get_id().to_string(), - listeners: self - .global_ctx - .get_running_listeners() - .iter() - .map(|x| x.to_string()) - .collect(), - config: self.global_ctx.config.dump(), - version: EASYTIER_VERSION.to_string(), - feature_flag: Some(self.global_ctx.get_feature_flags()), - ip_list: Some(self.global_ctx.get_ip_collector().collect_ip_addrs().await), - public_ipv6_addr: self.get_my_public_ipv6_addr().await.map(Into::into), - ipv6_public_addr_prefix: self - .global_ctx - .get_advertised_ipv6_public_addr_prefix() - .map(|prefix| { - cidr::Ipv6Inet::new(prefix.first_address(), prefix.network_length()) - .unwrap() - .into() - }), - } - } - - pub async fn wait(&self) { - while !self.tasks.lock().await.is_empty() { - tokio::time::sleep(std::time::Duration::from_secs(1)).await; - } - } - - pub async fn clear_resources(&self) { - let mut peer_pipeline = self.peer_packet_process_pipeline.write().await; - peer_pipeline.clear(); - let mut nic_pipeline = self.nic_packet_process_pipeline.write().await; - nic_pipeline.clear(); - - self.peer_rpc_mgr.rpc_server().registry().unregister_all(); - } - - pub async fn close_peer_conn( - &self, - peer_id: PeerId, - conn_id: &PeerConnId, - ) -> Result<(), Error> { - let ret = self.peers.close_peer_conn(peer_id, conn_id).await; - tracing::info!("close_peer_conn in peer map: {:?}", ret); - if ret.is_ok() || !matches!(ret.as_ref().unwrap_err(), Error::NotFound) { - return ret; - } - - let ret = self - .foreign_network_client - .get_peer_map() - .close_peer_conn(peer_id, conn_id) - .await; - tracing::info!("close_peer_conn in foreign network client: {:?}", ret); - if ret.is_ok() || !matches!(ret.as_ref().unwrap_err(), Error::NotFound) { - return ret; - } - - let ret = self - .foreign_network_manager - .close_peer_conn(peer_id, conn_id) - .await; - tracing::info!("close_peer_conn in foreign network manager done: {:?}", ret); - ret - } - - pub async fn check_allow_kcp_to_dst(&self, dst_ip: &IpAddr) -> bool { - let route = self.get_route(); - let Some(dst_peer_id) = route.get_peer_id_by_ip(dst_ip).await else { - return false; - }; - let Some(peer_info) = route.get_peer_info(dst_peer_id).await else { - return false; - }; - - // check dst allow kcp input - if !peer_info.feature_flag.map(|x| x.kcp_input).unwrap_or(false) { - return false; - } - - let next_hop_policy = Self::get_next_hop_policy(self.global_ctx.get_flags().latency_first); - // check relay node allow relay kcp. - let Some(next_hop_id) = route - .get_next_hop_with_policy(dst_peer_id, next_hop_policy) - .await - else { - return false; - }; - - if next_hop_id == dst_peer_id { - // dst p2p, no need to relay - return true; - } - - let Some(next_hop_info) = route.get_peer_info(next_hop_id).await else { - return false; - }; - - // check next hop allow kcp relay - if next_hop_info - .feature_flag - .map(|x| x.no_relay_kcp) - .unwrap_or(false) - { - return false; - } - - true - } - - pub async fn check_allow_quic_to_dst(&self, dst_ip: &IpAddr) -> bool { - let route = self.get_route(); - let Some(dst_peer_id) = route.get_peer_id_by_ip(dst_ip).await else { - return false; - }; - let Some(peer_info) = route.get_peer_info(dst_peer_id).await else { - return false; - }; - - // check dst allow quic input - if !peer_info - .feature_flag - .map(|x| x.quic_input) - .unwrap_or(false) - { - return false; - } - - let next_hop_policy = Self::get_next_hop_policy(self.global_ctx.get_flags().latency_first); - // check relay node allow relay quic. - let Some(next_hop_id) = route - .get_next_hop_with_policy(dst_peer_id, next_hop_policy) - .await - else { - return false; - }; - - if next_hop_id == dst_peer_id { - // dst p2p, no need to relay - return true; - } - - let Some(next_hop_info) = route.get_peer_info(next_hop_id).await else { - return false; - }; - - // check next hop allow quic relay - if next_hop_info - .feature_flag - .map(|x| x.no_relay_quic) - .unwrap_or(false) - { - return false; - } - - true - } - - pub async fn update_exit_nodes(&self) { - let exit_nodes = self.global_ctx.config.get_exit_nodes(); - *self.exit_nodes.write().await = exit_nodes; - } -} - -#[cfg(test)] -mod tests { - use base64::Engine; - use std::{collections::HashMap, fmt::Debug, sync::Arc, time::Duration}; - - use quanta::Instant; - - use crate::{ - common::{ - PeerId, - config::Flags, - global_ctx::{NetworkIdentity, tests::get_mock_global_ctx}, - stats_manager::{LabelSet, LabelType, MetricName}, - }, - connector::{ - create_connector_by_url, direct::PeerManagerForDirectConnector, - udp_hole_punch::tests::create_mock_peer_manager_with_mock_stun, - }, - instance::listeners::create_listener_by_url, - peers::{ - create_packet_recv_chan, - peer_conn::tests::set_secure_mode_cfg, - peer_manager::RouteAlgoType, - peer_rpc::tests::register_service, - route_trait::{NextHopPolicy, RouteCostCalculatorInterface}, - tests::{ - connect_peer_manager, create_mock_peer_manager_with_name, wait_route_appear, - wait_route_appear_with_cost, - }, - }, - proto::{ - common::{CompressionAlgoPb, NatType, SecureModeConfig}, - peer_rpc::SecureAuthLevel, - }, - tunnel::{ - TunnelConnector, TunnelListener, - common::tests::wait_for_condition, - filter::{TunnelWithFilter, tests::DropSendTunnelFilter}, - packet_def::{PacketType, ZCPacket}, - ring::create_ring_tunnel_pair, - }, - }; - - use super::PeerManager; - - async fn create_lazy_peer_manager() -> Arc { - let peer_mgr = create_mock_peer_manager_with_mock_stun(NatType::Unknown).await; - let mut flags = peer_mgr.get_global_ctx().get_flags(); - flags.lazy_p2p = true; - peer_mgr.get_global_ctx().set_flags(flags); - peer_mgr - } - - fn metric_value(peer_mgr: &PeerManager, metric: MetricName, labels: &LabelSet) -> u64 { - peer_mgr - .get_global_ctx() - .stats_manager() - .get_metric(metric, labels) - .map(|metric| metric.value) - .unwrap_or(0) - } - - fn network_labels(peer_mgr: &PeerManager) -> LabelSet { - LabelSet::new().with_label_type(LabelType::NetworkName( - peer_mgr.get_global_ctx().get_network_name(), - )) - } - - struct TestCostCalculator { - costs: HashMap<(PeerId, PeerId), i32>, - } - - impl RouteCostCalculatorInterface for TestCostCalculator { - fn calculate_cost(&self, src: PeerId, dst: PeerId) -> i32 { - *self.costs.get(&(src, dst)).unwrap_or(&1) - } - } - - #[test] - fn recent_traffic_fanout_policy_only_marks_single_peer() { - assert!(PeerManager::should_mark_recent_traffic_for_fanout(0)); - assert!(PeerManager::should_mark_recent_traffic_for_fanout(1)); - assert!(!PeerManager::should_mark_recent_traffic_for_fanout(2)); - } - - fn route_with_ipv4( - peer_id: u32, - ipv4_addr: Option, - ) -> crate::proto::api::instance::Route { - crate::proto::api::instance::Route { - peer_id, - ipv4_addr: ipv4_addr.map(|addr| cidr::Ipv4Inet::new(addr, 24).unwrap().into()), - ..Default::default() - } - } - - #[test] - fn ipv4_broadcast_peer_selection_skips_peers_without_ipv4() { - let routes = vec![ - route_with_ipv4(1, Some(std::net::Ipv4Addr::new(10, 126, 126, 1))), - route_with_ipv4(2, None), - route_with_ipv4(3, Some(std::net::Ipv4Addr::new(10, 126, 126, 3))), - route_with_ipv4(4, None), - ]; - - assert_eq!( - PeerManager::select_ipv4_broadcast_peers(&routes, 3), - vec![1] - ); - } - - #[test] - fn gc_recent_traffic_removes_expired_and_connected_entries() { - let stale_peer = 1; - let direct_peer = 2; - let active_peer = 3; - let recent_have_traffic = dashmap::DashMap::new(); - - recent_have_traffic.insert( - stale_peer, - Instant::now() - PeerManager::RECENT_HAVE_TRAFFIC_TTL - Duration::from_millis(1), - ); - recent_have_traffic.insert(direct_peer, Instant::now()); - recent_have_traffic.insert(active_peer, Instant::now()); - - let future_peer = 4; - - recent_have_traffic.insert(future_peer, Instant::now() + Duration::from_secs(1)); - - PeerManager::gc_recent_traffic_entries(&recent_have_traffic, Instant::now(), |peer_id| { - peer_id == direct_peer - }); - - assert!(!recent_have_traffic.contains_key(&stale_peer)); - assert!(!recent_have_traffic.contains_key(&direct_peer)); - assert!(recent_have_traffic.contains_key(&active_peer)); - assert!(recent_have_traffic.contains_key(&future_peer)); - } - - #[tokio::test] - async fn recent_traffic_skips_direct_peers_and_clears_after_direct_connect() { - let peer_mgr_a = create_lazy_peer_manager().await; - let peer_mgr_b = create_mock_peer_manager_with_mock_stun(NatType::Unknown).await; - let peer_b_id = peer_mgr_b.my_peer_id(); - - peer_mgr_a.mark_recent_traffic(peer_b_id); - assert!(peer_mgr_a.has_recent_traffic(peer_b_id, Instant::now())); - - let (a_ring, b_ring) = create_ring_tunnel_pair(); - let (client_ret, server_ret) = tokio::join!( - peer_mgr_a.add_client_tunnel(a_ring, true), - peer_mgr_b.add_tunnel_as_server(b_ring, true) - ); - client_ret.unwrap(); - server_ret.unwrap(); - - wait_for_condition( - || { - let peer_mgr_a = peer_mgr_a.clone(); - async move { peer_mgr_a.has_directly_connected_conn(peer_b_id) } - }, - Duration::from_secs(5), - ) - .await; - - wait_for_condition( - || { - let peer_mgr_a = peer_mgr_a.clone(); - async move { !peer_mgr_a.has_recent_traffic(peer_b_id, Instant::now()) } - }, - Duration::from_secs(5), - ) - .await; - - peer_mgr_a.mark_recent_traffic(peer_b_id); - assert!( - !peer_mgr_a.has_recent_traffic(peer_b_id, Instant::now()), - "directly connected peers should not be tracked as lazy-p2p demand" - ); - } - - #[tokio::test] - async fn recent_traffic_notifies_only_when_demand_becomes_active() { - let peer_mgr_a = create_lazy_peer_manager().await; - let peer_mgr_b = create_mock_peer_manager_with_mock_stun(NatType::Unknown).await; - let peer_b_id = peer_mgr_b.my_peer_id(); - let signal = peer_mgr_a.p2p_demand_notify(); - - let initial_version = signal.version(); - peer_mgr_a.mark_recent_traffic(peer_b_id); - assert_eq!(signal.version(), initial_version + 1); - - let first_seen = *peer_mgr_a.recent_have_traffic.get(&peer_b_id).unwrap(); - tokio::time::sleep(Duration::from_millis(5)).await; - peer_mgr_a.mark_recent_traffic(peer_b_id); - assert_eq!( - signal.version(), - initial_version + 1, - "fresh demand should not wake all p2p workers again" - ); - let refreshed_seen = *peer_mgr_a.recent_have_traffic.get(&peer_b_id).unwrap(); - assert!(refreshed_seen > first_seen); - - if let Some(mut last_seen) = peer_mgr_a.recent_have_traffic.get_mut(&peer_b_id) { - *last_seen = - Instant::now() - PeerManager::RECENT_HAVE_TRAFFIC_TTL - Duration::from_millis(1); - } - peer_mgr_a.mark_recent_traffic(peer_b_id); - assert_eq!(signal.version(), initial_version + 2); - } - - #[test] - fn disable_relay_data_classifies_data_plane_packets_only() { - for packet_type in [ - PacketType::Data, - PacketType::KcpSrc, - PacketType::KcpDst, - PacketType::QuicSrc, - PacketType::QuicDst, - PacketType::DataWithKcpSrcModified, - PacketType::DataWithQuicSrcModified, - PacketType::ForeignNetworkPacket, - ] { - assert!(PeerManager::is_relay_data_packet(packet_type as u8)); - } - - for packet_type in [ - PacketType::RpcReq, - PacketType::RpcResp, - PacketType::Ping, - PacketType::Pong, - PacketType::HandShake, - PacketType::NoiseHandshakeMsg1, - PacketType::NoiseHandshakeMsg2, - PacketType::NoiseHandshakeMsg3, - PacketType::RelayHandshake, - PacketType::RelayHandshakeAck, - ] { - assert!(!PeerManager::is_relay_data_packet(packet_type as u8)); - } - } - - #[test] - fn disable_relay_data_inspects_foreign_network_inner_packet_type() { - let network_name = "net1".to_string(); - - let mut rpc_packet = ZCPacket::new_with_payload(b"rpc"); - rpc_packet.fill_peer_manager_hdr(1, 2, PacketType::RpcReq as u8); - let mut foreign_rpc_packet = - ZCPacket::new_for_foreign_network(&network_name, 2, &rpc_packet); - foreign_rpc_packet.fill_peer_manager_hdr(10, 20, PacketType::ForeignNetworkPacket as u8); - - assert_eq!( - foreign_rpc_packet.foreign_network_inner_packet_type(), - Some(PacketType::RpcReq as u8) - ); - assert!(!PeerManager::is_relay_data_zc_packet(&foreign_rpc_packet)); - - let mut data_packet = ZCPacket::new_with_payload(b"data"); - data_packet.fill_peer_manager_hdr(1, 2, PacketType::Data as u8); - let mut foreign_data_packet = - ZCPacket::new_for_foreign_network(&network_name, 2, &data_packet); - foreign_data_packet.fill_peer_manager_hdr(10, 20, PacketType::ForeignNetworkPacket as u8); - - assert_eq!( - foreign_data_packet.foreign_network_inner_packet_type(), - Some(PacketType::Data as u8) - ); - assert!(PeerManager::is_relay_data_zc_packet(&foreign_data_packet)); - } - - #[tokio::test] - async fn non_whitelisted_network_avoid_relay_survives_disable_relay_data_toggle() { - let global_ctx = get_mock_global_ctx(); - let mut flags = global_ctx.get_flags(); - flags.disable_relay_data = true; - flags.relay_network_whitelist = "other-network".to_string(); - global_ctx.set_flags(flags); - - let (packet_send, _packet_recv) = create_packet_recv_chan(); - let _peer_mgr = PeerManager::new(RouteAlgoType::Ospf, global_ctx.clone(), packet_send); - - let mut flags = global_ctx.get_flags(); - flags.disable_relay_data = false; - global_ctx.set_flags(flags); - - assert!(global_ctx.get_feature_flags().avoid_relay_data); - } - - #[tokio::test] - async fn send_msg_internal_does_not_record_tx_metrics_on_failed_delivery() { - let peer_mgr = create_mock_peer_manager_with_mock_stun(NatType::Unknown).await; - let dst_peer_id = peer_mgr.my_peer_id().wrapping_add(1); - let network_labels = LabelSet::new().with_label_type(LabelType::NetworkName( - peer_mgr.get_global_ctx().get_network_name(), - )); - - let mut pkt = ZCPacket::new_with_payload(b"tx"); - pkt.fill_peer_manager_hdr(peer_mgr.my_peer_id(), dst_peer_id, PacketType::Data as u8); - - let result = PeerManager::send_msg_internal( - &peer_mgr.peers, - &peer_mgr.foreign_network_client, - &peer_mgr.relay_peer_map, - Some(&peer_mgr.traffic_metrics), - pkt, - dst_peer_id, - ) - .await; - - assert!(result.is_err()); - assert_eq!( - peer_mgr - .get_global_ctx() - .stats_manager() - .get_metric(MetricName::TrafficBytesTx, &network_labels) - .unwrap() - .value, - 0 - ); - assert_eq!( - peer_mgr - .get_global_ctx() - .stats_manager() - .get_metric(MetricName::TrafficPacketsTx, &network_labels) - .unwrap() - .value, - 0 - ); - assert!( - peer_mgr - .get_global_ctx() - .stats_manager() - .get_metric( - MetricName::TrafficBytesTxByInstance, - &network_labels - .clone() - .with_label_type(LabelType::ToInstanceId("unknown".to_string())), - ) - .is_none() - ); - assert!( - peer_mgr - .get_global_ctx() - .stats_manager() - .get_metric( - MetricName::TrafficPacketsTxByInstance, - &network_labels.with_label_type(LabelType::ToInstanceId("unknown".to_string())), - ) - .is_none() - ); - } - - #[tokio::test] - async fn send_msg_internal_does_not_record_tx_metrics_for_self_loop() { - let (s, _r) = create_packet_recv_chan(); - let peer_mgr = Arc::new(PeerManager::new( - RouteAlgoType::None, - get_mock_global_ctx(), - s, - )); - let dst_peer_id = peer_mgr.my_peer_id(); - let network_labels = LabelSet::new().with_label_type(LabelType::NetworkName( - peer_mgr.get_global_ctx().get_network_name(), - )); - - let mut pkt = ZCPacket::new_with_payload(b"tx"); - pkt.fill_peer_manager_hdr(peer_mgr.my_peer_id(), dst_peer_id, PacketType::Data as u8); - - PeerManager::send_msg_internal( - &peer_mgr.peers, - &peer_mgr.foreign_network_client, - &peer_mgr.relay_peer_map, - Some(&peer_mgr.traffic_metrics), - pkt, - dst_peer_id, - ) - .await - .unwrap(); - - assert_eq!( - metric_value(&peer_mgr, MetricName::TrafficBytesTx, &network_labels), - 0 - ); - assert_eq!( - metric_value(&peer_mgr, MetricName::TrafficPacketsTx, &network_labels), - 0 - ); - assert_eq!( - metric_value( - &peer_mgr, - MetricName::TrafficControlBytesTx, - &network_labels - ), - 0 - ); - assert_eq!( - metric_value( - &peer_mgr, - MetricName::TrafficControlPacketsTx, - &network_labels - ), - 0 - ); - assert!( - peer_mgr - .get_global_ctx() - .stats_manager() - .get_metric( - MetricName::TrafficBytesTxByInstance, - &network_labels - .clone() - .with_label_type(LabelType::ToInstanceId("unknown".to_string())), - ) - .is_none() - ); - assert!( - peer_mgr - .get_global_ctx() - .stats_manager() - .get_metric( - MetricName::TrafficControlBytesTxByInstance, - &network_labels.with_label_type(LabelType::ToInstanceId("unknown".to_string())), - ) - .is_none() - ); - } - - #[tokio::test] - async fn send_msg_internal_records_data_metrics_for_direct_peer() { - let peer_mgr_a = create_mock_peer_manager_with_mock_stun(NatType::Unknown).await; - let peer_mgr_b = create_mock_peer_manager_with_mock_stun(NatType::Unknown).await; - connect_peer_manager(peer_mgr_a.clone(), peer_mgr_b.clone()).await; - wait_route_appear(peer_mgr_a.clone(), peer_mgr_b.clone()) - .await - .unwrap(); - - let a_network_labels = LabelSet::new().with_label_type(LabelType::NetworkName( - peer_mgr_a.get_global_ctx().get_network_name(), - )); - let b_network_labels = LabelSet::new().with_label_type(LabelType::NetworkName( - peer_mgr_b.get_global_ctx().get_network_name(), - )); - - let a_data_tx_before = - metric_value(&peer_mgr_a, MetricName::TrafficBytesTx, &a_network_labels); - let b_data_rx_before = - metric_value(&peer_mgr_b, MetricName::TrafficBytesRx, &b_network_labels); - let mut pkt = ZCPacket::new_with_payload(b"data"); - pkt.fill_peer_manager_hdr( - peer_mgr_a.my_peer_id(), - peer_mgr_b.my_peer_id(), - PacketType::Data as u8, - ); - let pkt_len = pkt.buf_len() as u64; - - PeerManager::send_msg_internal( - &peer_mgr_a.peers, - &peer_mgr_a.foreign_network_client, - &peer_mgr_a.relay_peer_map, - Some(&peer_mgr_a.traffic_metrics), - pkt, - peer_mgr_b.my_peer_id(), - ) - .await - .unwrap(); - - wait_for_condition( - || { - let peer_mgr_a = peer_mgr_a.clone(); - let peer_mgr_b = peer_mgr_b.clone(); - let a_network_labels = a_network_labels.clone(); - let b_network_labels = b_network_labels.clone(); - async move { - metric_value(&peer_mgr_a, MetricName::TrafficBytesTx, &a_network_labels) - >= a_data_tx_before + pkt_len - && metric_value(&peer_mgr_b, MetricName::TrafficBytesRx, &b_network_labels) - >= b_data_rx_before + pkt_len - } - }, - Duration::from_secs(5), - ) - .await; - } - - #[tokio::test] - async fn send_msg_internal_uses_latency_first_gateway_for_direct_peer() { - let peer_mgr_a = create_mock_peer_manager_with_mock_stun(NatType::Unknown).await; - let peer_mgr_b = create_mock_peer_manager_with_mock_stun(NatType::Unknown).await; - let peer_mgr_c = create_mock_peer_manager_with_mock_stun(NatType::Unknown).await; - - connect_peer_manager(peer_mgr_a.clone(), peer_mgr_b.clone()).await; - connect_peer_manager(peer_mgr_b.clone(), peer_mgr_c.clone()).await; - connect_peer_manager(peer_mgr_a.clone(), peer_mgr_c.clone()).await; - wait_route_appear(peer_mgr_a.clone(), peer_mgr_b.clone()) - .await - .unwrap(); - wait_route_appear(peer_mgr_b.clone(), peer_mgr_c.clone()) - .await - .unwrap(); - wait_route_appear(peer_mgr_a.clone(), peer_mgr_c.clone()) - .await - .unwrap(); - - peer_mgr_a - .get_route() - .set_route_cost_fn(Box::new(TestCostCalculator { - costs: HashMap::from([ - ((peer_mgr_a.my_peer_id(), peer_mgr_c.my_peer_id()), 100), - ((peer_mgr_a.my_peer_id(), peer_mgr_b.my_peer_id()), 1), - ((peer_mgr_b.my_peer_id(), peer_mgr_c.my_peer_id()), 1), - ]), - })) - .await; - - wait_for_condition( - || { - let peer_mgr_a = peer_mgr_a.clone(); - let peer_mgr_b = peer_mgr_b.clone(); - let peer_mgr_c = peer_mgr_c.clone(); - async move { - peer_mgr_a - .get_route() - .get_next_hop_with_policy(peer_mgr_c.my_peer_id(), NextHopPolicy::LeastCost) - .await - == Some(peer_mgr_b.my_peer_id()) - } - }, - Duration::from_secs(5), - ) - .await; - - let b_network_labels = network_labels(&peer_mgr_b); - let forwarded_bytes_before = metric_value( - &peer_mgr_b, - MetricName::TrafficBytesForwarded, - &b_network_labels, - ); - let forwarded_packets_before = metric_value( - &peer_mgr_b, - MetricName::TrafficPacketsForwarded, - &b_network_labels, - ); - - let mut pkt = ZCPacket::new_with_payload(b"latency-first"); - pkt.fill_peer_manager_hdr( - peer_mgr_a.my_peer_id(), - peer_mgr_c.my_peer_id(), - PacketType::Data as u8, - ); - pkt.mut_peer_manager_header() - .unwrap() - .set_latency_first(true); - let pkt_len = pkt.buf_len() as u64; - - PeerManager::send_msg_internal( - &peer_mgr_a.peers, - &peer_mgr_a.foreign_network_client, - &peer_mgr_a.relay_peer_map, - Some(&peer_mgr_a.traffic_metrics), - pkt, - peer_mgr_c.my_peer_id(), - ) - .await - .unwrap(); - - wait_for_condition( - || { - let peer_mgr_b = peer_mgr_b.clone(); - let b_network_labels = b_network_labels.clone(); - async move { - metric_value( - &peer_mgr_b, - MetricName::TrafficBytesForwarded, - &b_network_labels, - ) >= forwarded_bytes_before + pkt_len - && metric_value( - &peer_mgr_b, - MetricName::TrafficPacketsForwarded, - &b_network_labels, - ) > forwarded_packets_before - } - }, - Duration::from_secs(5), - ) - .await; - } - - #[tokio::test] - async fn send_msg_internal_records_control_metrics_for_direct_peer() { - let peer_mgr_a = create_mock_peer_manager_with_mock_stun(NatType::Unknown).await; - let peer_mgr_b = create_mock_peer_manager_with_mock_stun(NatType::Unknown).await; - connect_peer_manager(peer_mgr_a.clone(), peer_mgr_b.clone()).await; - wait_route_appear(peer_mgr_a.clone(), peer_mgr_b.clone()) - .await - .unwrap(); - - let a_network_labels = LabelSet::new().with_label_type(LabelType::NetworkName( - peer_mgr_a.get_global_ctx().get_network_name(), - )); - let b_network_labels = LabelSet::new().with_label_type(LabelType::NetworkName( - peer_mgr_b.get_global_ctx().get_network_name(), - )); - - let a_control_tx_before = metric_value( - &peer_mgr_a, - MetricName::TrafficControlBytesTx, - &a_network_labels, - ); - let b_control_rx_before = metric_value( - &peer_mgr_b, - MetricName::TrafficControlBytesRx, - &b_network_labels, - ); - let a_data_tx_before = - metric_value(&peer_mgr_a, MetricName::TrafficBytesTx, &a_network_labels); - let b_data_rx_before = - metric_value(&peer_mgr_b, MetricName::TrafficBytesRx, &b_network_labels); - - let mut pkt = ZCPacket::new_with_payload(b"ctrl"); - pkt.fill_peer_manager_hdr( - peer_mgr_a.my_peer_id(), - peer_mgr_b.my_peer_id(), - PacketType::RpcReq as u8, - ); - let pkt_len = pkt.buf_len() as u64; - - PeerManager::send_msg_internal( - &peer_mgr_a.peers, - &peer_mgr_a.foreign_network_client, - &peer_mgr_a.relay_peer_map, - Some(&peer_mgr_a.traffic_metrics), - pkt, - peer_mgr_b.my_peer_id(), - ) - .await - .unwrap(); - - wait_for_condition( - || { - let peer_mgr_a = peer_mgr_a.clone(); - let peer_mgr_b = peer_mgr_b.clone(); - let a_network_labels = a_network_labels.clone(); - let b_network_labels = b_network_labels.clone(); - async move { - metric_value( - &peer_mgr_a, - MetricName::TrafficControlBytesTx, - &a_network_labels, - ) >= a_control_tx_before + pkt_len - && metric_value( - &peer_mgr_b, - MetricName::TrafficControlBytesRx, - &b_network_labels, - ) >= b_control_rx_before + pkt_len - } - }, - Duration::from_secs(5), - ) - .await; - - assert_eq!( - metric_value(&peer_mgr_a, MetricName::TrafficBytesTx, &a_network_labels), - a_data_tx_before - ); - assert_eq!( - metric_value(&peer_mgr_b, MetricName::TrafficBytesRx, &b_network_labels), - b_data_rx_before - ); - } - - #[tokio::test] - async fn send_msg_internal_records_data_forwarded_metrics_for_transit_peer() { - let peer_mgr_a = create_mock_peer_manager_with_mock_stun(NatType::Unknown).await; - let peer_mgr_b = create_mock_peer_manager_with_mock_stun(NatType::Unknown).await; - let peer_mgr_c = create_mock_peer_manager_with_mock_stun(NatType::Unknown).await; - - connect_peer_manager(peer_mgr_a.clone(), peer_mgr_b.clone()).await; - connect_peer_manager(peer_mgr_b.clone(), peer_mgr_c.clone()).await; - wait_route_appear(peer_mgr_a.clone(), peer_mgr_c.clone()) - .await - .unwrap(); - - let b_network_labels = network_labels(&peer_mgr_b); - let forwarded_bytes_before = metric_value( - &peer_mgr_b, - MetricName::TrafficBytesForwarded, - &b_network_labels, - ); - let forwarded_packets_before = metric_value( - &peer_mgr_b, - MetricName::TrafficPacketsForwarded, - &b_network_labels, - ); - - let mut pkt = ZCPacket::new_with_payload(b"forward-data"); - pkt.fill_peer_manager_hdr( - peer_mgr_a.my_peer_id(), - peer_mgr_c.my_peer_id(), - PacketType::Data as u8, - ); - let pkt_len = pkt.buf_len() as u64; - - PeerManager::send_msg_internal( - &peer_mgr_a.peers, - &peer_mgr_a.foreign_network_client, - &peer_mgr_a.relay_peer_map, - Some(&peer_mgr_a.traffic_metrics), - pkt, - peer_mgr_c.my_peer_id(), - ) - .await - .unwrap(); - - wait_for_condition( - || { - let peer_mgr_b = peer_mgr_b.clone(); - let b_network_labels = b_network_labels.clone(); - async move { - metric_value( - &peer_mgr_b, - MetricName::TrafficBytesForwarded, - &b_network_labels, - ) >= forwarded_bytes_before + pkt_len - && metric_value( - &peer_mgr_b, - MetricName::TrafficPacketsForwarded, - &b_network_labels, - ) > forwarded_packets_before - } - }, - Duration::from_secs(5), - ) - .await; - } - - #[tokio::test] - async fn send_msg_internal_records_control_forwarded_metrics_for_transit_peer() { - let peer_mgr_a = create_mock_peer_manager_with_mock_stun(NatType::Unknown).await; - let peer_mgr_b = create_mock_peer_manager_with_mock_stun(NatType::Unknown).await; - let peer_mgr_c = create_mock_peer_manager_with_mock_stun(NatType::Unknown).await; - - connect_peer_manager(peer_mgr_a.clone(), peer_mgr_b.clone()).await; - connect_peer_manager(peer_mgr_b.clone(), peer_mgr_c.clone()).await; - wait_route_appear(peer_mgr_a.clone(), peer_mgr_c.clone()) - .await - .unwrap(); - - let b_network_labels = network_labels(&peer_mgr_b); - let forwarded_bytes_before = metric_value( - &peer_mgr_b, - MetricName::TrafficControlBytesForwarded, - &b_network_labels, - ); - let forwarded_packets_before = metric_value( - &peer_mgr_b, - MetricName::TrafficControlPacketsForwarded, - &b_network_labels, - ); - - let mut pkt = ZCPacket::new_with_payload(b"forward-control"); - pkt.fill_peer_manager_hdr( - peer_mgr_a.my_peer_id(), - peer_mgr_c.my_peer_id(), - PacketType::RpcReq as u8, - ); - let pkt_len = pkt.buf_len() as u64; - - PeerManager::send_msg_internal( - &peer_mgr_a.peers, - &peer_mgr_a.foreign_network_client, - &peer_mgr_a.relay_peer_map, - Some(&peer_mgr_a.traffic_metrics), - pkt, - peer_mgr_c.my_peer_id(), - ) - .await - .unwrap(); - - wait_for_condition( - || { - let peer_mgr_b = peer_mgr_b.clone(); - let b_network_labels = b_network_labels.clone(); - async move { - metric_value( - &peer_mgr_b, - MetricName::TrafficControlBytesForwarded, - &b_network_labels, - ) >= forwarded_bytes_before + pkt_len - && metric_value( - &peer_mgr_b, - MetricName::TrafficControlPacketsForwarded, - &b_network_labels, - ) > forwarded_packets_before - } - }, - Duration::from_secs(5), - ) - .await; - } - - #[tokio::test] - async fn recent_traffic_tolerates_future_timestamps() { - let peer_mgr_a = create_lazy_peer_manager().await; - let peer_mgr_b = create_mock_peer_manager_with_mock_stun(NatType::Unknown).await; - let peer_b_id = peer_mgr_b.my_peer_id(); - - peer_mgr_a - .recent_have_traffic - .insert(peer_b_id, Instant::now() + Duration::from_secs(1)); - - assert!(peer_mgr_a.has_recent_traffic(peer_b_id, Instant::now())); - peer_mgr_a.mark_recent_traffic(peer_b_id); - } - - #[tokio::test] - async fn drop_peer_manager() { - let peer_mgr_a = create_mock_peer_manager_with_mock_stun(NatType::Unknown).await; - let peer_mgr_b = create_mock_peer_manager_with_mock_stun(NatType::Unknown).await; - let peer_mgr_c = create_mock_peer_manager_with_mock_stun(NatType::Unknown).await; - connect_peer_manager(peer_mgr_a.clone(), peer_mgr_b.clone()).await; - connect_peer_manager(peer_mgr_b.clone(), peer_mgr_c.clone()).await; - connect_peer_manager(peer_mgr_a.clone(), peer_mgr_c.clone()).await; - - wait_route_appear(peer_mgr_a.clone(), peer_mgr_b.clone()) - .await - .unwrap(); - wait_route_appear(peer_mgr_a.clone(), peer_mgr_c.clone()) - .await - .unwrap(); - - // wait mgr_a have 2 peers - wait_for_condition( - || async { peer_mgr_a.get_peer_map().list_peers_with_conn().await.len() == 2 }, - std::time::Duration::from_secs(5), - ) - .await; - - drop(peer_mgr_b); - - wait_for_condition( - || async { peer_mgr_a.get_peer_map().list_peers_with_conn().await.len() == 1 }, - std::time::Duration::from_secs(5), - ) - .await; - } - - #[tokio::test] - async fn peer_manager_safe_mode_connect_between_peers() { - let peer_mgr_a = create_mock_peer_manager_with_mock_stun(NatType::Unknown).await; - let peer_mgr_b = create_mock_peer_manager_with_mock_stun(NatType::Unknown).await; - - peer_mgr_a - .get_global_ctx() - .config - .set_network_identity(NetworkIdentity::new("net1".to_string(), "sec1".to_string())); - peer_mgr_b - .get_global_ctx() - .config - .set_network_identity(NetworkIdentity::new("net1".to_string(), "sec1".to_string())); - - set_secure_mode_cfg(&peer_mgr_a.get_global_ctx(), true); - set_secure_mode_cfg(&peer_mgr_b.get_global_ctx(), true); - - let (a_ring, b_ring) = create_ring_tunnel_pair(); - let (a_ret, b_ret) = tokio::join!( - peer_mgr_a.add_client_tunnel(a_ring, false), - peer_mgr_b.add_tunnel_as_server(b_ring, true) - ); - let (peer_b_id, _) = a_ret.unwrap(); - b_ret.unwrap(); - - wait_for_condition( - || { - let peer_mgr_a = peer_mgr_a.clone(); - async move { - if !peer_mgr_a - .get_peer_map() - .list_peers_with_conn() - .await - .contains(&peer_b_id) - { - return false; - } - let Some(conns) = peer_mgr_a.get_peer_map().list_peer_conns(peer_b_id).await - else { - return false; - }; - conns.iter().any(|c| { - c.noise_local_static_pubkey.len() == 32 - && c.noise_remote_static_pubkey.len() == 32 - && c.secure_auth_level == SecureAuthLevel::NetworkSecretConfirmed as i32 - }) - } - }, - Duration::from_secs(10), - ) - .await; - - let peer_a_id = peer_mgr_a.my_peer_id(); - wait_for_condition( - || { - let peer_mgr_b = peer_mgr_b.clone(); - async move { - if !peer_mgr_b - .get_peer_map() - .list_peers_with_conn() - .await - .contains(&peer_a_id) - { - return false; - } - let Some(conns) = peer_mgr_b.get_peer_map().list_peer_conns(peer_a_id).await - else { - return false; - }; - conns.iter().any(|c| { - c.noise_local_static_pubkey.len() == 32 - && c.noise_remote_static_pubkey.len() == 32 - && c.secure_auth_level == SecureAuthLevel::NetworkSecretConfirmed as i32 - }) - } - }, - Duration::from_secs(10), - ) - .await; - } - - #[tokio::test] - async fn peer_manager_same_network_secure_mode_mismatch_rejected() { - let peer_mgr_client = create_mock_peer_manager_with_mock_stun(NatType::Unknown).await; - let peer_mgr_server = create_mock_peer_manager_with_mock_stun(NatType::Unknown).await; - - peer_mgr_client - .get_global_ctx() - .config - .set_network_identity(NetworkIdentity::new("net1".to_string(), "sec1".to_string())); - peer_mgr_server - .get_global_ctx() - .config - .set_network_identity(NetworkIdentity::new("net1".to_string(), "sec1".to_string())); - - set_secure_mode_cfg(&peer_mgr_server.get_global_ctx(), true); - - let (c_ring, s_ring) = create_ring_tunnel_pair(); - let (c_ret, s_ret) = tokio::join!( - peer_mgr_client.add_client_tunnel(c_ring, false), - peer_mgr_server.add_tunnel_as_server(s_ring, true) - ); - let _ = c_ret; - assert!( - s_ret.is_err(), - "same-network peer with mismatched secure mode should be rejected" - ); - - wait_for_condition( - || { - let peer_mgr_server = peer_mgr_server.clone(); - async move { - peer_mgr_server - .get_peer_map() - .list_peers_with_conn() - .await - .is_empty() - } - }, - Duration::from_secs(5), - ) - .await; - } - - #[tokio::test] - async fn credential_node_rejects_legacy_client() { - let peer_mgr_client = create_mock_peer_manager_with_mock_stun(NatType::Unknown).await; - let peer_mgr_server = create_mock_peer_manager_with_mock_stun(NatType::Unknown).await; - - peer_mgr_client - .get_global_ctx() - .config - .set_network_identity(NetworkIdentity::new("net1".to_string(), "sec1".to_string())); - peer_mgr_server - .get_global_ctx() - .config - .set_network_identity(NetworkIdentity::new_credential("net1".to_string())); - - set_secure_mode_cfg(&peer_mgr_server.get_global_ctx(), true); - - let (c_ring, s_ring) = create_ring_tunnel_pair(); - let (c_ret, s_ret) = tokio::join!( - peer_mgr_client.add_client_tunnel(c_ring, false), - peer_mgr_server.add_tunnel_as_server(s_ring, true) - ); - - let _ = c_ret; - assert!( - s_ret.is_err(), - "credential server should reject legacy client" - ); - - wait_for_condition( - || { - let peer_mgr_server = peer_mgr_server.clone(); - async move { - peer_mgr_server - .get_peer_map() - .list_peers_with_conn() - .await - .is_empty() - } - }, - Duration::from_secs(5), - ) - .await; - } - - #[tokio::test] - async fn peer_manager_safe_mode_shared_node_pinning_connect() { - let peer_mgr_client = create_mock_peer_manager_with_mock_stun(NatType::Unknown).await; - let peer_mgr_server = create_mock_peer_manager_with_mock_stun(NatType::Unknown).await; - - peer_mgr_client - .get_global_ctx() - .config - .set_network_identity(NetworkIdentity::new("user".to_string(), "sec1".to_string())); - peer_mgr_server - .get_global_ctx() - .config - .set_network_identity(NetworkIdentity { - network_name: "shared".to_string(), - network_secret: None, - network_secret_digest: None, - }); - - set_secure_mode_cfg(&peer_mgr_client.get_global_ctx(), true); - set_secure_mode_cfg(&peer_mgr_server.get_global_ctx(), true); - - let server_pub_b64 = peer_mgr_server - .get_global_ctx() - .config - .get_secure_mode() - .unwrap() - .local_public_key - .unwrap(); - - let (a_ring, b_ring) = create_ring_tunnel_pair(); - let server_remote_url: url::Url = a_ring - .info() - .unwrap() - .remote_addr - .unwrap() - .url - .parse() - .unwrap(); - peer_mgr_client.get_global_ctx().config.set_peers(vec![ - crate::common::config::PeerConfig { - uri: server_remote_url, - peer_public_key: Some(server_pub_b64.clone()), - }, - ]); - - let (c_ret, s_ret) = tokio::join!( - peer_mgr_client.add_client_tunnel(a_ring, false), - peer_mgr_server.add_tunnel_as_server(b_ring, true) - ); - c_ret.unwrap(); - s_ret.unwrap(); - - wait_for_condition( - || { - let peer_mgr_client = peer_mgr_client.clone(); - async move { - let foreign_peer_map = - peer_mgr_client.get_foreign_network_client().get_peer_map(); - if foreign_peer_map.list_peers_with_conn().await.len() != 1 { - return false; - } - let Some(peer_id) = foreign_peer_map - .list_peers_with_conn() - .await - .into_iter() - .next() - else { - return false; - }; - let Some(conns) = foreign_peer_map.list_peer_conns(peer_id).await else { - return false; - }; - conns.iter().any(|c| { - c.secure_auth_level == SecureAuthLevel::PeerVerified as i32 - && c.noise_local_static_pubkey.len() == 32 - && c.noise_remote_static_pubkey.len() == 32 - }) - } - }, - Duration::from_secs(10), - ) - .await; - - wait_for_condition( - || { - let peer_mgr_server = peer_mgr_server.clone(); - async move { - let foreigns = peer_mgr_server - .get_foreign_network_manager() - .list_foreign_networks() - .await; - let Some(entry) = foreigns.foreign_networks.get("user") else { - return false; - }; - entry.peers.iter().any(|p| { - p.conns - .iter() - .any(|c| c.noise_local_static_pubkey.len() == 32) - }) - } - }, - Duration::from_secs(10), - ) - .await; - } - - async fn connect_peer_manager_with( - client_mgr: Arc, - server_mgr: &Arc, - mut client: C, - server: &mut L, - ) { - server.listen().await.unwrap(); - - tokio::spawn(async move { - client.set_bind_addrs(vec![]); - client_mgr.try_direct_connect(client).await.unwrap(); - }); - - server_mgr - .add_client_tunnel(server.accept().await.unwrap(), false) - .await - .unwrap(); - } - - #[rstest::rstest] - #[tokio::test] - #[serial_test::serial(forward_packet_test)] - async fn forward_packet( - #[values("tcp", "udp", "wg", "quic")] proto1: &str, - #[values("tcp", "udp", "wg", "quic")] proto2: &str, - ) { - use crate::proto::{ - rpc_impl::RpcController, - tests::{GreetingClientFactory, SayHelloRequest}, - }; - - let peer_mgr_a = create_mock_peer_manager_with_mock_stun(NatType::Unknown).await; - register_service(&peer_mgr_a.peer_rpc_mgr, "", 0, "hello a"); - - let peer_mgr_b = create_mock_peer_manager_with_mock_stun(NatType::Unknown).await; - - let peer_mgr_c = create_mock_peer_manager_with_mock_stun(NatType::Unknown).await; - register_service(&peer_mgr_c.peer_rpc_mgr, "", 0, "hello c"); - - let mut listener1 = create_listener_by_url( - &format!("{}://0.0.0.0:31013", proto1).parse().unwrap(), - peer_mgr_b.get_global_ctx(), - ) - .unwrap(); - let connector1 = create_connector_by_url( - format!("{}://127.0.0.1:31013", proto1).as_str(), - &peer_mgr_a.get_global_ctx(), - crate::tunnel::IpVersion::Both, - ) - .await - .unwrap(); - connect_peer_manager_with(peer_mgr_a.clone(), &peer_mgr_b, connector1, &mut listener1) - .await; - - wait_route_appear(peer_mgr_a.clone(), peer_mgr_b.clone()) - .await - .unwrap(); - - let mut listener2 = create_listener_by_url( - &format!("{}://0.0.0.0:31014", proto2).parse().unwrap(), - peer_mgr_c.get_global_ctx(), - ) - .unwrap(); - let connector2 = create_connector_by_url( - format!("{}://127.0.0.1:31014", proto2).as_str(), - &peer_mgr_b.get_global_ctx(), - crate::tunnel::IpVersion::Both, - ) - .await - .unwrap(); - connect_peer_manager_with(peer_mgr_b.clone(), &peer_mgr_c, connector2, &mut listener2) - .await; - - wait_route_appear(peer_mgr_a.clone(), peer_mgr_c.clone()) - .await - .unwrap(); - - let stub = peer_mgr_a - .peer_rpc_mgr - .rpc_client() - .scoped_client::>( - peer_mgr_a.my_peer_id, - peer_mgr_c.my_peer_id, - "".to_string(), - ); - - let ret = stub - .say_hello( - RpcController::default(), - SayHelloRequest { - name: "abc".to_string(), - }, - ) - .await - .unwrap(); - - assert_eq!(ret.greeting, "hello c abc!"); - } - - #[tokio::test] - async fn communicate_between_enc_and_non_enc() { - let create_mgr = |enable_encryption| async move { - let (s, _r) = create_packet_recv_chan(); - let mock_global_ctx = get_mock_global_ctx(); - mock_global_ctx.set_flags(Flags { - enable_encryption, - data_compress_algo: CompressionAlgoPb::Zstd.into(), - ..Default::default() - }); - let peer_mgr = Arc::new(PeerManager::new(RouteAlgoType::Ospf, mock_global_ctx, s)); - peer_mgr.run().await.unwrap(); - peer_mgr - }; - - let peer_mgr_a = create_mgr(true).await; - let peer_mgr_b = create_mgr(false).await; - - connect_peer_manager(peer_mgr_a.clone(), peer_mgr_b.clone()).await; - - // wait 5sec should not crash. - tokio::time::sleep(Duration::from_secs(5)).await; - - // both mgr should alive - let mgr_c = create_mgr(true).await; - connect_peer_manager(peer_mgr_a.clone(), mgr_c.clone()).await; - wait_route_appear(mgr_c, peer_mgr_a).await.unwrap(); - - let mgr_d = create_mgr(false).await; - connect_peer_manager(peer_mgr_b.clone(), mgr_d.clone()).await; - wait_route_appear(mgr_d, peer_mgr_b).await.unwrap(); - } - - #[tokio::test] - async fn test_avoid_relay_data() { - // a->b->c - // a->d->e->c - let peer_mgr_a = create_mock_peer_manager_with_mock_stun(NatType::Unknown).await; - let peer_mgr_b = create_mock_peer_manager_with_mock_stun(NatType::Unknown).await; - let peer_mgr_c = create_mock_peer_manager_with_mock_stun(NatType::Unknown).await; - let peer_mgr_d = create_mock_peer_manager_with_mock_stun(NatType::Unknown).await; - let peer_mgr_e = create_mock_peer_manager_with_mock_stun(NatType::Unknown).await; - - println!("peer_mgr_a: {}", peer_mgr_a.my_peer_id); - println!("peer_mgr_b: {}", peer_mgr_b.my_peer_id); - println!("peer_mgr_c: {}", peer_mgr_c.my_peer_id); - println!("peer_mgr_d: {}", peer_mgr_d.my_peer_id); - println!("peer_mgr_e: {}", peer_mgr_e.my_peer_id); - - connect_peer_manager(peer_mgr_a.clone(), peer_mgr_b.clone()).await; - connect_peer_manager(peer_mgr_b.clone(), peer_mgr_c.clone()).await; - - connect_peer_manager(peer_mgr_a.clone(), peer_mgr_d.clone()).await; - connect_peer_manager(peer_mgr_d.clone(), peer_mgr_e.clone()).await; - connect_peer_manager(peer_mgr_e.clone(), peer_mgr_c.clone()).await; - - // when b's avoid_relay_data is false, a->c should route through b and cost is 2 - wait_route_appear_with_cost(peer_mgr_a.clone(), peer_mgr_c.my_peer_id, Some(2)) - .await - .unwrap(); - let ret = peer_mgr_a - .get_route() - .get_next_hop_with_policy(peer_mgr_c.my_peer_id, NextHopPolicy::LeastCost) - .await; - assert_eq!(ret, Some(peer_mgr_b.my_peer_id)); - - // when b's avoid_relay_data is true, a->c should route through d and e, cost is 3 - peer_mgr_b - .get_global_ctx() - .set_avoid_relay_data_preference(true); - tokio::time::sleep(Duration::from_secs(2)).await; - if wait_route_appear_with_cost(peer_mgr_a.clone(), peer_mgr_c.my_peer_id, Some(3)) - .await - .is_err() - { - panic!( - "route not appear, a route table: {}, table: {:#?}", - peer_mgr_a.get_route().dump().await, - peer_mgr_a.get_route().list_routes().await - ) - } - - let ret = peer_mgr_a - .get_route() - .get_next_hop_with_policy(peer_mgr_c.my_peer_id, NextHopPolicy::LeastCost) - .await; - assert_eq!(ret, Some(peer_mgr_d.my_peer_id)); - - println!("route table: {:#?}", peer_mgr_a.list_routes().await); - - // drop e, path should go back to through b - drop(peer_mgr_e); - wait_route_appear_with_cost(peer_mgr_a.clone(), peer_mgr_c.my_peer_id, Some(2)) - .await - .unwrap(); - let ret = peer_mgr_a - .get_route() - .get_next_hop_with_policy(peer_mgr_c.my_peer_id, NextHopPolicy::LeastCost) - .await; - assert_eq!(ret, Some(peer_mgr_b.my_peer_id)); - } - - #[tokio::test] - async fn test_client_inbound_blackhole() { - let peer_mgr_a = create_mock_peer_manager_with_mock_stun(NatType::Unknown).await; - let peer_mgr_b = create_mock_peer_manager_with_mock_stun(NatType::Unknown).await; - - // a is client, b is server - - let (a_ring, b_ring) = create_ring_tunnel_pair(); - let a_ring = Box::new(TunnelWithFilter::new( - a_ring, - DropSendTunnelFilter::new(2, 50000), - )); - - let a_mgr_copy = peer_mgr_a.clone(); - tokio::spawn(async move { - a_mgr_copy.add_client_tunnel(a_ring, false).await.unwrap(); - }); - let b_mgr_copy = peer_mgr_b.clone(); - tokio::spawn(async move { - b_mgr_copy.add_tunnel_as_server(b_ring, true).await.unwrap(); - }); - - wait_for_condition( - || async { - let peers = peer_mgr_a.list_peers().await; - peers.is_empty() - }, - Duration::from_secs(10), - ) - .await; - } - - #[tokio::test] - async fn close_conn_in_peer_map() { - let peer_mgr_a = create_mock_peer_manager_with_mock_stun(NatType::Unknown).await; - let peer_mgr_b = create_mock_peer_manager_with_mock_stun(NatType::Unknown).await; - connect_peer_manager(peer_mgr_a.clone(), peer_mgr_b.clone()).await; - wait_route_appear(peer_mgr_a.clone(), peer_mgr_b.clone()) - .await - .unwrap(); - - let conns = peer_mgr_a - .get_peer_map() - .list_peer_conns(peer_mgr_b.my_peer_id) - .await; - assert!(conns.is_some()); - let conn_info = conns.as_ref().unwrap().first().unwrap(); - - peer_mgr_a - .close_peer_conn(peer_mgr_b.my_peer_id, &conn_info.conn_id.parse().unwrap()) - .await - .unwrap(); - - wait_for_condition( - || async { - let peers = peer_mgr_a.list_peers().await; - peers.is_empty() - }, - Duration::from_secs(10), - ) - .await; - // a is client, b is server - } - - #[tokio::test] - async fn expired_credential_peer_conn_is_closed_without_ospf() { - let (admin_ch, _admin_rx) = create_packet_recv_chan(); - let admin_ctx = get_mock_global_ctx(); - admin_ctx.config.set_network_identity(NetworkIdentity::new( - "net1".to_string(), - "secret".to_string(), - )); - set_secure_mode_cfg(&admin_ctx, true); - let admin = Arc::new(PeerManager::new( - RouteAlgoType::None, - admin_ctx.clone(), - admin_ch, - )); - admin.run().await.unwrap(); - - let (_cred_id, cred_secret) = admin_ctx.get_credential_manager().generate_credential( - vec![], - false, - vec![], - Duration::from_secs(1), - ); - let privkey_bytes: [u8; 32] = base64::engine::general_purpose::STANDARD - .decode(&cred_secret) - .unwrap() - .try_into() - .unwrap(); - let private = x25519_dalek::StaticSecret::from(privkey_bytes); - let public = x25519_dalek::PublicKey::from(&private); - let (credential_ch, _credential_rx) = create_packet_recv_chan(); - let credential_ctx = get_mock_global_ctx(); - credential_ctx - .config - .set_network_identity(NetworkIdentity::new_credential("net1".to_string())); - credential_ctx - .config - .set_secure_mode(Some(SecureModeConfig { - enabled: true, - local_private_key: Some( - base64::engine::general_purpose::STANDARD.encode(private.as_bytes()), - ), - local_public_key: Some( - base64::engine::general_purpose::STANDARD.encode(public.as_bytes()), - ), - })); - let credential = Arc::new(PeerManager::new( - RouteAlgoType::None, - credential_ctx, - credential_ch, - )); - credential.run().await.unwrap(); - let credential_peer_id = credential.my_peer_id(); - - connect_peer_manager(credential.clone(), admin.clone()).await; - - wait_for_condition( - || { - let admin = admin.clone(); - async move { - admin - .get_peer_map() - .list_peer_conns(credential_peer_id) - .await - .is_some_and(|conns| !conns.is_empty()) - } - }, - Duration::from_secs(5), - ) - .await; - - wait_for_condition( - || { - let admin = admin.clone(); - async move { - admin - .get_peer_map() - .list_peer_conns(credential_peer_id) - .await - .is_none_or(|conns| conns.is_empty()) - } - }, - Duration::from_secs(5), - ) - .await; - } - - #[tokio::test] - async fn close_conn_in_foreign_network_client() { - let peer_mgr_server = create_mock_peer_manager_with_name("server".to_string()).await; - let peer_mgr_client = create_mock_peer_manager_with_name("client".to_string()).await; - connect_peer_manager(peer_mgr_client.clone(), peer_mgr_server.clone()).await; - wait_for_condition( - || async { - peer_mgr_client - .get_foreign_network_client() - .list_public_peers() - .await - .len() - == 1 - }, - Duration::from_secs(3), - ) - .await; - - let peer_id = peer_mgr_client - .foreign_network_client - .list_public_peers() - .await[0]; - let conns = peer_mgr_client - .foreign_network_client - .get_peer_map() - .list_peer_conns(peer_id) - .await; - assert!(conns.is_some()); - let conn_info = conns.as_ref().unwrap().first().unwrap(); - peer_mgr_client - .close_peer_conn(peer_id, &conn_info.conn_id.parse().unwrap()) - .await - .unwrap(); - - wait_for_condition( - || async { - peer_mgr_client - .get_foreign_network_client() - .list_public_peers() - .await - .is_empty() - }, - Duration::from_secs(10), - ) - .await; - } - - #[tokio::test] - async fn close_conn_in_foreign_network_manager() { - let peer_mgr_server = create_mock_peer_manager_with_name("server".to_string()).await; - let peer_mgr_client = create_mock_peer_manager_with_name("client".to_string()).await; - connect_peer_manager(peer_mgr_client.clone(), peer_mgr_server.clone()).await; - wait_for_condition( - || async { - peer_mgr_client - .get_foreign_network_client() - .list_public_peers() - .await - .len() - == 1 - }, - Duration::from_secs(3), - ) - .await; - - let conns = peer_mgr_server - .foreign_network_manager - .list_foreign_networks() - .await; - let client_info = conns.foreign_networks["client"].peers[0].clone(); - let conn_info = client_info.conns[0].clone(); - peer_mgr_server - .close_peer_conn(client_info.peer_id, &conn_info.conn_id.parse().unwrap()) - .await - .unwrap(); - - wait_for_condition( - || async { - peer_mgr_client - .get_foreign_network_client() - .list_public_peers() - .await - .is_empty() - }, - Duration::from_secs(10), - ) - .await; - } -} diff --git a/easytier/src/peers/peer_rpc.rs b/easytier/src/peers/peer_rpc.rs deleted file mode 100644 index 37da48cd..00000000 --- a/easytier/src/peers/peer_rpc.rs +++ /dev/null @@ -1,347 +0,0 @@ -use std::sync::{Arc, Mutex}; - -use futures::{SinkExt as _, StreamExt}; -use tokio::task::JoinSet; - -use crate::{ - common::{PeerId, error::Error, stats_manager::StatsManager}, - proto::rpc_impl::{self, bidirect::BidirectRpcManager}, - tunnel::packet_def::ZCPacket, -}; - -const RPC_PACKET_CONTENT_MTU: usize = 1300; - -type PeerRpcServiceId = u32; -type PeerRpcTransactId = u32; - -#[async_trait::async_trait] -#[auto_impl::auto_impl(Arc)] -pub trait PeerRpcManagerTransport: Send + Sync + 'static { - fn my_peer_id(&self) -> PeerId; - async fn send(&self, msg: ZCPacket, dst_peer_id: PeerId) -> Result<(), Error>; - async fn recv(&self) -> Result; -} - -// handle rpc request from one peer -pub struct PeerRpcManager { - tspt: Arc>, - bidirect_rpc: BidirectRpcManager, - tasks: Mutex>, -} - -impl std::fmt::Debug for PeerRpcManager { - fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { - f.debug_struct("PeerRpcManager") - .field("node_id", &self.tspt.my_peer_id()) - .finish() - } -} - -impl PeerRpcManager { - pub fn new(tspt: impl PeerRpcManagerTransport) -> Self { - Self { - tspt: Arc::new(Box::new(tspt)), - bidirect_rpc: BidirectRpcManager::new(), - - tasks: Mutex::new(JoinSet::new()), - } - } - - pub fn new_with_stats_manager( - tspt: impl PeerRpcManagerTransport, - stats_manager: Arc, - ) -> Self { - Self { - tspt: Arc::new(Box::new(tspt)), - bidirect_rpc: BidirectRpcManager::new_with_stats_manager(stats_manager), - - tasks: Mutex::new(JoinSet::new()), - } - } - - pub fn run(&self) { - let ret = self.bidirect_rpc.run_and_create_tunnel(); - let (mut rx, mut tx) = ret.split(); - let tspt = self.tspt.clone(); - self.tasks.lock().unwrap().spawn(async move { - while let Some(Ok(packet)) = rx.next().await { - let dst_peer_id = packet.peer_manager_header().unwrap().to_peer_id.into(); - if let Err(e) = tspt.send(packet, dst_peer_id).await { - tracing::error!("send to rpc tspt error: {:?}", e); - } - } - }); - - let tspt = self.tspt.clone(); - self.tasks.lock().unwrap().spawn(async move { - while let Ok(packet) = tspt.recv().await { - if let Err(e) = tx.send(packet).await { - tracing::error!("send to rpc tspt error: {:?}", e); - } - } - }); - } - - pub fn rpc_client(&self) -> &rpc_impl::client::Client { - self.bidirect_rpc.rpc_client() - } - - pub fn rpc_server(&self) -> &rpc_impl::server::Server { - self.bidirect_rpc.rpc_server() - } - - pub fn my_peer_id(&self) -> PeerId { - self.tspt.my_peer_id() - } -} - -impl Drop for PeerRpcManager { - fn drop(&mut self) { - tracing::debug!("PeerRpcManager drop, my_peer_id: {:?}", self.my_peer_id()); - } -} - -#[cfg(test)] -pub mod tests { - use std::{pin::Pin, sync::Arc}; - - use futures::{SinkExt, StreamExt}; - use tokio::sync::Mutex; - - use crate::{ - common::{PeerId, error::Error, new_peer_id}, - peers::{ - peer_rpc::PeerRpcManager, - tests::{connect_peer_manager, create_mock_peer_manager, wait_route_appear}, - }, - proto::{ - rpc_impl::RpcController, - tests::{GreetingClientFactory, GreetingServer, GreetingService, SayHelloRequest}, - }, - tunnel::{ - Tunnel, ZCPacketSink, ZCPacketStream, packet_def::ZCPacket, - ring::create_ring_tunnel_pair, - }, - }; - - use super::PeerRpcManagerTransport; - - fn random_string(len: usize) -> String { - use rand::Rng; - use rand::distributions::Alphanumeric; - let mut rng = rand::thread_rng(); - let s: Vec = std::iter::repeat(()) - .map(|()| rng.sample(Alphanumeric)) - .take(len) - .collect(); - String::from_utf8(s).unwrap() - } - - pub fn register_service(rpc_mgr: &PeerRpcManager, domain: &str, delay_ms: u64, prefix: &str) { - rpc_mgr.rpc_server().registry().register( - GreetingServer::new(GreetingService { - delay_ms, - prefix: prefix.to_string(), - }), - domain, - ); - } - - #[tokio::test] - async fn peer_rpc_basic_test() { - struct MockTransport { - sink: Arc>>>, - stream: Arc>>>, - my_peer_id: PeerId, - } - - #[async_trait::async_trait] - impl PeerRpcManagerTransport for MockTransport { - fn my_peer_id(&self) -> PeerId { - self.my_peer_id - } - async fn send(&self, msg: ZCPacket, _dst_peer_id: PeerId) -> Result<(), Error> { - println!("rpc mgr send: {:?}", msg); - self.sink.lock().await.send(msg).await.unwrap(); - Ok(()) - } - async fn recv(&self) -> Result { - let ret = self.stream.lock().await.next().await.unwrap(); - println!("rpc mgr recv: {:?}", ret); - return ret.map_err(|e| e.into()); - } - } - - let (ct, st) = create_ring_tunnel_pair(); - let (cts, ctsr) = ct.split(); - let (sts, stsr) = st.split(); - - let server_rpc_mgr = PeerRpcManager::new(MockTransport { - sink: Arc::new(Mutex::new(ctsr)), - stream: Arc::new(Mutex::new(cts)), - my_peer_id: new_peer_id(), - }); - server_rpc_mgr.run(); - register_service(&server_rpc_mgr, "test", 0, "Hello"); - - let client_rpc_mgr = PeerRpcManager::new(MockTransport { - sink: Arc::new(Mutex::new(stsr)), - stream: Arc::new(Mutex::new(sts)), - my_peer_id: new_peer_id(), - }); - client_rpc_mgr.run(); - - let stub = client_rpc_mgr - .rpc_client() - .scoped_client::>(1, 1, "test".to_string()); - - let msg = random_string(8192); - let ret = stub - .say_hello( - RpcController::default(), - SayHelloRequest { name: msg.clone() }, - ) - .await - .unwrap(); - - println!("ret: {:?}", ret); - assert_eq!(ret.greeting, format!("Hello {}!", msg)); - - let msg = random_string(10); - let ret = stub - .say_hello( - RpcController::default(), - SayHelloRequest { name: msg.clone() }, - ) - .await - .unwrap(); - - println!("ret: {:?}", ret); - assert_eq!(ret.greeting, format!("Hello {}!", msg)); - } - - #[tokio::test] - async fn test_rpc_with_peer_manager() { - let peer_mgr_a = create_mock_peer_manager().await; - let peer_mgr_b = create_mock_peer_manager().await; - let peer_mgr_c = create_mock_peer_manager().await; - connect_peer_manager(peer_mgr_a.clone(), peer_mgr_b.clone()).await; - connect_peer_manager(peer_mgr_b.clone(), peer_mgr_c.clone()).await; - - wait_route_appear(peer_mgr_a.clone(), peer_mgr_b.clone()) - .await - .unwrap(); - wait_route_appear(peer_mgr_a.clone(), peer_mgr_c.clone()) - .await - .unwrap(); - - assert_eq!(peer_mgr_a.get_peer_map().list_peers().len(), 1); - assert_eq!( - peer_mgr_a.get_peer_map().list_peers()[0], - peer_mgr_b.my_peer_id() - ); - - assert_eq!(peer_mgr_c.get_peer_map().list_peers().len(), 1); - assert_eq!( - peer_mgr_c.get_peer_map().list_peers()[0], - peer_mgr_b.my_peer_id() - ); - - register_service(&peer_mgr_b.get_peer_rpc_mgr(), "test", 0, "Hello"); - - let msg = random_string(16 * 1024); - let stub = peer_mgr_a - .get_peer_rpc_mgr() - .rpc_client() - .scoped_client::>( - peer_mgr_a.my_peer_id(), - peer_mgr_b.my_peer_id(), - "test".to_string(), - ); - - let ret = stub - .say_hello( - RpcController::default(), - SayHelloRequest { name: msg.clone() }, - ) - .await - .unwrap(); - assert_eq!(ret.greeting, format!("Hello {}!", msg)); - - // call again - let msg = random_string(16 * 1024); - let ret = stub - .say_hello( - RpcController::default(), - SayHelloRequest { name: msg.clone() }, - ) - .await - .unwrap(); - assert_eq!(ret.greeting, format!("Hello {}!", msg)); - - let msg = random_string(16 * 1024); - let ret = stub - .say_hello( - RpcController::default(), - SayHelloRequest { name: msg.clone() }, - ) - .await - .unwrap(); - assert_eq!(ret.greeting, format!("Hello {}!", msg)); - } - - #[tokio::test] - async fn test_multi_domain_with_peer_manager() { - let peer_mgr_a = create_mock_peer_manager().await; - let peer_mgr_b = create_mock_peer_manager().await; - connect_peer_manager(peer_mgr_a.clone(), peer_mgr_b.clone()).await; - wait_route_appear(peer_mgr_a.clone(), peer_mgr_b.clone()) - .await - .unwrap(); - - assert_eq!(peer_mgr_a.get_peer_map().list_peers().len(), 1); - assert_eq!( - peer_mgr_a.get_peer_map().list_peers()[0], - peer_mgr_b.my_peer_id() - ); - - register_service(&peer_mgr_b.get_peer_rpc_mgr(), "test1", 0, "Hello"); - register_service(&peer_mgr_b.get_peer_rpc_mgr(), "test2", 20000, "Hello2"); - - let stub1 = peer_mgr_a - .get_peer_rpc_mgr() - .rpc_client() - .scoped_client::>( - peer_mgr_a.my_peer_id(), - peer_mgr_b.my_peer_id(), - "test1".to_string(), - ); - - let stub2 = peer_mgr_a - .get_peer_rpc_mgr() - .rpc_client() - .scoped_client::>( - peer_mgr_a.my_peer_id(), - peer_mgr_b.my_peer_id(), - "test2".to_string(), - ); - - let msg = random_string(16 * 1024); - let ret = stub1 - .say_hello( - RpcController::default(), - SayHelloRequest { name: msg.clone() }, - ) - .await - .unwrap(); - assert_eq!(ret.greeting, format!("Hello {}!", msg)); - - let ret = stub2 - .say_hello( - RpcController::default(), - SayHelloRequest { name: msg.clone() }, - ) - .await; - assert!(ret.is_err() && ret.unwrap_err().to_string().contains("Timeout")); - } -} diff --git a/easytier/src/peers/peer_rpc_service.rs b/easytier/src/peers/peer_rpc_service.rs deleted file mode 100644 index ec9db867..00000000 --- a/easytier/src/peers/peer_rpc_service.rs +++ /dev/null @@ -1,288 +0,0 @@ -use std::net::{IpAddr, Ipv6Addr, SocketAddr}; - -use crate::{ - common::{global_ctx::ArcGlobalCtx, network::IPCollector}, - proto::{ - common::Void, - peer_rpc::{ - DirectConnectorRpc, GetIpListRequest, GetIpListResponse, SendUdpHolePunchPacketRequest, - }, - rpc_types::{self, controller::BaseController}, - }, - tunnel::udp, -}; - -const MAX_UDP_HOLE_PUNCH_CONNECTOR_ADDRS: usize = 16; - -fn remove_easytier_managed_ipv6s(ret: &mut GetIpListResponse, global_ctx: &ArcGlobalCtx) { - ret.interface_ipv6s.retain(|ip| { - let ip = std::net::Ipv6Addr::from(*ip); - !global_ctx.is_ip_easytier_managed_ipv6(&ip) - }); - - if ret - .public_ipv6 - .as_ref() - .map(|ip| std::net::Ipv6Addr::from(*ip)) - .is_some_and(|ip| global_ctx.is_ip_easytier_managed_ipv6(&ip)) - { - ret.public_ipv6 = None; - } -} - -fn is_usable_preferred_src_ipv6(ip: &Ipv6Addr, global_ctx: &ArcGlobalCtx) -> bool { - !global_ctx.is_ip_easytier_managed_ipv6(ip) - && !ip.is_loopback() - && !ip.is_unspecified() - && !ip.is_unique_local() - && !ip.is_unicast_link_local() - && !ip.is_multicast() -} - -async fn local_preferred_src_ipv6( - global_ctx: &ArcGlobalCtx, - preferred_src_ipv6: Option, -) -> Option { - let preferred_src_ipv6 = preferred_src_ipv6.map(Ipv6Addr::from)?; - if !is_usable_preferred_src_ipv6(&preferred_src_ipv6, global_ctx) { - tracing::debug!( - ?preferred_src_ipv6, - "ignore unusable preferred IPv6 source for udp hole punch" - ); - return None; - } - - let ifaces = IPCollector::collect_interfaces(global_ctx.net_ns.clone(), false).await; - for iface in ifaces { - let is_local = iface.ips.iter().any(|ip| match ip.ip() { - IpAddr::V6(v6) => v6 == preferred_src_ipv6, - IpAddr::V4(_) => false, - }); - if is_local { - tracing::debug!( - ?preferred_src_ipv6, - ifindex = iface.index, - "use preferred IPv6 source for udp hole punch" - ); - return Some(udp::PreferredIpv6Source { - ip: preferred_src_ipv6, - ifindex: iface.index, - }); - } - } - - tracing::debug!( - ?preferred_src_ipv6, - "ignore non-local preferred IPv6 source for udp hole punch" - ); - None -} - -fn connector_addrs_from_request( - req: SendUdpHolePunchPacketRequest, -) -> rpc_types::error::Result<(u16, Vec, Option)> { - let listener_port = u16::try_from(req.listener_port) - .map_err(|_| anyhow::anyhow!("listener_port is out of range: {}", req.listener_port))?; - let mut connector_addrs = req - .connector_addrs - .into_iter() - .map(SocketAddr::from) - .collect::>(); - - if connector_addrs.is_empty() { - connector_addrs.push( - req.connector_addr - .ok_or(anyhow::anyhow!("connector_addr is required"))? - .into(), - ); - } - - let mut deduped = Vec::with_capacity(connector_addrs.len()); - for addr in connector_addrs { - if !deduped.contains(&addr) { - deduped.push(addr); - } - if deduped.len() >= MAX_UDP_HOLE_PUNCH_CONNECTOR_ADDRS { - break; - } - } - - Ok((listener_port, deduped, req.preferred_src_ipv6)) -} - -#[derive(Clone)] -pub struct DirectConnectorManagerRpcServer { - // TODO: this only cache for one src peer, should make it global - global_ctx: ArcGlobalCtx, -} - -#[async_trait::async_trait] -impl DirectConnectorRpc for DirectConnectorManagerRpcServer { - type Controller = BaseController; - - async fn get_ip_list( - &self, - _: BaseController, - _: GetIpListRequest, - ) -> rpc_types::error::Result { - let mut ret = self.global_ctx.get_ip_collector().collect_ip_addrs().await; - ret.listeners = self - .global_ctx - .config - .get_mapped_listeners() - .into_iter() - .chain(self.global_ctx.get_running_listeners()) - .map(Into::into) - .collect(); - remove_easytier_managed_ipv6s(&mut ret, &self.global_ctx); - tracing::trace!( - "get_ip_list: public_ipv4: {:?}, public_ipv6: {:?}, listeners: {:?}", - ret.public_ipv4, - ret.public_ipv6, - ret.listeners - ); - Ok(ret) - } - - async fn send_udp_hole_punch_packet( - &self, - _: BaseController, - req: SendUdpHolePunchPacketRequest, - ) -> rpc_types::error::Result { - let (listener_port, connector_addrs, preferred_src_ipv6) = - connector_addrs_from_request(req)?; - let preferred_src_ipv6 = - local_preferred_src_ipv6(&self.global_ctx, preferred_src_ipv6).await; - - tracing::info!( - ?connector_addrs, - ?preferred_src_ipv6, - listener_port, - "Sending udp hole punch packet" - ); - - // send 3 packets to the connector - for _ in 0..3 { - for connector_addr in &connector_addrs { - let ret = match connector_addr { - SocketAddr::V4(addr) => { - udp::send_v4_hole_punch_packet(listener_port, *addr).await - } - SocketAddr::V6(addr) => { - udp::send_v6_hole_punch_packet(listener_port, *addr, preferred_src_ipv6) - .await - } - }; - if let Err(e) = ret { - tracing::debug!( - ?e, - ?connector_addr, - listener_port, - "send udp hole punch packet failed" - ); - } - } - tokio::time::sleep(std::time::Duration::from_millis(30)).await; - } - Ok(Default::default()) - } -} - -impl DirectConnectorManagerRpcServer { - pub fn new(global_ctx: ArcGlobalCtx) -> Self { - Self { global_ctx } - } -} - -#[cfg(test)] -mod tests { - use std::{collections::BTreeSet, net::SocketAddr}; - - use crate::{ - common::global_ctx::tests::get_mock_global_ctx, - peers::peer_rpc_service::{connector_addrs_from_request, remove_easytier_managed_ipv6s}, - proto::peer_rpc::{GetIpListResponse, SendUdpHolePunchPacketRequest}, - }; - - #[tokio::test] - async fn get_ip_list_sanitizer_removes_managed_ipv6_from_all_sources() { - let global_ctx = get_mock_global_ctx(); - let virtual_ipv6 = "fd00::1/64".parse().unwrap(); - let public_ipv6 = "2001:db8::2/128".parse().unwrap(); - let physical_ipv6: std::net::Ipv6Addr = "2001:db8::3".parse().unwrap(); - let routed_ipv6: cidr::Ipv6Inet = "2001:db8::4/128".parse().unwrap(); - global_ctx.set_ipv6(Some(virtual_ipv6)); - global_ctx.set_public_ipv6_lease(Some(public_ipv6)); - global_ctx.set_public_ipv6_routes(BTreeSet::from([routed_ipv6])); - - let mut ip_list = GetIpListResponse { - public_ipv6: Some(public_ipv6.address().into()), - interface_ipv6s: vec![ - virtual_ipv6.address().into(), - public_ipv6.address().into(), - routed_ipv6.address().into(), - physical_ipv6.into(), - ], - ..Default::default() - }; - - remove_easytier_managed_ipv6s(&mut ip_list, &global_ctx); - - assert_eq!(ip_list.public_ipv6, None); - assert_eq!(ip_list.interface_ipv6s, vec![physical_ipv6.into()]); - } - - #[test] - fn hole_punch_request_prefers_batch_connector_addrs() { - let old_addr: SocketAddr = "[2001:db8::1]:10001".parse().unwrap(); - let first_batch_addr: SocketAddr = "[2001:db8::2]:10002".parse().unwrap(); - let second_batch_addr: SocketAddr = "[2001:db8::3]:10003".parse().unwrap(); - let preferred_src_ipv6: std::net::Ipv6Addr = "2001:db8::4".parse().unwrap(); - - let (listener_port, connector_addrs, preferred_src) = - connector_addrs_from_request(SendUdpHolePunchPacketRequest { - connector_addr: Some(old_addr.into()), - listener_port: 11010, - preferred_src_ipv6: Some(preferred_src_ipv6.into()), - connector_addrs: vec![ - first_batch_addr.into(), - first_batch_addr.into(), - second_batch_addr.into(), - ], - }) - .unwrap(); - - assert_eq!(listener_port, 11010); - assert_eq!(connector_addrs, vec![first_batch_addr, second_batch_addr]); - assert_eq!(preferred_src, Some(preferred_src_ipv6.into())); - } - - #[test] - fn hole_punch_request_falls_back_to_legacy_connector_addr() { - let old_addr: SocketAddr = "[2001:db8::1]:10001".parse().unwrap(); - - let (_, connector_addrs, _) = connector_addrs_from_request(SendUdpHolePunchPacketRequest { - connector_addr: Some(old_addr.into()), - listener_port: 11010, - preferred_src_ipv6: None, - connector_addrs: vec![], - }) - .unwrap(); - - assert_eq!(connector_addrs, vec![old_addr]); - } - - #[test] - fn hole_punch_request_rejects_out_of_range_listener_port() { - let old_addr: SocketAddr = "[2001:db8::1]:10001".parse().unwrap(); - - let ret = connector_addrs_from_request(SendUdpHolePunchPacketRequest { - connector_addr: Some(old_addr.into()), - listener_port: u16::MAX as u32 + 1, - preferred_src_ipv6: None, - connector_addrs: vec![], - }); - - assert!(ret.is_err()); - } -} diff --git a/easytier/src/peers/peer_task.rs b/easytier/src/peers/peer_task.rs deleted file mode 100644 index 47fd7fcb..00000000 --- a/easytier/src/peers/peer_task.rs +++ /dev/null @@ -1,205 +0,0 @@ -use std::{ - result::Result, - sync::{Arc, Mutex, atomic::Ordering}, -}; - -use atomic_shim::AtomicU64; - -use async_trait::async_trait; -use dashmap::DashMap; -use tokio::select; -use tokio::sync::Notify; -use tokio::task::JoinHandle; - -use anyhow::Error; -use tokio_util::task::AbortOnDropHandle; - -use super::peer_manager::PeerManager; - -pub struct ExternalTaskSignal { - version: AtomicU64, - notify: Notify, -} - -impl Default for ExternalTaskSignal { - fn default() -> Self { - Self::new() - } -} - -impl ExternalTaskSignal { - pub fn new() -> Self { - Self { - version: AtomicU64::new(0), - notify: Notify::new(), - } - } - - pub fn notify(&self) { - self.version.fetch_add(1, Ordering::Relaxed); - self.notify.notify_waiters(); - } - - pub fn version(&self) -> u64 { - self.version.load(Ordering::Relaxed) - } - - pub fn notified(&self) -> impl std::future::Future + '_ { - self.notify.notified() - } -} - -#[async_trait] -pub trait PeerTaskLauncher: Send + Sync + Clone + 'static { - type Data; - type CollectPeerItem; - type TaskRet; - - fn new_data(&self, peer_mgr: Arc) -> Self::Data; - async fn collect_peers_need_task(&self, data: &Self::Data) -> Vec; - async fn launch_task( - &self, - data: &Self::Data, - item: Self::CollectPeerItem, - ) -> JoinHandle>; - - async fn all_task_done(&self, _data: &Self::Data) {} - - fn loop_interval_ms(&self) -> u64 { - 5000 - } -} - -pub struct PeerTaskManager { - launcher: Launcher, - main_loop_task: Mutex>>, - run_signal: Arc, - external_signal: Option>, - data: Launcher::Data, -} - -impl PeerTaskManager -where - D: Send + Sync + Clone + 'static, - C: std::fmt::Debug + Send + Sync + Clone + core::hash::Hash + Eq + 'static, - T: Send + 'static, - L: PeerTaskLauncher + 'static, -{ - pub fn new(launcher: L, peer_mgr: Arc) -> Self { - Self::new_with_external_signal(launcher, peer_mgr, None) - } - - pub fn new_with_external_signal( - launcher: L, - peer_mgr: Arc, - external_signal: Option>, - ) -> Self { - let data = launcher.new_data(peer_mgr.clone()); - Self { - launcher, - main_loop_task: Mutex::new(None), - run_signal: Arc::new(Notify::new()), - external_signal, - data, - } - } - - pub fn start(&self) { - let task = AbortOnDropHandle::new(tokio::spawn(Self::main_loop( - self.launcher.clone(), - self.data.clone(), - self.run_signal.clone(), - self.external_signal.clone(), - ))); - self.main_loop_task.lock().unwrap().replace(task); - } - - async fn main_loop( - launcher: L, - data: D, - signal: Arc, - external_signal: Option>, - ) { - let peer_task_map = Arc::new(DashMap::>>::new()); - let mut external_signal_version = external_signal.as_ref().map(|signal| signal.version()); - - loop { - let peers_to_connect = launcher.collect_peers_need_task(&data).await; - - // remove task not in peers_to_connect - let mut to_remove = vec![]; - for item in peer_task_map.iter() { - if !peers_to_connect.contains(item.key()) || item.value().is_finished() { - to_remove.push(item.key().clone()); - } - } - - for key in to_remove { - if let Some((_, task)) = peer_task_map.remove(&key) { - task.abort(); - match task.await { - Ok(Ok(_)) => {} - Ok(Err(task_ret)) => { - tracing::error!(?task_ret, "hole punching task failed"); - } - Err(e) => { - tracing::error!(?e, "hole punching task aborted"); - } - } - } - peer_task_map.shrink_to_fit(); - } - - if !peers_to_connect.is_empty() { - for item in peers_to_connect { - if peer_task_map.contains_key(&item) { - continue; - } - - tracing::debug!(?item, "launch hole punching task"); - peer_task_map.insert( - item.clone(), - AbortOnDropHandle::new(launcher.launch_task(&data, item).await), - ); - } - } else if peer_task_map.is_empty() { - launcher.all_task_done(&data).await; - } - - if let Some(external_signal) = external_signal.as_ref() { - let notified = external_signal.notified(); - tokio::pin!(notified); - let cur_version = external_signal.version(); - if external_signal_version != Some(cur_version) { - external_signal_version = Some(cur_version); - continue; - } - - select! { - _ = tokio::time::sleep(std::time::Duration::from_millis( - launcher.loop_interval_ms(), - )) => {}, - _ = signal.notified() => {}, - _ = &mut notified => { - external_signal_version = Some(external_signal.version()); - } - } - } else { - select! { - _ = tokio::time::sleep(std::time::Duration::from_millis( - launcher.loop_interval_ms(), - )) => {}, - _ = signal.notified() => {} - } - } - } - } - - pub async fn run_immediately(&self) { - self.run_signal.notify_one(); - } - - pub fn data(&self) -> D { - self.data.clone() - } -} diff --git a/easytier/src/peers/rpc_service.rs b/easytier/src/peers/rpc_service.rs deleted file mode 100644 index 941d03f4..00000000 --- a/easytier/src/peers/rpc_service.rs +++ /dev/null @@ -1,300 +0,0 @@ -use std::{ - ops::Deref, - sync::{Arc, Weak}, - time::Duration, -}; - -use crate::{ - proto::{ - api::instance::{ - AclManageRpc, CredentialManageRpc, DumpRouteRequest, DumpRouteResponse, - GenerateCredentialRequest, GenerateCredentialResponse, GetAclStatsRequest, - GetAclStatsResponse, GetForeignNetworkSummaryRequest, GetForeignNetworkSummaryResponse, - GetWhitelistRequest, GetWhitelistResponse, ListCredentialsRequest, - ListCredentialsResponse, ListForeignNetworkRequest, ListForeignNetworkResponse, - ListGlobalForeignNetworkRequest, ListGlobalForeignNetworkResponse, ListPeerRequest, - ListPeerResponse, ListPublicIpv6InfoRequest, ListPublicIpv6InfoResponse, - ListRouteRequest, ListRouteResponse, PeerInfo, PeerManageRpc, RevokeCredentialRequest, - RevokeCredentialResponse, ShowNodeInfoRequest, ShowNodeInfoResponse, - }, - rpc_types::{self, controller::BaseController}, - }, - utils::weak_upgrade, -}; - -use super::peer_manager::PeerManager; - -#[derive(Clone)] -pub struct PeerManagerRpcService { - peer_manager: Weak, -} - -impl PeerManagerRpcService { - pub fn new(peer_manager: Arc) -> Self { - PeerManagerRpcService { - peer_manager: Arc::downgrade(&peer_manager), - } - } - - pub async fn list_peers(peer_manager: &PeerManager) -> Vec { - let mut peers = peer_manager.get_peer_map().list_peers(); - peers.extend( - peer_manager - .get_foreign_network_client() - .get_peer_map() - .list_peers() - .iter(), - ); - let peer_map = peer_manager.get_peer_map(); - let mut peer_infos = Vec::new(); - for peer in peers { - let mut peer_info = PeerInfo { - peer_id: peer, - default_conn_id: peer_map - .get_peer_default_conn_id(peer) - .await - .map(Into::into), - directly_connected_conns: peer_map - .get_directly_connections_by_peer_id(peer) - .into_iter() - .map(Into::into) - .collect(), - ..Default::default() - }; - - if let Some(conns) = peer_map.list_peer_conns(peer).await { - peer_info.conns = conns; - } else if let Some(conns) = peer_manager - .get_foreign_network_client() - .get_peer_map() - .list_peer_conns(peer) - .await - { - peer_info.conns = conns; - } - - peer_infos.push(peer_info); - } - - peer_infos - } -} - -#[async_trait::async_trait] -impl PeerManageRpc for PeerManagerRpcService { - type Controller = BaseController; - async fn list_peer( - &self, - _: BaseController, - _request: ListPeerRequest, // Accept request of type HelloRequest - ) -> Result { - let mut reply = ListPeerResponse::default(); - - let peers = - PeerManagerRpcService::list_peers(weak_upgrade(&self.peer_manager)?.deref()).await; - for peer in peers { - reply.peer_infos.push(peer); - } - - Ok(reply) - } - - async fn list_public_ipv6_info( - &self, - _: BaseController, - _request: ListPublicIpv6InfoRequest, - ) -> Result { - Ok(weak_upgrade(&self.peer_manager)? - .get_local_public_ipv6_info() - .await) - } - - async fn list_route( - &self, - _: BaseController, - _request: ListRouteRequest, // Accept request of type HelloRequest - ) -> Result { - let reply = ListRouteResponse { - routes: weak_upgrade(&self.peer_manager)?.list_routes().await, - }; - Ok(reply) - } - - async fn dump_route( - &self, - _: BaseController, - _request: DumpRouteRequest, // Accept request of type HelloRequest - ) -> Result { - let reply = DumpRouteResponse { - result: weak_upgrade(&self.peer_manager)?.dump_route().await, - }; - Ok(reply) - } - - async fn list_foreign_network( - &self, - _: BaseController, - request: ListForeignNetworkRequest, - ) -> Result { - let reply = weak_upgrade(&self.peer_manager)? - .get_foreign_network_manager() - .list_foreign_networks_with_options(request.include_trusted_keys) - .await; - Ok(reply) - } - - async fn list_global_foreign_network( - &self, - _: BaseController, - _request: ListGlobalForeignNetworkRequest, - ) -> Result { - Ok(weak_upgrade(&self.peer_manager)? - .list_global_foreign_network() - .await) - } - - async fn get_foreign_network_summary( - &self, - _: BaseController, - _request: GetForeignNetworkSummaryRequest, - ) -> Result { - Ok(GetForeignNetworkSummaryResponse { - summary: Some( - weak_upgrade(&self.peer_manager)? - .get_foreign_network_summary() - .await, - ), - }) - } - - async fn show_node_info( - &self, - _: BaseController, - _request: ShowNodeInfoRequest, // Accept request of type HelloRequest - ) -> Result { - Ok(ShowNodeInfoResponse { - node_info: Some(weak_upgrade(&self.peer_manager)?.get_my_info().await), - }) - } -} - -#[async_trait::async_trait] -impl AclManageRpc for PeerManagerRpcService { - type Controller = BaseController; - - async fn get_acl_stats( - &self, - _: BaseController, - _request: GetAclStatsRequest, - ) -> Result { - let acl_stats = weak_upgrade(&self.peer_manager)? - .get_global_ctx() - .get_acl_filter() - .get_stats(); - Ok(GetAclStatsResponse { - acl_stats: Some(acl_stats), - }) - } - - async fn get_whitelist( - &self, - _: BaseController, - _request: GetWhitelistRequest, - ) -> Result { - let global_ctx = weak_upgrade(&self.peer_manager)?.get_global_ctx(); - let tcp_ports = global_ctx.config.get_tcp_whitelist(); - let udp_ports = global_ctx.config.get_udp_whitelist(); - tracing::info!( - "Getting whitelist - TCP: {:?}, UDP: {:?}", - tcp_ports, - udp_ports - ); - Ok(GetWhitelistResponse { - tcp_ports, - udp_ports, - }) - } -} - -#[async_trait::async_trait] -impl CredentialManageRpc for PeerManagerRpcService { - type Controller = BaseController; - - async fn generate_credential( - &self, - _: BaseController, - request: GenerateCredentialRequest, - ) -> Result { - let pm = weak_upgrade(&self.peer_manager)?; - let global_ctx = pm.get_global_ctx(); - - if global_ctx.get_network_identity().network_secret.is_none() { - return Err(rpc_types::error::Error::ExecutionError(anyhow::anyhow!( - "only admin nodes (with network_secret) can generate credentials" - ))); - } - - let ttl = if request.ttl_seconds > 0 { - Duration::from_secs(request.ttl_seconds as u64) - } else { - return Err(rpc_types::error::Error::ExecutionError(anyhow::anyhow!( - "ttl_seconds must be positive" - ))); - }; - - let (id, secret) = global_ctx - .get_credential_manager() - .generate_credential_with_options( - request.groups, - request.allow_relay, - request.allowed_proxy_cidrs, - ttl, - request.credential_id, - request.reusable.unwrap_or(true), - ); - - global_ctx.issue_event(crate::common::global_ctx::GlobalCtxEvent::CredentialChanged); - - Ok(GenerateCredentialResponse { - credential_id: id, - credential_secret: secret, - }) - } - - async fn revoke_credential( - &self, - _: BaseController, - request: RevokeCredentialRequest, - ) -> Result { - let pm = weak_upgrade(&self.peer_manager)?; - let global_ctx = pm.get_global_ctx(); - if global_ctx.get_network_identity().network_secret.is_none() { - return Err(rpc_types::error::Error::ExecutionError(anyhow::anyhow!( - "only admin nodes (with network_secret) can revoke credentials" - ))); - } - - let success = global_ctx - .get_credential_manager() - .revoke_credential(&request.credential_id); - - if success { - global_ctx.issue_event(crate::common::global_ctx::GlobalCtxEvent::CredentialChanged); - } - - Ok(RevokeCredentialResponse { success }) - } - - async fn list_credentials( - &self, - _: BaseController, - _request: ListCredentialsRequest, - ) -> Result { - let pm = weak_upgrade(&self.peer_manager)?; - let global_ctx = pm.get_global_ctx(); - - Ok(ListCredentialsResponse { - credentials: global_ctx.get_credential_manager().list_credentials(), - }) - } -} diff --git a/easytier/src/peers/tests.rs b/easytier/src/peers/tests.rs deleted file mode 100644 index 5987b3f5..00000000 --- a/easytier/src/peers/tests.rs +++ /dev/null @@ -1,1627 +0,0 @@ -use std::sync::Arc; -use std::time::Duration; - -use base64::Engine as _; - -use crate::{ - common::{ - PeerId, - error::Error, - global_ctx::{ - NetworkIdentity, TrustedKeySource, - tests::{get_mock_global_ctx, get_mock_global_ctx_with_network}, - }, - stats_manager::{LabelSet, LabelType, MetricName}, - }, - proto::api::instance::TrustedKeySourcePb, - tunnel::{ - common::tests::wait_for_condition, - packet_def::{PacketType, ZCPacket}, - ring::create_ring_tunnel_pair, - }, -}; - -use super::{ - create_packet_recv_chan, - peer_conn::tests::set_secure_mode_cfg, - peer_manager::{PeerManager, RouteAlgoType}, - peer_map::PeerMap, - peer_session::{PeerSession, PeerSessionStore, SessionKey}, - relay_peer_map::RelayPeerMap, - route_trait::NextHopPolicy, -}; - -pub async fn create_mock_peer_manager() -> Arc { - let (s, _r) = create_packet_recv_chan(); - let peer_mgr = Arc::new(PeerManager::new( - RouteAlgoType::Ospf, - get_mock_global_ctx(), - s, - )); - peer_mgr.run().await.unwrap(); - peer_mgr -} - -pub async fn create_mock_peer_manager_with_name(network_name: String) -> Arc { - let (s, _r) = create_packet_recv_chan(); - let g = - get_mock_global_ctx_with_network(Some(NetworkIdentity::new(network_name, "".to_string()))); - let peer_mgr = Arc::new(PeerManager::new(RouteAlgoType::Ospf, g, s)); - peer_mgr.run().await.unwrap(); - peer_mgr -} - -pub async fn create_mock_peer_manager_secure( - network_name: String, - network_secret: String, -) -> Arc { - let (s, _r) = create_packet_recv_chan(); - let g = - get_mock_global_ctx_with_network(Some(NetworkIdentity::new(network_name, network_secret))); - set_secure_mode_cfg(&g, true); - let peer_mgr = Arc::new(PeerManager::new(RouteAlgoType::Ospf, g, s)); - peer_mgr.run().await.unwrap(); - peer_mgr -} - -fn set_private_mode(peer_mgr: &PeerManager, enabled: bool) { - let global_ctx = peer_mgr.get_global_ctx(); - let mut flags = global_ctx.get_flags(); - flags.private_mode = enabled; - global_ctx.set_flags(flags); -} - -async fn connect_client_and_server( - client: Arc, - server: Arc, -) -> (Result<(), Error>, Result<(), Error>) { - let (client_ring, server_ring) = create_ring_tunnel_pair(); - tokio::join!( - { - let client = client.clone(); - async move { - client.add_client_tunnel(client_ring, false).await?; - Ok(()) - } - }, - { - let server = server.clone(); - async move { server.add_tunnel_as_server(server_ring, true).await } - } - ) -} - -async fn wait_for_foreign_network(server: Arc, network_name: &'static str) { - wait_for_condition( - || { - let server = server.clone(); - async move { - server - .get_foreign_network_manager() - .list_foreign_networks() - .await - .foreign_networks - .contains_key(network_name) - } - }, - Duration::from_secs(10), - ) - .await; -} - -async fn wait_for_foreign_network_peer_count_at_least( - server: Arc, - network_name: &'static str, - min_peer_count: usize, -) { - wait_for_condition( - || { - let server = server.clone(); - async move { - server - .get_foreign_network_manager() - .list_foreign_networks() - .await - .foreign_networks - .get(network_name) - .map(|entry| entry.peers.len() >= min_peer_count) - .unwrap_or(false) - } - }, - Duration::from_secs(10), - ) - .await; -} - -async fn wait_for_public_peers_empty(client: Arc) { - wait_for_condition( - || { - let client = client.clone(); - async move { - client - .get_foreign_network_client() - .list_public_peers() - .await - .is_empty() - } - }, - Duration::from_secs(5), - ) - .await; -} - -pub async fn connect_peer_manager(client: Arc, server: Arc) { - let (a_ring, b_ring) = create_ring_tunnel_pair(); - let a_mgr_copy = client; - tokio::spawn(async move { - a_mgr_copy.add_client_tunnel(a_ring, false).await.unwrap(); - }); - let b_mgr_copy = server; - tokio::spawn(async move { - b_mgr_copy.add_tunnel_as_server(b_ring, true).await.unwrap(); - }); -} - -pub async fn wait_route_appear_with_cost( - peer_mgr: Arc, - node_id: PeerId, - cost: Option, -) -> Result<(), Error> { - let now = std::time::Instant::now(); - while now.elapsed().as_secs() < 5 { - let route = peer_mgr.list_routes().await; - if route - .iter() - .any(|r| r.peer_id == node_id && (cost.is_none() || r.cost == cost.unwrap())) - { - return Ok(()); - } - tokio::time::sleep(std::time::Duration::from_millis(50)).await; - } - Err(Error::NotFound) -} - -pub async fn wait_route_appear( - peer_mgr: Arc, - target_peer: Arc, -) -> Result<(), Error> { - wait_route_appear_with_cost(peer_mgr.clone(), target_peer.my_peer_id(), None).await?; - wait_route_appear_with_cost(target_peer, peer_mgr.my_peer_id(), None).await -} - -fn metric_value(peer_mgr: &PeerManager, metric: MetricName, network_name: &str) -> u64 { - peer_mgr - .get_global_ctx() - .stats_manager() - .get_metric( - metric, - &LabelSet::new().with_label_type(LabelType::NetworkName(network_name.to_string())), - ) - .map(|metric| metric.value) - .unwrap_or(0) -} - -#[tokio::test] -async fn foreign_mgr_stress_test() { - const FOREIGN_NETWORK_COUNT: i32 = 20; - const PEER_PER_NETWORK: i32 = 3; - const PUBLIC_PEER_COUNT: i32 = 3; - - let mut public_peers = Vec::new(); - for _ in 0..PUBLIC_PEER_COUNT { - public_peers.push(create_mock_peer_manager().await); - } - connect_peer_manager(public_peers[0].clone(), public_peers[1].clone()).await; - connect_peer_manager(public_peers[0].clone(), public_peers[2].clone()).await; - connect_peer_manager(public_peers[1].clone(), public_peers[2].clone()).await; - - let mut foreigns = Vec::new(); - - for i in 0..FOREIGN_NETWORK_COUNT { - let mut peers = Vec::new(); - - let name = format!("foreign-network-test-{}", i); - - for _ in 0..PEER_PER_NETWORK { - let mgr = create_mock_peer_manager_with_name(name.clone()).await; - let public_peer_idx = rand::random::() % public_peers.len(); - connect_peer_manager(mgr.clone(), public_peers[public_peer_idx].clone()).await; - peers.push(mgr); - } - - foreigns.push(peers); - } - - for _ in 0..5 { - for i in 0..PUBLIC_PEER_COUNT { - let p = public_peers[i as usize].clone(); - println!( - "public peer {} routes: {:?}, global_foreign_network: {:?}, peers: {:?}", - i, - p.list_routes().await, - p.list_global_foreign_network().await.foreign_networks.len(), - p.get_peer_map().list_peers() - ); - } - tokio::time::sleep(std::time::Duration::from_secs(1)).await; - - let new_peer = create_mock_peer_manager().await; - connect_peer_manager(new_peer.clone(), public_peers[0].clone()).await; - while let Err(e) = wait_route_appear(public_peers[1].clone(), new_peer.clone()).await { - println!("wait route ret: {:?}", e); - } - } -} - -#[tokio::test] -async fn relay_peer_map_secure_session_decrypt() { - let (s, _r) = create_packet_recv_chan(); - let ctx = get_mock_global_ctx_with_network(Some(NetworkIdentity::new( - "net1".to_string(), - "sec1".to_string(), - ))); - set_secure_mode_cfg(&ctx, true); - let peer_map = Arc::new(PeerMap::new(s, ctx.clone(), 10)); - let store = Arc::new(PeerSessionStore::new()); - let relay_map = RelayPeerMap::new(peer_map, None, ctx.clone(), 10, store.clone()); - - let algo = ctx.get_flags().encryption_algorithm.clone(); - let root_key = [7u8; 32]; - let session = Arc::new(PeerSession::new( - 20, - root_key, - 1, - 1, - algo.clone(), - algo.clone(), - None, - )); - let key = SessionKey::new(ctx.get_network_identity().network_name, 20); - store.insert_session(key.clone(), session.clone()); - - relay_map - .ensure_session(20, NextHopPolicy::LeastHop) - .await - .unwrap(); - assert!(relay_map.has_session(20)); - - let mut packet = ZCPacket::new_with_payload(b"relay-hello"); - packet.fill_peer_manager_hdr(20, 10, PacketType::Data as u8); - session.encrypt_payload(20, 10, &mut packet).unwrap(); - assert!(relay_map.decrypt_if_needed(&mut packet).await.unwrap()); - assert_eq!(packet.payload(), b"relay-hello"); -} - -#[tokio::test] -async fn private_mode_allows_foreign_network_with_same_secret() { - let server = create_mock_peer_manager_secure("public".to_string(), "shared".to_string()).await; - let client = - create_mock_peer_manager_secure("tenant-a".to_string(), "shared".to_string()).await; - set_private_mode(&server, true); - - let (client_ret, server_ret) = connect_client_and_server(client, server.clone()).await; - - assert!(client_ret.is_ok(), "client should connect in private mode"); - assert!( - server_ret.is_ok(), - "server should accept foreign network with matching secret: {:?}", - server_ret - ); - wait_for_foreign_network(server, "tenant-a").await; -} - -#[tokio::test] -async fn private_mode_rejects_foreign_network_with_different_secret() { - let server = create_mock_peer_manager_secure("public".to_string(), "shared".to_string()).await; - let client = create_mock_peer_manager_secure("tenant-a".to_string(), "other".to_string()).await; - set_private_mode(&server, true); - - let (client_ret, server_ret) = connect_client_and_server(client.clone(), server.clone()).await; - - assert!( - server_ret.is_err(), - "server should reject foreign network with mismatched secret in private mode" - ); - let _ = client_ret; - wait_for_public_peers_empty(client).await; - assert!( - server - .get_foreign_network_manager() - .list_foreign_networks() - .await - .foreign_networks - .is_empty() - ); -} - -#[tokio::test] -async fn private_mode_allows_trusted_foreign_credential() { - let server = create_mock_peer_manager_secure("public".to_string(), "shared".to_string()).await; - let admin = create_mock_peer_manager_secure("tenant-a".to_string(), "shared".to_string()).await; - set_private_mode(&server, true); - - let (_cred_id, cred_secret) = admin - .get_global_ctx() - .get_credential_manager() - .generate_credential(vec![], false, vec![], Duration::from_secs(3600)); - - let privkey_bytes: [u8; 32] = base64::engine::general_purpose::STANDARD - .decode(&cred_secret) - .unwrap() - .try_into() - .unwrap(); - let private = x25519_dalek::StaticSecret::from(privkey_bytes); - let public = x25519_dalek::PublicKey::from(&private); - let credential = create_mock_peer_manager_credential("tenant-a".to_string(), &private).await; - - connect_peer_manager(admin.clone(), server.clone()).await; - wait_for_condition( - || { - let server = server.clone(); - let pubkey = public.as_bytes().to_vec(); - async move { - server - .get_foreign_network_manager() - .list_foreign_networks_with_options(true) - .await - .foreign_networks - .get("tenant-a") - .map(|entry| { - entry.trusted_keys.iter().any(|trusted_key| { - trusted_key.pubkey == pubkey - && trusted_key.source == TrustedKeySourcePb::OspfCredential as i32 - }) - }) - .unwrap_or(false) - } - }, - Duration::from_secs(10), - ) - .await; - - let (client_ret, server_ret) = connect_client_and_server(credential, server.clone()).await; - - assert!( - client_ret.is_ok(), - "trusted foreign credential client should connect in private mode" - ); - assert!( - server_ret.is_ok(), - "server should allow trusted foreign credential in private mode: {:?}", - server_ret - ); - wait_for_foreign_network_peer_count_at_least(server, "tenant-a", 2).await; -} - -#[tokio::test] -async fn private_mode_rejects_untrusted_foreign_credential() { - let server = create_mock_peer_manager_secure("public".to_string(), "shared".to_string()).await; - let admin = create_mock_peer_manager_secure("tenant-a".to_string(), "shared".to_string()).await; - set_private_mode(&server, true); - - let random_private = x25519_dalek::StaticSecret::random_from_rng(rand::rngs::OsRng); - let unknown_credential = - create_mock_peer_manager_credential("tenant-a".to_string(), &random_private).await; - - connect_peer_manager(admin.clone(), server.clone()).await; - wait_for_foreign_network(server.clone(), "tenant-a").await; - - let (client_ret, server_ret) = - connect_client_and_server(unknown_credential, server.clone()).await; - - let _ = client_ret; - assert!( - server_ret.is_err(), - "server should reject untrusted foreign credential in private mode" - ); - wait_for_condition( - || { - let server = server.clone(); - async move { - server - .get_foreign_network_manager() - .list_foreign_networks() - .await - .foreign_networks - .get("tenant-a") - .map(|entry| entry.peers.len() == 1) - .unwrap_or(false) - } - }, - Duration::from_secs(10), - ) - .await; -} - -#[tokio::test] -async fn relay_peer_map_retry_backoff_and_evict() { - let (s, _r) = create_packet_recv_chan(); - let ctx_secure = get_mock_global_ctx(); - set_secure_mode_cfg(&ctx_secure, true); - let peer_map = Arc::new(PeerMap::new(s, ctx_secure.clone(), 10)); - let relay_map = RelayPeerMap::new( - peer_map, - None, - ctx_secure.clone(), - 10, - Arc::new(PeerSessionStore::new()), - ); - - let ret = relay_map - .handshake_session(20, NextHopPolicy::LeastHop, None) - .await; - assert!(ret.is_err()); - assert!(relay_map.failure_count(20).unwrap_or(0) >= 1); - assert!(relay_map.is_backoff_active(20)); - - let (s2, _r2) = create_packet_recv_chan(); - let ctx_plain = get_mock_global_ctx(); - let peer_map_plain = Arc::new(PeerMap::new(s2, ctx_plain.clone(), 30)); - let relay_map_plain = RelayPeerMap::new( - peer_map_plain, - None, - ctx_plain.clone(), - 30, - Arc::new(PeerSessionStore::new()), - ); - - let mut pkt = ZCPacket::new_with_payload(b"evict"); - pkt.fill_peer_manager_hdr(30, 40, PacketType::Data as u8); - let _ = relay_map_plain - .send_msg(pkt, 40, NextHopPolicy::LeastHop) - .await; - assert!(relay_map_plain.has_state(40)); - relay_map_plain.evict_idle_sessions(Duration::from_millis(0)); - assert!(!relay_map_plain.has_state(40)); -} - -#[tokio::test] -async fn relay_peer_map_pending_packet_buffer() { - // Verify that packets sent during handshake are buffered (not dropped), - // and flushed after handshake completes. - let (s, _r) = create_packet_recv_chan(); - let ctx = get_mock_global_ctx_with_network(Some(NetworkIdentity::new( - "net1".to_string(), - "sec1".to_string(), - ))); - set_secure_mode_cfg(&ctx, true); - let peer_map = Arc::new(PeerMap::new(s, ctx.clone(), 10)); - let store = Arc::new(PeerSessionStore::new()); - let relay_map = RelayPeerMap::new(peer_map, None, ctx.clone(), 10, store.clone()); - - // Send multiple packets while no session exists (handshake will fail, but packets should be buffered) - for i in 0..5u8 { - let mut pkt = ZCPacket::new_with_payload(&[i]); - pkt.fill_peer_manager_hdr(10, 20, PacketType::Data as u8); - let _ = relay_map.send_msg(pkt, 20, NextHopPolicy::LeastHop).await; - } - - // Verify packets were buffered - assert_eq!( - relay_map - .pending_packets - .get(&20) - .map(|v| v.len()) - .unwrap_or(0), - 5, - "5 packets should be buffered during handshake" - ); - - // Verify buffer respects capacity limit - for i in 0..50u8 { - let mut pkt = ZCPacket::new_with_payload(&[i]); - pkt.fill_peer_manager_hdr(10, 20, PacketType::Data as u8); - let _ = relay_map.send_msg(pkt, 20, NextHopPolicy::LeastHop).await; - } - - let buffered = relay_map - .pending_packets - .get(&20) - .map(|v| v.len()) - .unwrap_or(0); - assert!( - buffered <= 32, - "buffer should not exceed MAX_PENDING_PACKETS_PER_PEER, got {buffered}" - ); - - // Verify remove_peer clears pending packets - relay_map.remove_peer(20); - assert_eq!( - relay_map - .pending_packets - .get(&20) - .map(|v| v.len()) - .unwrap_or(0), - 0, - "pending packets should be cleared on peer removal" - ); -} - -#[tokio::test] -async fn relay_peer_map_pending_packets_flushed_on_handshake_success() { - // Test that pending packets are flushed after handshake succeeds. - // We pre-populate the buffer, then run handshake, and verify it's cleared. - let peer_a = create_mock_peer_manager_secure("net1".to_string(), "sec1".to_string()).await; - let peer_b = create_mock_peer_manager_secure("net1".to_string(), "sec1".to_string()).await; - let peer_c = create_mock_peer_manager_secure("net1".to_string(), "sec1".to_string()).await; - - connect_peer_manager(peer_a.clone(), peer_b.clone()).await; - connect_peer_manager(peer_b.clone(), peer_c.clone()).await; - - let peer_a_id = peer_a.my_peer_id(); - let peer_c_id = peer_c.my_peer_id(); - - // Wait for routes to propagate - wait_for_condition( - || { - let peer_a = peer_a.clone(); - let peer_c = peer_c.clone(); - async move { wait_route_appear(peer_a.clone(), peer_c).await.is_ok() } - }, - Duration::from_secs(10), - ) - .await; - - // Wait for noise_static_pubkey to be available on both sides - wait_for_condition( - || { - let peer_a = peer_a.clone(); - async move { - peer_a - .get_peer_map() - .get_route_peer_info(peer_c_id) - .await - .map(|info| !info.noise_static_pubkey.is_empty()) - .unwrap_or(false) - } - }, - Duration::from_secs(10), - ) - .await; - - let relay_a = peer_a.get_relay_peer_map(); - - // Pre-populate pending packets buffer (simulating what send_msg does during handshake) - for i in 0..3u8 { - let mut pkt = ZCPacket::new_with_payload(&[i]); - pkt.fill_peer_manager_hdr(peer_a_id, peer_c_id, PacketType::Data as u8); - relay_a - .pending_packets - .entry(peer_c_id) - .or_default() - .push((pkt, NextHopPolicy::LeastHop)); - } - - assert_eq!( - relay_a - .pending_packets - .get(&peer_c_id) - .map(|v| v.len()) - .unwrap_or(0), - 3, - "3 packets should be in the buffer" - ); - - // Run handshake — on success it should flush the buffer - relay_a - .handshake_session(peer_c_id, NextHopPolicy::LeastHop, None) - .await - .unwrap(); - - // Verify session established and buffer cleared - assert!(relay_a.has_session(peer_c_id)); - assert_eq!( - relay_a - .pending_packets - .get(&peer_c_id) - .map(|v| v.len()) - .unwrap_or(0), - 0, - "pending packets should be flushed after successful handshake" - ); -} - -#[tokio::test] -async fn relay_peer_map_real_link_handshake_success() { - let peer_a = create_mock_peer_manager_secure("net1".to_string(), "sec1".to_string()).await; - let peer_b = create_mock_peer_manager_secure("net1".to_string(), "sec1".to_string()).await; - let peer_c = create_mock_peer_manager_secure("net1".to_string(), "sec1".to_string()).await; - - connect_peer_manager(peer_a.clone(), peer_b.clone()).await; - connect_peer_manager(peer_b.clone(), peer_c.clone()).await; - - let peer_a_id = peer_a.my_peer_id(); - let peer_b_id = peer_b.my_peer_id(); - let peer_c_id = peer_c.my_peer_id(); - let a_control_tx_before = metric_value(&peer_a, MetricName::TrafficControlBytesTx, "net1"); - let a_control_rx_before = metric_value(&peer_a, MetricName::TrafficControlBytesRx, "net1"); - let c_control_tx_before = metric_value(&peer_c, MetricName::TrafficControlBytesTx, "net1"); - let c_control_rx_before = metric_value(&peer_c, MetricName::TrafficControlBytesRx, "net1"); - - wait_for_condition( - || { - let peer_a = peer_a.clone(); - let peer_c = peer_c.clone(); - async move { wait_route_appear(peer_a.clone(), peer_c).await.is_ok() } - }, - Duration::from_secs(10), - ) - .await; - - wait_for_condition( - || { - let peer_a = peer_a.clone(); - async move { - peer_a - .get_peer_map() - .get_gateway_peer_id(peer_c_id, NextHopPolicy::LeastHop) - .await - == Some(peer_b_id) - } - }, - Duration::from_secs(5), - ) - .await; - - wait_for_condition( - || { - let peer_a = peer_a.clone(); - async move { - peer_a - .get_peer_map() - .get_route_peer_info(peer_c_id) - .await - .map(|info| !info.noise_static_pubkey.is_empty()) - .unwrap_or(false) - } - }, - Duration::from_secs(10), - ) - .await; - - let relay_a = peer_a.get_relay_peer_map(); - let relay_c = peer_c.get_relay_peer_map(); - - relay_a - .handshake_session(peer_c_id, NextHopPolicy::LeastHop, None) - .await - .unwrap(); - - wait_for_condition( - || { - let relay_a = relay_a.clone(); - async move { relay_a.has_session(peer_c_id) } - }, - Duration::from_secs(5), - ) - .await; - - wait_for_condition( - || { - let relay_c = relay_c.clone(); - async move { relay_c.has_session(peer_a_id) } - }, - Duration::from_secs(5), - ) - .await; - - assert!(metric_value(&peer_a, MetricName::TrafficControlBytesTx, "net1") > a_control_tx_before); - assert!(metric_value(&peer_a, MetricName::TrafficControlBytesRx, "net1") > a_control_rx_before); - assert!(metric_value(&peer_c, MetricName::TrafficControlBytesTx, "net1") > c_control_tx_before); - assert!(metric_value(&peer_c, MetricName::TrafficControlBytesRx, "net1") > c_control_rx_before); -} - -#[tokio::test] -async fn relay_peer_map_responder_rejects_mismatched_pubkey() { - // Create three peers: A -> B -> C - let peer_a = create_mock_peer_manager_secure("net1".to_string(), "sec1".to_string()).await; - let peer_b = create_mock_peer_manager_secure("net1".to_string(), "sec1".to_string()).await; - let peer_c = create_mock_peer_manager_secure("net1".to_string(), "sec1".to_string()).await; - - connect_peer_manager(peer_a.clone(), peer_b.clone()).await; - connect_peer_manager(peer_b.clone(), peer_c.clone()).await; - - let peer_a_id = peer_a.my_peer_id(); - let peer_c_id = peer_c.my_peer_id(); - - // Wait for routes to propagate - wait_for_condition( - || { - let peer_a = peer_a.clone(); - let peer_c = peer_c.clone(); - async move { wait_route_appear(peer_a.clone(), peer_c).await.is_ok() } - }, - Duration::from_secs(10), - ) - .await; - - // Wait for noise_static_pubkey to be available - wait_for_condition( - || { - let peer_a = peer_a.clone(); - async move { - peer_a - .get_peer_map() - .get_route_peer_info(peer_c_id) - .await - .map(|info| !info.noise_static_pubkey.is_empty()) - .unwrap_or(false) - } - }, - Duration::from_secs(10), - ) - .await; - - // Get the original correct pubkey to verify it exists - let original_info = peer_a - .get_peer_map() - .get_route_peer_info(peer_c_id) - .await - .expect("should have route info for peer_c"); - assert!( - !original_info.noise_static_pubkey.is_empty(), - "noise_static_pubkey should be present" - ); - - // Attempt handshake - this should succeed because pubkeys match - let relay_a = peer_a.get_relay_peer_map(); - let result = relay_a - .handshake_session(peer_c_id, NextHopPolicy::LeastHop, None) - .await; - - // The handshake should succeed because the pubkeys match - assert!( - result.is_ok(), - "handshake should succeed with matching pubkeys" - ); - - // Verify session was established on both sides - wait_for_condition( - || { - let relay_a = relay_a.clone(); - async move { relay_a.has_session(peer_c_id) } - }, - Duration::from_secs(5), - ) - .await; - - let relay_c = peer_c.get_relay_peer_map(); - wait_for_condition( - || { - let relay_c = relay_c.clone(); - async move { relay_c.has_session(peer_a_id) } - }, - Duration::from_secs(5), - ) - .await; -} - -#[tokio::test] -async fn relay_peer_map_remove_peer() { - let (s, _r) = create_packet_recv_chan(); - let ctx = get_mock_global_ctx_with_network(Some(NetworkIdentity::new( - "net1".to_string(), - "sec1".to_string(), - ))); - set_secure_mode_cfg(&ctx, true); - let peer_map = Arc::new(PeerMap::new(s, ctx.clone(), 10)); - let store = Arc::new(PeerSessionStore::new()); - let relay_map = RelayPeerMap::new(peer_map, None, ctx.clone(), 10, store.clone()); - - let peer_1: PeerId = 100; - - // Add session for peer_1 - let root_key = [1u8; 32]; - let session = Arc::new(PeerSession::new( - peer_1, - root_key, - 1, - 0, - "aes-256-gcm".to_string(), - "aes-256-gcm".to_string(), - None, - )); - let key = SessionKey::new(ctx.get_network_name(), peer_1); - store.insert_session(key.clone(), session); - - assert!(store.get(&key).is_some()); - - // Remove the peer relay state - relay_map.remove_peer(peer_1); - - // Session should still be in the store (lifecycle is independent of relay state) - assert!( - store.get(&key).is_some(), - "session should persist after relay peer removal" - ); -} - -/// Test bidirectional handshake race resolution. -/// When both peers simultaneously initiate handshake, the one with smaller peer_id -/// should become initiator, and the other should yield and become responder. -#[tokio::test] -async fn relay_peer_map_bidirectional_handshake_race() { - // Create three peers: A -> B -> C - let peer_a = create_mock_peer_manager_secure("net1".to_string(), "sec1".to_string()).await; - let peer_b = create_mock_peer_manager_secure("net1".to_string(), "sec1".to_string()).await; - let peer_c = create_mock_peer_manager_secure("net1".to_string(), "sec1".to_string()).await; - - connect_peer_manager(peer_a.clone(), peer_b.clone()).await; - connect_peer_manager(peer_b.clone(), peer_c.clone()).await; - - let peer_a_id = peer_a.my_peer_id(); - let peer_c_id = peer_c.my_peer_id(); - - // Wait for routes to propagate - wait_for_condition( - || { - let peer_a = peer_a.clone(); - let peer_c = peer_c.clone(); - async move { wait_route_appear(peer_a.clone(), peer_c).await.is_ok() } - }, - Duration::from_secs(10), - ) - .await; - - // Wait for noise_static_pubkey to be available - wait_for_condition( - || { - let peer_a = peer_a.clone(); - async move { - peer_a - .get_peer_map() - .get_route_peer_info(peer_c_id) - .await - .map(|info| !info.noise_static_pubkey.is_empty()) - .unwrap_or(false) - } - }, - Duration::from_secs(10), - ) - .await; - - wait_for_condition( - || { - let peer_c = peer_c.clone(); - async move { - peer_c - .get_peer_map() - .get_route_peer_info(peer_a_id) - .await - .map(|info| !info.noise_static_pubkey.is_empty()) - .unwrap_or(false) - } - }, - Duration::from_secs(10), - ) - .await; - - // Simulate bidirectional handshake race by having both sides initiate simultaneously - let relay_a = peer_a.get_relay_peer_map(); - let relay_c = peer_c.get_relay_peer_map(); - - // Both sides initiate handshake at the same time - let handle_a = tokio::spawn({ - let relay_a = relay_a.clone(); - async move { - relay_a - .handshake_session(peer_c_id, NextHopPolicy::LeastHop, None) - .await - } - }); - - let handle_c = tokio::spawn({ - let relay_c = relay_c.clone(); - async move { - relay_c - .handshake_session(peer_a_id, NextHopPolicy::LeastHop, None) - .await - } - }); - - // Wait for both handshakes to complete - let (result_a, result_c) = tokio::join!(handle_a, handle_c); - - // At least one should succeed (the initiator with smaller peer_id) - // Both could succeed if race resolution worked correctly - tracing::info!( - ?peer_a_id, - ?peer_c_id, - ?result_a, - ?result_c, - "bidirectional handshake results" - ); - - // Wait for sessions to be established - wait_for_condition( - || { - let relay_a = relay_a.clone(); - async move { relay_a.has_session(peer_c_id) } - }, - Duration::from_secs(5), - ) - .await; - - wait_for_condition( - || { - let relay_c = relay_c.clone(); - async move { relay_c.has_session(peer_a_id) } - }, - Duration::from_secs(5), - ) - .await; - - // Both sides should have sessions after race resolution - assert!( - relay_a.has_session(peer_c_id), - "peer_a should have session with peer_c" - ); - assert!( - relay_c.has_session(peer_a_id), - "peer_c should have session with peer_a" - ); -} - -/// Helper: create a secure peer manager for a credential node. -/// Uses the given X25519 private key as the Noise static key, with no network_secret. -pub async fn create_mock_peer_manager_credential( - network_name: String, - private_key: &x25519_dalek::StaticSecret, -) -> Arc { - use crate::common::config::NetworkIdentity; - use crate::proto::common::SecureModeConfig; - use base64::Engine; - use base64::engine::general_purpose::STANDARD as BASE64_STANDARD; - - let (s, _r) = create_packet_recv_chan(); - let g = get_mock_global_ctx_with_network(Some(NetworkIdentity::new_credential(network_name))); - - let public = x25519_dalek::PublicKey::from(private_key); - g.config.set_secure_mode(Some(SecureModeConfig { - enabled: true, - local_private_key: Some(BASE64_STANDARD.encode(private_key.as_bytes())), - local_public_key: Some(BASE64_STANDARD.encode(public.as_bytes())), - })); - - let peer_mgr = Arc::new(PeerManager::new(RouteAlgoType::Ospf, g, s)); - peer_mgr.run().await.unwrap(); - peer_mgr -} - -/// Test: credential node joins a 2-admin network and routes appear. -/// Topology: Admin_A -- Credential_C, Admin_A -- Admin_B -/// Credential node connects to the admin that generated the credential. -#[tokio::test] -async fn credential_node_joins_network() { - let admin_a = create_mock_peer_manager_secure("net1".to_string(), "secret".to_string()).await; - let admin_b = create_mock_peer_manager_secure("net1".to_string(), "secret".to_string()).await; - - // Generate credential on admin_a - let (_cred_id, cred_secret) = admin_a - .get_global_ctx() - .get_credential_manager() - .generate_credential( - vec!["guest".to_string()], - false, - vec![], - std::time::Duration::from_secs(3600), - ); - - // Create credential node using the generated key - let privkey_bytes: [u8; 32] = base64::engine::general_purpose::STANDARD - .decode(&cred_secret) - .unwrap() - .try_into() - .unwrap(); - let private = x25519_dalek::StaticSecret::from(privkey_bytes); - let cred_c = create_mock_peer_manager_credential("net1".to_string(), &private).await; - - // Connect admins first - connect_peer_manager(admin_a.clone(), admin_b.clone()).await; - - // Admin A and B should discover each other - wait_route_appear(admin_a.clone(), admin_b.clone()) - .await - .unwrap(); - - // Now connect credential node to admin A (credential as client) - connect_peer_manager(cred_c.clone(), admin_a.clone()).await; - - // Credential node C should be reachable from admin B (via A) - let cred_c_id = cred_c.my_peer_id(); - wait_for_condition( - || { - let admin_b = admin_b.clone(); - async move { - admin_b - .list_routes() - .await - .iter() - .any(|r| r.peer_id == cred_c_id) - } - }, - Duration::from_secs(10), - ) - .await; - - // Credential node C should see admin B - wait_for_condition( - || { - let cred_c = cred_c.clone(); - let admin_b_id = admin_b.my_peer_id(); - async move { - cred_c - .list_routes() - .await - .iter() - .any(|r| r.peer_id == admin_b_id) - } - }, - Duration::from_secs(10), - ) - .await; -} - -/// Test: credential node is rejected when its pubkey is not in any admin's trusted list. -/// Topology: Admin_A -- Unknown_B (random key, not in trusted list) -#[tokio::test] -async fn unknown_credential_node_rejected() { - let admin_a = create_mock_peer_manager_secure("net1".to_string(), "secret".to_string()).await; - - // Create a credential node with a random key (NOT generated by admin) - let random_private = x25519_dalek::StaticSecret::random_from_rng(rand::rngs::OsRng); - let unknown_c = create_mock_peer_manager_credential("net1".to_string(), &random_private).await; - - // Try to connect: C -> A (unknown credential as client, admin as server) - connect_peer_manager(unknown_c.clone(), admin_a.clone()).await; - - // The handshake should fail so the connection won't establish. - // Wait a bit and verify no route appears. - tokio::time::sleep(Duration::from_secs(3)).await; - - let routes = admin_a.list_routes().await; - assert!( - !routes.iter().any(|r| r.peer_id == unknown_c.my_peer_id()), - "unknown credential node should NOT appear in admin's routes" - ); -} - -/// Test: after revocation, the credential node disappears from routes. -/// Topology: Admin_A -- Credential_C, Admin_A -- Admin_B -/// After revocation on A, C should be removed from B's route table. -#[tokio::test] -async fn credential_revocation_removes_from_routes() { - let admin_a = create_mock_peer_manager_secure("net1".to_string(), "secret".to_string()).await; - let admin_b = create_mock_peer_manager_secure("net1".to_string(), "secret".to_string()).await; - - let (cred_id, cred_secret) = admin_a - .get_global_ctx() - .get_credential_manager() - .generate_credential(vec![], false, vec![], std::time::Duration::from_secs(3600)); - - let privkey_bytes: [u8; 32] = base64::engine::general_purpose::STANDARD - .decode(&cred_secret) - .unwrap() - .try_into() - .unwrap(); - let private = x25519_dalek::StaticSecret::from(privkey_bytes); - let cred_c = create_mock_peer_manager_credential("net1".to_string(), &private).await; - - // Connect: A -- B, C -> A (credential node as client, admin as server) - connect_peer_manager(admin_a.clone(), admin_b.clone()).await; - connect_peer_manager(cred_c.clone(), admin_a.clone()).await; - - // Wait for credential node to appear in admin_b's routes - let cred_c_id = cred_c.my_peer_id(); - wait_for_condition( - || { - let admin_b = admin_b.clone(); - async move { - admin_b - .list_routes() - .await - .iter() - .any(|r| r.peer_id == cred_c_id) - } - }, - Duration::from_secs(10), - ) - .await; - - // Now revoke the credential - assert!( - admin_a - .get_global_ctx() - .get_credential_manager() - .revoke_credential(&cred_id) - ); - // Issue event to trigger OSPF sync - admin_a - .get_global_ctx() - .issue_event(crate::common::global_ctx::GlobalCtxEvent::CredentialChanged); - - // Wait for credential node to disappear from admin_b's routes - wait_for_condition( - || { - let admin_b = admin_b.clone(); - async move { - !admin_b - .list_routes() - .await - .iter() - .any(|r| r.peer_id == cred_c_id) - } - }, - Duration::from_secs(15), - ) - .await; -} - -#[tokio::test] -async fn credential_expiry_disconnects_from_all_admins() { - let admin_a = create_mock_peer_manager_secure("net1".to_string(), "secret".to_string()).await; - let admin_b = create_mock_peer_manager_secure("net1".to_string(), "secret".to_string()).await; - - connect_peer_manager(admin_a.clone(), admin_b.clone()).await; - wait_route_appear(admin_a.clone(), admin_b.clone()) - .await - .unwrap(); - - let (_cred_id, cred_secret) = admin_a - .get_global_ctx() - .get_credential_manager() - .generate_credential(vec![], false, vec![], std::time::Duration::from_secs(2)); - - admin_a - .get_global_ctx() - .issue_event(crate::common::global_ctx::GlobalCtxEvent::CredentialChanged); - - let privkey_bytes: [u8; 32] = base64::engine::general_purpose::STANDARD - .decode(&cred_secret) - .unwrap() - .try_into() - .unwrap(); - let private = x25519_dalek::StaticSecret::from(privkey_bytes); - let cred_c = create_mock_peer_manager_credential("net1".to_string(), &private).await; - let cred_c_id = cred_c.my_peer_id(); - - connect_peer_manager(cred_c.clone(), admin_a.clone()).await; - - wait_for_condition( - || { - let admin_b = admin_b.clone(); - async move { - admin_b - .list_routes() - .await - .iter() - .any(|r| r.peer_id == cred_c_id) - } - }, - Duration::from_secs(10), - ) - .await; - - connect_peer_manager(cred_c.clone(), admin_b.clone()).await; - - wait_for_condition( - || { - let admin_b = admin_b.clone(); - async move { - admin_b - .get_peer_map() - .list_peer_conns(cred_c_id) - .await - .is_some_and(|conns| !conns.is_empty()) - } - }, - Duration::from_secs(10), - ) - .await; - - tokio::time::sleep(Duration::from_secs(3)).await; - - wait_for_condition( - || { - let admin_b = admin_b.clone(); - async move { - !admin_b - .list_routes() - .await - .iter() - .any(|r| r.peer_id == cred_c_id) - } - }, - Duration::from_secs(20), - ) - .await; - - wait_for_condition( - || { - let admin_b = admin_b.clone(); - async move { - admin_b - .get_peer_map() - .list_peer_conns(cred_c_id) - .await - .is_none_or(|conns| conns.is_empty()) - } - }, - Duration::from_secs(20), - ) - .await; -} - -/// Test: admin node with credential — credential node gets group assignment. -/// Verify that the credential node's groups appear in the OSPF sync data. -#[tokio::test] -async fn credential_node_group_assignment() { - let admin_a = create_mock_peer_manager_secure("net1".to_string(), "secret".to_string()).await; - let admin_b = create_mock_peer_manager_secure("net1".to_string(), "secret".to_string()).await; - - let (_cred_id, cred_secret) = admin_a - .get_global_ctx() - .get_credential_manager() - .generate_credential( - vec!["guest".to_string(), "limited".to_string()], - false, - vec![], - std::time::Duration::from_secs(3600), - ); - - let privkey_bytes: [u8; 32] = base64::engine::general_purpose::STANDARD - .decode(&cred_secret) - .unwrap() - .try_into() - .unwrap(); - let private = x25519_dalek::StaticSecret::from(privkey_bytes); - let cred_c = create_mock_peer_manager_credential("net1".to_string(), &private).await; - - connect_peer_manager(admin_a.clone(), admin_b.clone()).await; - connect_peer_manager(cred_c.clone(), admin_a.clone()).await; - - // Wait for credential node route to appear on admin_b (via OSPF through admin_a) - let cred_c_id = cred_c.my_peer_id(); - wait_for_condition( - || { - let admin_b = admin_b.clone(); - async move { - admin_b - .list_routes() - .await - .iter() - .any(|r| r.peer_id == cred_c_id) - } - }, - Duration::from_secs(10), - ) - .await; - - // Verify the credential node's groups are assigned via OSPF on admin_b - // (admin_b gets the groups from admin_a's TrustedCredentialPubkey via OSPF sync) - wait_for_condition( - || { - let admin_b = admin_b.clone(); - async move { - let g = admin_b.get_route().get_peer_groups(cred_c_id); - g.contains(&"guest".to_string()) && g.contains(&"limited".to_string()) - } - }, - Duration::from_secs(10), - ) - .await; -} - -#[tokio::test] -async fn credential_node_connected_via_admin_b_trusts_admin_a_groups() { - use crate::proto::acl::{Acl, AclV1, GroupIdentity, GroupInfo}; - - let admin_a = create_mock_peer_manager_secure("net1".to_string(), "secret".to_string()).await; - let admin_b = create_mock_peer_manager_secure("net1".to_string(), "secret".to_string()).await; - - let group_declares = vec![GroupIdentity { - group_name: "platform-admin".to_string(), - group_secret: "platform-admin-secret".to_string(), - }]; - admin_a.get_global_ctx().config.set_acl(Some(Acl { - acl_v1: Some(AclV1 { - group: Some(GroupInfo { - declares: group_declares.clone(), - members: vec!["platform-admin".to_string()], - }), - ..Default::default() - }), - })); - admin_b.get_global_ctx().config.set_acl(Some(Acl { - acl_v1: Some(AclV1 { - group: Some(GroupInfo { - declares: group_declares, - members: vec![], - }), - ..Default::default() - }), - })); - - connect_peer_manager(admin_a.clone(), admin_b.clone()).await; - wait_route_appear(admin_a.clone(), admin_b.clone()) - .await - .unwrap(); - - let (_cred_id, cred_secret) = admin_a - .get_global_ctx() - .get_credential_manager() - .generate_credential(vec![], false, vec![], std::time::Duration::from_secs(3600)); - admin_a - .get_global_ctx() - .issue_event(crate::common::global_ctx::GlobalCtxEvent::CredentialChanged); - - let privkey_bytes: [u8; 32] = base64::engine::general_purpose::STANDARD - .decode(&cred_secret) - .unwrap() - .try_into() - .unwrap(); - let private = x25519_dalek::StaticSecret::from(privkey_bytes); - let credential_pubkey = x25519_dalek::PublicKey::from(&private).as_bytes().to_vec(); - - wait_for_condition( - || { - let admin_b = admin_b.clone(); - let credential_pubkey = credential_pubkey.clone(); - async move { - admin_b.get_global_ctx().is_pubkey_trusted_with_source( - &credential_pubkey, - "net1", - TrustedKeySource::OspfCredential, - ) - } - }, - Duration::from_secs(10), - ) - .await; - - let cred_c = create_mock_peer_manager_credential("net1".to_string(), &private).await; - connect_peer_manager(cred_c.clone(), admin_b.clone()).await; - - let admin_a_id = admin_a.my_peer_id(); - wait_for_condition( - || { - let cred_c = cred_c.clone(); - async move { - cred_c - .list_routes() - .await - .iter() - .any(|r| r.peer_id == admin_a_id) - } - }, - Duration::from_secs(10), - ) - .await; - - wait_for_condition( - || { - let cred_c = cred_c.clone(); - async move { - cred_c - .get_route() - .get_peer_groups(admin_a_id) - .contains(&"platform-admin".to_string()) - } - }, - Duration::from_secs(10), - ) - .await; -} - -/// Minimal test: two secure peers connect and discover each other's route. -#[tokio::test] -async fn two_secure_peers_route_appear() { - let peer_a = create_mock_peer_manager_secure("net1".to_string(), "sec1".to_string()).await; - let peer_b = create_mock_peer_manager_secure("net1".to_string(), "sec1".to_string()).await; - - connect_peer_manager(peer_a.clone(), peer_b.clone()).await; - - wait_route_appear(peer_a.clone(), peer_b.clone()) - .await - .unwrap(); -} - -#[tokio::test] -async fn multi_admin_multi_credential_route_and_revocation_isolation() { - let admin_a = create_mock_peer_manager_secure("net1".to_string(), "secret".to_string()).await; - let admin_b = create_mock_peer_manager_secure("net1".to_string(), "secret".to_string()).await; - let admin_d = create_mock_peer_manager_secure("net1".to_string(), "secret".to_string()).await; - - connect_peer_manager(admin_a.clone(), admin_b.clone()).await; - connect_peer_manager(admin_b.clone(), admin_d.clone()).await; - connect_peer_manager(admin_a.clone(), admin_d.clone()).await; - - wait_route_appear(admin_a.clone(), admin_b.clone()) - .await - .unwrap(); - wait_route_appear(admin_b.clone(), admin_d.clone()) - .await - .unwrap(); - wait_route_appear(admin_a.clone(), admin_d.clone()) - .await - .unwrap(); - - let (cred1_id, cred1_secret) = admin_a - .get_global_ctx() - .get_credential_manager() - .generate_credential( - vec!["guest-a".to_string()], - false, - vec![], - std::time::Duration::from_secs(3600), - ); - let (_cred2_id, cred2_secret) = admin_b - .get_global_ctx() - .get_credential_manager() - .generate_credential( - vec!["guest-b".to_string()], - false, - vec![], - std::time::Duration::from_secs(3600), - ); - - let cred1_private: [u8; 32] = base64::engine::general_purpose::STANDARD - .decode(&cred1_secret) - .unwrap() - .try_into() - .unwrap(); - let cred2_private: [u8; 32] = base64::engine::general_purpose::STANDARD - .decode(&cred2_secret) - .unwrap() - .try_into() - .unwrap(); - let cred_1 = create_mock_peer_manager_credential( - "net1".to_string(), - &x25519_dalek::StaticSecret::from(cred1_private), - ) - .await; - let cred_2 = create_mock_peer_manager_credential( - "net1".to_string(), - &x25519_dalek::StaticSecret::from(cred2_private), - ) - .await; - - connect_peer_manager(cred_1.clone(), admin_a.clone()).await; - connect_peer_manager(cred_2.clone(), admin_b.clone()).await; - - let cred_1_id = cred_1.my_peer_id(); - let cred_2_id = cred_2.my_peer_id(); - - wait_for_condition( - || { - let admin_d = admin_d.clone(); - async move { - let routes = admin_d.list_routes().await; - routes.iter().any(|r| r.peer_id == cred_1_id) - && routes.iter().any(|r| r.peer_id == cred_2_id) - } - }, - Duration::from_secs(15), - ) - .await; - - wait_for_condition( - || { - let admin_d = admin_d.clone(); - async move { - let g1 = admin_d.get_route().get_peer_groups(cred_1_id); - let g2 = admin_d.get_route().get_peer_groups(cred_2_id); - g1.contains(&"guest-a".to_string()) && g2.contains(&"guest-b".to_string()) - } - }, - Duration::from_secs(15), - ) - .await; - - assert!( - admin_a - .get_global_ctx() - .get_credential_manager() - .revoke_credential(&cred1_id) - ); - admin_a - .get_global_ctx() - .issue_event(crate::common::global_ctx::GlobalCtxEvent::CredentialChanged); - - wait_for_condition( - || { - let admin_d = admin_d.clone(); - async move { - let routes = admin_d.list_routes().await; - !routes.iter().any(|r| r.peer_id == cred_1_id) - && routes.iter().any(|r| r.peer_id == cred_2_id) - } - }, - Duration::from_secs(20), - ) - .await; -} - -#[tokio::test] -async fn unknown_credential_rejected_while_valid_credential_survives() { - let admin_a = create_mock_peer_manager_secure("net1".to_string(), "secret".to_string()).await; - let admin_b = create_mock_peer_manager_secure("net1".to_string(), "secret".to_string()).await; - - connect_peer_manager(admin_a.clone(), admin_b.clone()).await; - wait_route_appear(admin_a.clone(), admin_b.clone()) - .await - .unwrap(); - - let (_cred_id, cred_secret) = admin_a - .get_global_ctx() - .get_credential_manager() - .generate_credential( - vec!["stable".to_string()], - false, - vec![], - std::time::Duration::from_secs(3600), - ); - - let valid_private: [u8; 32] = base64::engine::general_purpose::STANDARD - .decode(&cred_secret) - .unwrap() - .try_into() - .unwrap(); - let valid_cred = create_mock_peer_manager_credential( - "net1".to_string(), - &x25519_dalek::StaticSecret::from(valid_private), - ) - .await; - let unknown_private = x25519_dalek::StaticSecret::random_from_rng(rand::rngs::OsRng); - let unknown_cred = - create_mock_peer_manager_credential("net1".to_string(), &unknown_private).await; - - connect_peer_manager(valid_cred.clone(), admin_a.clone()).await; - let (unknown_ring_client, unknown_ring_server) = create_ring_tunnel_pair(); - let unknown_connect_client = tokio::spawn({ - let unknown_cred = unknown_cred.clone(); - async move { - unknown_cred - .add_client_tunnel(unknown_ring_client, false) - .await - } - }); - let unknown_connect_server = tokio::spawn({ - let admin_a = admin_a.clone(); - async move { - admin_a - .add_tunnel_as_server(unknown_ring_server, true) - .await - } - }); - let (unknown_client_ret, unknown_server_ret) = - tokio::join!(unknown_connect_client, unknown_connect_server); - assert!( - unknown_client_ret.unwrap().is_err() || unknown_server_ret.unwrap().is_err(), - "unknown credential connection should fail on at least one side" - ); - - let valid_id = valid_cred.my_peer_id(); - let unknown_id = unknown_cred.my_peer_id(); - - wait_for_condition( - || { - let admin_b = admin_b.clone(); - async move { - admin_b - .list_routes() - .await - .iter() - .any(|r| r.peer_id == valid_id) - } - }, - Duration::from_secs(15), - ) - .await; - - tokio::time::sleep(Duration::from_secs(5)).await; - - let routes = admin_b.list_routes().await; - assert!(routes.iter().any(|r| r.peer_id == valid_id)); - assert!(!routes.iter().any(|r| r.peer_id == unknown_id)); -} diff --git a/easytier/src/proto/mod.rs b/easytier/src/proto/mod.rs index 3315a5da..74444c5b 100644 --- a/easytier/src/proto/mod.rs +++ b/easytier/src/proto/mod.rs @@ -1,21 +1,14 @@ -pub mod rpc_impl; -pub mod rpc_types; +pub use easytier_proto::api; +#[cfg(feature = "management")] +pub use easytier_proto::web; +pub use easytier_proto::{ + ALL_DESCRIPTOR_BYTES, acl, common, core_config, error, peer_rpc, rpc_types, +}; -pub mod acl; -pub mod api; -pub mod common; -pub mod error; #[cfg(feature = "magic-dns")] -pub mod magic_dns; -pub mod peer_rpc; -pub mod web; +pub use easytier_proto::magic_dns; #[cfg(test)] pub mod tests; -pub mod utils; -pub const DESCRIPTOR_POOL_BYTES: &[u8] = - include_bytes!(concat!(env!("OUT_DIR"), "/file_descriptor_set.bin")); - -pub const ALL_DESCRIPTOR_BYTES: &[u8] = - include_bytes!(concat!(env!("OUT_DIR"), "/descriptors.bin")); +pub mod rpc; diff --git a/easytier/src/proto/rpc/mod.rs b/easytier/src/proto/rpc/mod.rs new file mode 100644 index 00000000..02c6940e --- /dev/null +++ b/easytier/src/proto/rpc/mod.rs @@ -0,0 +1,3 @@ +pub mod standalone; + +pub use easytier_core::rpc::{RpcController, bidirect, client, server, service_registry}; diff --git a/easytier/src/proto/rpc/standalone.rs b/easytier/src/proto/rpc/standalone.rs new file mode 100644 index 00000000..372771e6 --- /dev/null +++ b/easytier/src/proto/rpc/standalone.rs @@ -0,0 +1,111 @@ +use easytier_core::{ + connectivity::protocol::raw::{ + TcpTunnelDialer, TcpTunnelListener, TunnelDialer, UdpTunnelDialer, UdpTunnelListener, + }, + rpc::standalone::StandAloneClient, + socket::SocketListener, + socket::udp::{UdpBindOptions, UdpSessionListenRequest}, + tunnel::Tunnel, +}; + +use crate::{ + host_runtime::{NativeHostRuntime, native_host_runtime}, + tunnel::TunnelUrl, +}; + +pub use easytier_core::rpc::standalone::{RpcServerHook, StandAloneServer}; + +pub type RuntimeRpcDialer = TcpTunnelDialer; +pub type RuntimeRpcListener = TcpTunnelListener; +pub type RuntimeRpcClient = StandAloneClient; + +pub fn runtime_rpc_dialer(remote_url: url::Url) -> RuntimeRpcDialer { + TcpTunnelDialer::new(remote_url, native_host_runtime(), native_host_runtime()) +} + +pub fn runtime_rpc_client(remote_url: url::Url) -> RuntimeRpcClient { + StandAloneClient::new(runtime_rpc_dialer(remote_url)) +} + +pub fn runtime_rpc_listener(local_addr: std::net::SocketAddr) -> RuntimeRpcListener { + TcpTunnelListener::new(local_addr, native_host_runtime()) +} + +pub fn runtime_udp_tunnel_dialer(remote_url: url::Url) -> impl TunnelDialer { + UdpTunnelDialer::new(remote_url, native_host_runtime(), native_host_runtime()) +} + +pub fn runtime_udp_tunnel_listener( + local_url: url::Url, + local_addr: std::net::SocketAddr, +) -> impl SocketListener> { + let bind = UdpBindOptions::port_bound_listener(local_addr) + .with_bind_device(TunnelUrl::from(local_url.clone()).bind_dev()) + .with_only_v6(true); + UdpTunnelListener::new_with_request( + local_url, + UdpSessionListenRequest::new(bind), + native_host_runtime(), + ) +} + +#[cfg(test)] +mod tests { + use easytier_core::{ + connectivity::protocol::raw::TunnelDialer as _, socket::SocketListener as _, + }; + + use crate::proto::rpc::standalone::{ + StandAloneServer, runtime_rpc_dialer, runtime_rpc_listener, runtime_udp_tunnel_dialer, + runtime_udp_tunnel_listener, + }; + + #[tokio::test] + async fn standalone_exit_on_drop() { + let addr = "0.0.0.0:53884".parse().unwrap(); + let tunnel = runtime_rpc_listener(addr); + let mut server = StandAloneServer::new(tunnel); + server.serve().await.unwrap(); + drop(server); + + // tcp should closed + let connector = runtime_rpc_dialer("tcp://0.0.0.0:53884".parse().unwrap()); + connector.connect().await.unwrap_err(); + } + + #[tokio::test] + async fn standalone_ipv4_and_ipv6_listeners_share_port() { + let mut ipv6 = runtime_rpc_listener("[::]:0".parse().unwrap()); + ipv6.listen().await.unwrap(); + let port = ipv6.local_url().port().unwrap(); + + let mut ipv4 = runtime_rpc_listener(format!("0.0.0.0:{port}").parse().unwrap()); + ipv4.listen().await.unwrap(); + } + + #[tokio::test] + async fn runtime_udp_tunnel_endpoints_connect() { + let local_url = "udp://127.0.0.1:0".parse().unwrap(); + let mut listener = runtime_udp_tunnel_listener(local_url, "127.0.0.1:0".parse().unwrap()); + listener.listen().await.unwrap(); + let listener_url = listener.local_url(); + let dialer = runtime_udp_tunnel_dialer(listener_url.clone()); + + let (accepted, connected) = + tokio::time::timeout(std::time::Duration::from_secs(5), async { + tokio::try_join!(listener.accept(), dialer.connect()) + }) + .await + .unwrap() + .unwrap(); + + assert_eq!( + accepted.info().unwrap().local_addr.unwrap().url, + listener_url.as_str() + ); + assert_eq!( + connected.info().unwrap().remote_addr.unwrap().url, + listener_url.as_str() + ); + } +} diff --git a/easytier/src/proto/rpc_impl/standalone.rs b/easytier/src/proto/rpc_impl/standalone.rs deleted file mode 100644 index 0988ad04..00000000 --- a/easytier/src/proto/rpc_impl/standalone.rs +++ /dev/null @@ -1,224 +0,0 @@ -use std::{ - sync::{Arc, Mutex, atomic::AtomicU32}, - time::Duration, -}; - -use anyhow::Context as _; -use tokio::task::JoinSet; - -use crate::{ - common::join_joinset_background, - proto::{ - common::TunnelInfo, - rpc_impl::bidirect::BidirectRpcManager, - rpc_types::{__rt::RpcClientFactory, error::Error}, - }, - tunnel::{Tunnel, TunnelConnector, TunnelListener}, -}; - -use super::service_registry::ServiceRegistry; - -#[async_trait::async_trait] -#[auto_impl::auto_impl(Arc, Box)] -pub trait RpcServerHook: Send + Sync { - async fn on_new_client( - &self, - tunnel_info: Option, - ) -> Result, anyhow::Error> { - Ok(tunnel_info) - } - async fn on_client_disconnected(&self, _tunnel_info: Option) {} -} - -struct DefaultHook; -impl RpcServerHook for DefaultHook {} - -pub struct StandAloneServer { - registry: Arc, - listener: Option, - inflight_server: Arc, - tasks: JoinSet<()>, - hook: Option>, - rx_timeout: Option, -} - -impl StandAloneServer { - pub fn new(listener: L) -> Self { - StandAloneServer { - registry: Arc::new(ServiceRegistry::new()), - listener: Some(listener), - inflight_server: Arc::new(AtomicU32::new(0)), - tasks: JoinSet::new(), - - hook: None, - rx_timeout: Some(Duration::from_secs(60)), - } - } - - pub fn set_rx_timeout(&mut self, timeout: Option) { - self.rx_timeout = timeout; - } - - pub fn set_hook(&mut self, hook: Arc) { - self.hook = Some(hook); - } - - pub fn registry(&self) -> &ServiceRegistry { - &self.registry - } - - async fn serve_loop( - listener: &mut L, - inflight: Arc, - registry: Arc, - hook: Arc, - rx_timeout: Option, - ) -> Result<(), Error> { - let tasks = Arc::new(Mutex::new(JoinSet::new())); - join_joinset_background(tasks.clone(), "standalone serve_loop".to_string()); - - loop { - let tunnel = listener.accept().await?; - let tunnel_info = tunnel.info(); - let registry = registry.clone(); - let inflight_server = inflight.clone(); - let hook = hook.clone(); - - let tunnel_info = match hook.on_new_client(tunnel_info).await { - Ok(info) => info, - Err(e) => { - tracing::warn!(?e, "standalone hook.on_new_client failed"); - continue; - } - }; - - inflight_server.fetch_add(1, std::sync::atomic::Ordering::Relaxed); - tasks.lock().unwrap().spawn(async move { - let server = BidirectRpcManager::new().set_rx_timeout(rx_timeout); - server.rpc_server().registry().replace_registry(®istry); - server.run_with_tunnel(tunnel); - server.wait().await; - hook.on_client_disconnected(tunnel_info.clone()).await; - inflight_server.fetch_sub(1, std::sync::atomic::Ordering::Relaxed); - }); - } - } - - pub async fn serve(&mut self) -> Result<(), Error> { - let mut listener = self.listener.take().unwrap(); - let hook = self.hook.take().unwrap_or_else(|| Arc::new(DefaultHook)); - let rx_timeout = self.rx_timeout; - - listener - .listen() - .await - .with_context(|| "failed to listen")?; - - let registry = self.registry.clone(); - - let inflight_server = self.inflight_server.clone(); - - self.tasks.spawn(async move { - loop { - let ret = Self::serve_loop( - &mut listener, - inflight_server.clone(), - registry.clone(), - hook.clone(), - rx_timeout, - ) - .await; - if let Err(e) = ret { - tracing::error!(?e, url = ?listener.local_url(), "serve_loop exit unexpectedly"); - println!("standalone serve_loop exit unexpectedly: {:?}", e); - } - - tokio::time::sleep(Duration::from_secs(1)).await; - } - }); - - Ok(()) - } - - pub fn inflight_server(&self) -> u32 { - self.inflight_server - .load(std::sync::atomic::Ordering::Relaxed) - } -} - -pub struct StandAloneClient { - connector: C, - client: Option, -} - -impl StandAloneClient { - pub fn new(connector: C) -> Self { - StandAloneClient { - connector, - client: None, - } - } - - async fn connect(&mut self) -> Result, Error> { - Ok(self.connector.connect().await.with_context(|| { - format!( - "failed to connect to server: {:?}", - self.connector.remote_url() - ) - })?) - } - - pub async fn scoped_client( - &mut self, - domain_name: String, - ) -> Result { - let mut c = self.client.take(); - let error = c.as_ref().and_then(|c| c.take_error()); - if c.is_none() || error.is_some() { - tracing::info!("reconnect due to error: {:?}", error); - let tunnel = self.connect().await?; - let mgr = BidirectRpcManager::new().set_rx_timeout(Some(Duration::from_secs(60))); - mgr.run_with_tunnel(tunnel); - c = Some(mgr); - } - - self.client = c; - - Ok(self - .client - .as_ref() - .unwrap() - .rpc_client() - .scoped_client::(1, 1, domain_name)) - } - - pub async fn wait(&mut self) { - if let Some(client) = self.client.take() { - client.wait().await; - } - } -} - -#[cfg(test)] -mod tests { - use crate::{ - proto::rpc_impl::standalone::StandAloneServer, - tunnel::{ - TunnelConnector as _, - tcp::{TcpTunnelConnector, TcpTunnelListener}, - }, - }; - - #[tokio::test] - async fn standalone_exit_on_drop() { - let addr: url::Url = "tcp://0.0.0.0:53884".parse().unwrap(); - let tunnel = TcpTunnelListener::new(addr.clone()); - let mut server = StandAloneServer::new(tunnel); - server.serve().await.unwrap(); - drop(server); - - // tcp should closed - let mut connector = TcpTunnelConnector::new(addr); - connector.connect().await.unwrap_err(); - } -} diff --git a/easytier/src/proto/tests.rs b/easytier/src/proto/tests.rs index 4bc4f6c0..2c97b1c6 100644 --- a/easytier/src/proto/tests.rs +++ b/easytier/src/proto/tests.rs @@ -1,12 +1,11 @@ -include!(concat!(env!("OUT_DIR"), "/tests.rs")); -include!(concat!(env!("OUT_DIR"), "/tests.serde.rs")); +pub use easytier_proto::tests::*; use std::sync::{Arc, Mutex}; use futures::StreamExt as _; use tokio::task::JoinSet; -use super::rpc_impl::RpcController; +use super::rpc::RpcController; #[derive(Clone, Default)] struct GreetingJsonCallHandler; @@ -138,13 +137,13 @@ impl Greeting for GreetingService { } use crate::proto::common::{CompressionAlgoPb, RpcCompressionInfo}; -use crate::proto::rpc_impl::client::Client; -use crate::proto::rpc_impl::server::Server; +use crate::proto::rpc::client::Client; +use crate::proto::rpc::server::Server; struct TestContext { client: Client, server: Server, - tasks: Arc>>, + _tasks: Arc>>, } impl TestContext { @@ -186,7 +185,7 @@ impl TestContext { Self { client, server: rpc_server, - tasks, + _tasks: tasks, } } } @@ -294,7 +293,7 @@ async fn rpc_timeout_test() { #[tokio::test] async fn rpc_tunnel_stuck_test() { use crate::proto::rpc_types; - use crate::tunnel::ring::RING_TUNNEL_CAP; + use easytier_core::tunnel::ring::RING_TUNNEL_CAP; let rpc_server = Server::new(); rpc_server.run(); @@ -372,62 +371,11 @@ async fn rpc_tunnel_stuck_test() { assert_eq!(ret.greeting, "Hello fuck world!"); } -#[tokio::test] -async fn standalone_rpc_test() { - use crate::proto::rpc_impl::standalone::{StandAloneClient, StandAloneServer}; - use crate::tunnel::tcp::{TcpTunnelConnector, TcpTunnelListener}; - - let mut server = StandAloneServer::new(TcpTunnelListener::new( - "tcp://0.0.0.0:33455".parse().unwrap(), - )); - let service = GreetingServer::new(GreetingService { - delay_ms: 0, - prefix: "Hello".to_string(), - }); - server.registry().register(service, "test"); - server.serve().await.unwrap(); - - tokio::time::sleep(std::time::Duration::from_millis(100)).await; - - let mut client = StandAloneClient::new(TcpTunnelConnector::new( - "tcp://127.0.0.1:33455".parse().unwrap(), - )); - - let out = client - .scoped_client::>("test".to_string()) - .await - .unwrap(); - - let ctrl = RpcController::default(); - let input = SayHelloRequest { - name: "world".to_string(), - }; - let ret = out.say_hello(ctrl, input).await; - assert_eq!(ret.unwrap().greeting, "Hello world!"); - - let out = client - .scoped_client::>("test".to_string()) - .await - .unwrap(); - - let ctrl = RpcController::default(); - let input = SayGoodbyeRequest { - name: "world".to_string(), - }; - let ret = out.say_goodbye(ctrl, input).await; - assert_eq!(ret.unwrap().greeting, "Goodbye, world!"); - - drop(client); - - tokio::time::sleep(std::time::Duration::from_secs(1)).await; - assert_eq!(0, server.inflight_server()); -} - #[tokio::test] async fn test_bidirect_rpc_manager() { - use crate::proto::rpc_impl::bidirect::BidirectRpcManager; - use crate::tunnel::tcp::{TcpTunnelConnector, TcpTunnelListener}; - use crate::tunnel::{TunnelConnector, TunnelListener}; + use crate::proto::rpc::bidirect::BidirectRpcManager; + use crate::proto::rpc::standalone::{runtime_rpc_dialer, runtime_rpc_listener}; + use easytier_core::{connectivity::protocol::raw::TunnelDialer, socket::SocketListener}; use tokio::sync::Notify; use tokio_util::task::AbortOnDropHandle; @@ -448,7 +396,7 @@ async fn test_bidirect_rpc_manager() { let server_test_done = Arc::new(Notify::new()); let server_test_done_clone = server_test_done.clone(); - let mut tcp_listener = TcpTunnelListener::new("tcp://0.0.0.0:55443".parse().unwrap()); + let mut tcp_listener = runtime_rpc_listener("0.0.0.0:55443".parse().unwrap()); let s_task = AbortOnDropHandle::new(tokio::spawn(async move { tcp_listener.listen().await.unwrap(); let tunnel = tcp_listener.accept().await.unwrap(); @@ -476,7 +424,7 @@ async fn test_bidirect_rpc_manager() { tokio::time::sleep(std::time::Duration::from_secs(1)).await; - let mut tcp_connector = TcpTunnelConnector::new("tcp://0.0.0.0:55443".parse().unwrap()); + let tcp_connector = runtime_rpc_dialer("tcp://0.0.0.0:55443".parse().unwrap()); let c_tunnel = tcp_connector.connect().await.unwrap(); c.run_with_tunnel(c_tunnel); diff --git a/easytier/src/proto/utils.rs b/easytier/src/proto/utils.rs deleted file mode 100644 index 951a9b2c..00000000 --- a/easytier/src/proto/utils.rs +++ /dev/null @@ -1,101 +0,0 @@ -use delegate::delegate; -use derivative::Derivative; -use derive_more::{AsMut, AsRef, Deref, DerefMut, From, IntoIterator}; -use prost::Message; -use serde::{Deserialize, Serialize}; -use sha2::{Digest, Sha256}; - -/// Generates a stable digest strictly within the lifecycle of the current process. -/// -/// ⚠️ WARNING: -/// - This digest is ONLY guaranteed to be deterministic within a **single process and the exact same binary build**. -pub trait TransientDigest: Message { - fn digest(&self) -> [u8; 32] - where - Self: Sized, - { - let buf = self.encode_to_vec(); - let mut hasher = Sha256::new(); - hasher.update(buf); - hasher.finalize().into() - } -} - -impl TransientDigest for S {} - -pub trait MessageModel: - Into + for<'m> TryFrom<&'m Message> -{ -} - -impl MessageModel for Model -where - Message: prost::Message, - Model: Into + for<'m> TryFrom<&'m Message>, -{ -} - -#[derive( - Derivative, - Debug, - Clone, - PartialEq, - Eq, - Hash, - From, - Deref, - DerefMut, - AsRef, - AsMut, - Serialize, - Deserialize, - IntoIterator, -)] -#[derivative(Default(bound = ""))] -#[as_ref(forward)] -#[as_mut(forward)] -#[serde(transparent)] -#[into_iterator(owned, ref, ref_mut)] -pub struct RepeatedMessageModel(Vec); - -impl RepeatedMessageModel { - pub fn into_inner(self) -> Vec { - self.0 - } -} - -impl FromIterator for RepeatedMessageModel { - fn from_iter>(iter: I) -> Self { - Self(iter.into_iter().collect()) - } -} - -impl Extend for RepeatedMessageModel { - delegate! { - to self.0 { - fn extend>(&mut self, iter: T); - } - } -} - -impl<'m, Message, Model> TryFrom<&'m [Message]> for RepeatedMessageModel -where - Message: prost::Message, - Model: MessageModel, -{ - type Error = >::Error; - - fn try_from(value: &'m [Message]) -> Result { - value.iter().map(TryInto::try_into).collect() - } -} - -impl From> for Vec -where - Message: prost::Message, - Model: MessageModel, -{ - fn from(value: RepeatedMessageModel) -> Self { - value.into_iter().map(Into::into).collect() - } -} diff --git a/easytier/src/rpc_service/acl_manage.rs b/easytier/src/rpc_service/acl_manage.rs deleted file mode 100644 index 452375fb..00000000 --- a/easytier/src/rpc_service/acl_manage.rs +++ /dev/null @@ -1,50 +0,0 @@ -use std::sync::Arc; - -use crate::{ - instance_manager::NetworkInstanceManager, - proto::{ - api::instance::{ - AclManageRpc, GetAclStatsRequest, GetAclStatsResponse, GetWhitelistRequest, - GetWhitelistResponse, - }, - rpc_types::controller::BaseController, - }, -}; - -#[derive(Clone)] -pub struct AclManageRpcService { - instance_manager: Arc, -} - -impl AclManageRpcService { - pub fn new(instance_manager: Arc) -> Self { - Self { instance_manager } - } -} - -#[async_trait::async_trait] -impl AclManageRpc for AclManageRpcService { - type Controller = BaseController; - - async fn get_acl_stats( - &self, - ctrl: Self::Controller, - req: GetAclStatsRequest, - ) -> crate::proto::rpc_types::error::Result { - super::get_instance_service(&self.instance_manager, &req.instance)? - .get_acl_manage_service() - .get_acl_stats(ctrl, req) - .await - } - - async fn get_whitelist( - &self, - ctrl: Self::Controller, - req: GetWhitelistRequest, - ) -> crate::proto::rpc_types::error::Result { - super::get_instance_service(&self.instance_manager, &req.instance)? - .get_acl_manage_service() - .get_whitelist(ctrl, req) - .await - } -} diff --git a/easytier/src/rpc_service/api.rs b/easytier/src/rpc_service/api.rs index 4aedca31..cb732e84 100644 --- a/easytier/src/rpc_service/api.rs +++ b/easytier/src/rpc_service/api.rs @@ -2,81 +2,68 @@ use std::{net::SocketAddr, sync::Arc}; use anyhow::Context; use cidr::IpCidr; +#[cfg(feature = "management")] +use easytier_core::management::ManagementServer; +use easytier_core::{management::ReadOnlyManagementServer, socket::SocketListener, tunnel::Tunnel}; +#[cfg(feature = "management")] use crate::{ - instance::instance::InstanceRpcServerHook, - instance_manager::NetworkInstanceManager, + instance::config_storage::NativeConfigFileStorage, rpc_service::logger::NativeLoggerControl, + web_client::DefaultHooks, +}; +use crate::{ + instance::factory::NativeInstanceManager, proto::{ - api::{ - config::ConfigRpcServer, - instance::{ - AclManageRpcServer, ConnectorManageRpcServer, CredentialManageRpcServer, - MappedListenerManageRpcServer, PeerManageRpcServer, PortForwardManageRpcServer, - StatsRpcServer, TcpProxyRpcServer, VpnPortalRpcServer, - }, - logger::LoggerRpcServer, - manage::WebClientServiceServer, - }, - peer_rpc::PeerCenterRpcServer, - rpc_impl::{service_registry::ServiceRegistry, standalone::StandAloneServer}, + rpc::standalone::{RuntimeRpcListener, runtime_rpc_listener}, rpc_types::error::Error, }, - rpc_service::{ - acl_manage::AclManageRpcService, config::ConfigRpcService, - connector_manage::ConnectorManageRpcService, credential_manage::CredentialManageRpcService, - instance_manage::InstanceManageRpcService, logger::LoggerRpcService, - mapped_listener_manage::MappedListenerManageRpcService, - peer_center::PeerCenterManageRpcService, peer_manage::PeerManageRpcService, - port_forward_manage::PortForwardManageRpcService, protected_port, - proxy::TcpProxyRpcService, stats::StatsRpcService, vpn_portal::VpnPortalRpcService, - }, - tunnel::{TunnelListener, tcp::TcpTunnelListener}, - web_client::{DefaultHooks, WebClientHooks}, }; -pub struct ApiRpcServer { - rpc_server: StandAloneServer, - protected_tcp_port: Option, +#[cfg(feature = "management")] +pub struct ApiRpcServer +where + T: SocketListener> + 'static, +{ + rpc_server: ManagementServer, } -impl ApiRpcServer { +#[cfg(feature = "management")] +impl ApiRpcServer { pub fn new( rpc_portal: Option, rpc_portal_whitelist: Option>, - instance_manager: Arc, + instance_manager: Arc, ) -> anyhow::Result { let rpc_addr = parse_rpc_portal(rpc_portal)?; - let mut server = Self::from_tunnel( - TcpTunnelListener::new( - format!("tcp://{}", rpc_addr) - .parse() - .context("failed to parse rpc portal address")?, - ), - instance_manager, - ); - protected_port::register_protected_tcp_port(rpc_addr.port()); - server.protected_tcp_port = Some(rpc_addr.port()); - - server - .rpc_server - .set_hook(Arc::new(InstanceRpcServerHook::new(rpc_portal_whitelist))); + let mut server = Self::from_tunnel(runtime_rpc_listener(rpc_addr), instance_manager); + server.rpc_server.set_whitelist(rpc_portal_whitelist); Ok(server) } } -impl ApiRpcServer { - pub fn from_tunnel(tunnel: T, instance_manager: Arc) -> Self { - let rpc_server = StandAloneServer::new(tunnel); - register_api_rpc_service(&instance_manager, rpc_server.registry(), None); - Self { - rpc_server, - protected_tcp_port: None, - } +#[cfg(feature = "management")] +impl ApiRpcServer +where + T: SocketListener> + 'static, +{ + pub fn from_tunnel(tunnel: T, instance_manager: Arc) -> Self { + let rpc_server = ManagementServer::new( + tunnel, + instance_manager, + Arc::new(DefaultHooks), + Arc::new(NativeConfigFileStorage), + Arc::new(NativeLoggerControl), + ); + Self { rpc_server } } } -impl ApiRpcServer { +#[cfg(feature = "management")] +impl ApiRpcServer +where + T: SocketListener> + 'static, +{ pub async fn serve(mut self) -> Result { self.rpc_server.serve().await?; Ok(self) @@ -88,106 +75,60 @@ impl ApiRpcServer { } } -impl Drop for ApiRpcServer { - fn drop(&mut self) { - if let Some(port) = self.protected_tcp_port.take() { - protected_port::unregister_protected_tcp_port(port); - } - self.rpc_server.registry().unregister_all(); +pub struct ReadOnlyApiRpcServer +where + T: SocketListener> + 'static, +{ + rpc_server: ReadOnlyManagementServer, +} + +impl ReadOnlyApiRpcServer { + pub fn new( + rpc_portal: Option, + rpc_portal_whitelist: Option>, + instance_manager: Arc, + ) -> anyhow::Result { + let rpc_addr = parse_rpc_portal(rpc_portal)?; + let mut server = Self::from_tunnel(runtime_rpc_listener(rpc_addr), instance_manager); + server.rpc_server.set_whitelist(rpc_portal_whitelist); + Ok(server) } } -pub fn register_api_rpc_service( - instance_manager: &Arc, - registry: &ServiceRegistry, - hooks: Option>, -) { - registry.register( - PeerManageRpcServer::new(PeerManageRpcService::new(instance_manager.clone())), - "", - ); - - registry.register( - ConnectorManageRpcServer::new(ConnectorManageRpcService::new(instance_manager.clone())), - "", - ); - - registry.register( - MappedListenerManageRpcServer::new(MappedListenerManageRpcService::new( - instance_manager.clone(), - )), - "", - ); - - registry.register( - VpnPortalRpcServer::new(VpnPortalRpcService::new(instance_manager.clone())), - "", - ); - - for client_type in ["tcp", "kcp_src", "kcp_dst", "quic_src", "quic_dst"] { - registry.register( - TcpProxyRpcServer::new(TcpProxyRpcService::new( - instance_manager.clone(), - client_type, - )), - client_type, - ); +impl ReadOnlyApiRpcServer +where + T: SocketListener> + 'static, +{ + pub fn from_tunnel(tunnel: T, instance_manager: Arc) -> Self { + Self { + rpc_server: ReadOnlyManagementServer::new(tunnel, instance_manager), + } } - registry.register( - AclManageRpcServer::new(AclManageRpcService::new(instance_manager.clone())), - "", - ); + pub async fn serve(mut self) -> Result { + self.rpc_server.serve().await?; + Ok(self) + } - registry.register( - PortForwardManageRpcServer::new(PortForwardManageRpcService::new(instance_manager.clone())), - "", - ); - - registry.register( - StatsRpcServer::new(StatsRpcService::new(instance_manager.clone())), - "", - ); - - registry.register(LoggerRpcServer::new(LoggerRpcService), ""); - - registry.register( - ConfigRpcServer::new(ConfigRpcService::new(instance_manager.clone())), - "", - ); - - registry.register( - WebClientServiceServer::new(InstanceManageRpcService::new( - instance_manager.clone(), - hooks.unwrap_or(Arc::new(DefaultHooks)), - )), - "", - ); - - registry.register( - PeerCenterRpcServer::new(PeerCenterManageRpcService::new(instance_manager.clone())), - "", - ); - - registry.register( - CredentialManageRpcServer::new(CredentialManageRpcService::new(instance_manager.clone())), - "", - ); + pub fn with_rx_timeout(mut self, timeout: Option) -> Self { + self.rpc_server.set_rx_timeout(timeout); + self + } } fn parse_rpc_portal(rpc_portal: Option) -> anyhow::Result { - if let Some(Ok(port)) = rpc_portal.as_ref().map(|s| s.parse::()) { - Ok(SocketAddr::from(([0, 0, 0, 0], port))) + let mut rpc_addr = if let Some(Ok(port)) = rpc_portal.as_ref().map(|s| s.parse::()) { + Some(SocketAddr::from(([0, 0, 0, 0], port))) } else { - let mut rpc_addr = rpc_portal + rpc_portal .map(|addr| { addr.parse::() .context("failed to parse rpc portal address") }) - .transpose()?; - select_proper_rpc_port(&mut rpc_addr)?; - rpc_addr.ok_or_else(|| anyhow::anyhow!("failed to parse rpc portal address")) - } + .transpose()? + }; + select_proper_rpc_port(&mut rpc_addr)?; + rpc_addr.ok_or_else(|| anyhow::anyhow!("failed to parse rpc portal address")) } fn select_proper_rpc_port(addr: &mut Option) -> anyhow::Result<()> { @@ -211,3 +152,90 @@ fn select_proper_rpc_port(addr: &mut Option) -> anyhow::Result<()> { } } } + +#[cfg(all(test, feature = "management"))] +mod tests { + use std::{fmt, sync::Arc, time::Duration}; + + use easytier_core::{ + rpc::bidirect::BidirectRpcManager, + socket::SocketListener, + tunnel::{Tunnel, ring::create_ring_tunnel_pair}, + }; + use tokio::sync::mpsc; + + use crate::{ + instance::factory::native_instance_manager, + proto::{ + api::logger::{GetLoggerConfigRequest, LoggerRpc, LoggerRpcClientFactory}, + rpc_types::controller::BaseController, + }, + }; + + use super::{ApiRpcServer, parse_rpc_portal}; + + #[test] + fn zero_rpc_portal_is_resolved_before_listener_binding() { + assert_ne!(parse_rpc_portal(Some("0".to_owned())).unwrap().port(), 0); + } + + struct RingListener { + accepted: mpsc::Receiver>, + } + + impl fmt::Debug for RingListener { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter.debug_struct("RingListener").finish() + } + } + + #[async_trait::async_trait] + impl SocketListener for RingListener { + type Accepted = Box; + + async fn listen(&mut self) -> anyhow::Result<()> { + Ok(()) + } + + async fn accept(&mut self) -> anyhow::Result { + self.accepted + .recv() + .await + .ok_or_else(|| anyhow::anyhow!("ring test listener closed")) + } + + fn local_url(&self) -> url::Url { + "ring://management-test".parse().unwrap() + } + } + + #[tokio::test] + async fn trusted_ring_management_transport_does_not_require_an_ip_host() { + let (client_tunnel, server_tunnel) = create_ring_tunnel_pair(); + let (accepted, receiver) = mpsc::channel(1); + accepted.send(server_tunnel).await.unwrap(); + let server = ApiRpcServer::from_tunnel( + RingListener { accepted: receiver }, + Arc::new(native_instance_manager()), + ) + .with_rx_timeout(Some(Duration::from_secs(1))) + .serve() + .await + .unwrap(); + let client = BidirectRpcManager::new().set_rx_timeout(Some(Duration::from_secs(1))); + client.run_with_tunnel(client_tunnel); + let logger = client + .rpc_client() + .scoped_client::>(1, 1, String::new()); + + tokio::time::timeout( + Duration::from_secs(1), + logger.get_logger_config(BaseController::default(), GetLoggerConfigRequest::default()), + ) + .await + .unwrap() + .unwrap(); + + drop(server); + } +} diff --git a/easytier/src/rpc_service/config.rs b/easytier/src/rpc_service/config.rs deleted file mode 100644 index 1f1b9214..00000000 --- a/easytier/src/rpc_service/config.rs +++ /dev/null @@ -1,49 +0,0 @@ -use std::sync::Arc; - -use crate::{ - instance_manager::NetworkInstanceManager, - proto::{ - api::config::{ - ConfigRpc, GetConfigRequest, GetConfigResponse, PatchConfigRequest, PatchConfigResponse, - }, - rpc_types::{self, controller::BaseController}, - }, -}; - -#[derive(Clone)] -pub struct ConfigRpcService { - instance_manager: Arc, -} - -impl ConfigRpcService { - pub fn new(instance_manager: Arc) -> Self { - Self { instance_manager } - } -} - -#[async_trait::async_trait] -impl ConfigRpc for ConfigRpcService { - type Controller = BaseController; - - async fn patch_config( - &self, - ctrl: Self::Controller, - input: PatchConfigRequest, - ) -> Result { - super::get_instance_service(&self.instance_manager, &input.instance)? - .get_config_service() - .patch_config(ctrl, input) - .await - } - - async fn get_config( - &self, - ctrl: Self::Controller, - input: GetConfigRequest, - ) -> Result { - super::get_instance_service(&self.instance_manager, &input.instance)? - .get_config_service() - .get_config(ctrl, input) - .await - } -} diff --git a/easytier/src/rpc_service/connector_manage.rs b/easytier/src/rpc_service/connector_manage.rs deleted file mode 100644 index 6030159d..00000000 --- a/easytier/src/rpc_service/connector_manage.rs +++ /dev/null @@ -1,36 +0,0 @@ -use std::sync::Arc; - -use crate::{ - instance_manager::NetworkInstanceManager, - proto::{ - api::instance::{ConnectorManageRpc, ListConnectorRequest, ListConnectorResponse}, - rpc_types::controller::BaseController, - }, -}; - -#[derive(Clone)] -pub struct ConnectorManageRpcService { - instance_manager: Arc, -} - -impl ConnectorManageRpcService { - pub fn new(instance_manager: Arc) -> Self { - Self { instance_manager } - } -} - -#[async_trait::async_trait] -impl ConnectorManageRpc for ConnectorManageRpcService { - type Controller = BaseController; - - async fn list_connector( - &self, - ctrl: Self::Controller, - req: ListConnectorRequest, - ) -> crate::proto::rpc_types::error::Result { - super::get_instance_service(&self.instance_manager, &req.instance)? - .get_connector_manage_service() - .list_connector(ctrl, req) - .await - } -} diff --git a/easytier/src/rpc_service/credential_manage.rs b/easytier/src/rpc_service/credential_manage.rs deleted file mode 100644 index 716d979a..00000000 --- a/easytier/src/rpc_service/credential_manage.rs +++ /dev/null @@ -1,62 +0,0 @@ -use std::sync::Arc; - -use crate::{ - instance_manager::NetworkInstanceManager, - proto::{ - api::instance::{ - CredentialManageRpc, GenerateCredentialRequest, GenerateCredentialResponse, - ListCredentialsRequest, ListCredentialsResponse, RevokeCredentialRequest, - RevokeCredentialResponse, - }, - rpc_types::controller::BaseController, - }, -}; - -#[derive(Clone)] -pub struct CredentialManageRpcService { - instance_manager: Arc, -} - -impl CredentialManageRpcService { - pub fn new(instance_manager: Arc) -> Self { - Self { instance_manager } - } -} - -#[async_trait::async_trait] -impl CredentialManageRpc for CredentialManageRpcService { - type Controller = BaseController; - - async fn generate_credential( - &self, - ctrl: Self::Controller, - req: GenerateCredentialRequest, - ) -> crate::proto::rpc_types::error::Result { - super::get_instance_service(&self.instance_manager, &req.instance)? - .get_credential_manage_service() - .generate_credential(ctrl, req) - .await - } - - async fn revoke_credential( - &self, - ctrl: Self::Controller, - req: RevokeCredentialRequest, - ) -> crate::proto::rpc_types::error::Result { - super::get_instance_service(&self.instance_manager, &req.instance)? - .get_credential_manage_service() - .revoke_credential(ctrl, req) - .await - } - - async fn list_credentials( - &self, - ctrl: Self::Controller, - req: ListCredentialsRequest, - ) -> crate::proto::rpc_types::error::Result { - super::get_instance_service(&self.instance_manager, &req.instance)? - .get_credential_manage_service() - .list_credentials(ctrl, req) - .await - } -} diff --git a/easytier/src/rpc_service/instance_manage.rs b/easytier/src/rpc_service/instance_manage.rs deleted file mode 100644 index 3dd038d4..00000000 --- a/easytier/src/rpc_service/instance_manage.rs +++ /dev/null @@ -1,1301 +0,0 @@ -use std::{collections::HashSet, sync::Arc}; - -use crate::{ - common::config::{ - ConfigFileControl, ConfigFilePermission, ConfigLoader, ConfigSource, TomlConfigLoader, - }, - instance_manager::NetworkInstanceManager, - proto::{ - api::{ - config::GetConfigRequest, - manage::{ - CollectNetworkInfoRequest, CollectNetworkInfoResponse, - DeleteNetworkInstanceRequest, DeleteNetworkInstanceResponse, - GetNetworkInstanceConfigRequest, GetNetworkInstanceConfigResponse, - ListNetworkInstanceMetaRequest, ListNetworkInstanceMetaResponse, - ListNetworkInstanceRequest, ListNetworkInstanceResponse, - NetworkInstanceRunningInfoMap, NetworkMeta, RetainNetworkInstanceRequest, - RetainNetworkInstanceResponse, RunNetworkInstanceRequest, - RunNetworkInstanceResponse, ValidateConfigRequest, ValidateConfigResponse, - WebClientService, - }, - }, - rpc_types::{self, controller::BaseController}, - }, - web_client::WebClientHooks, -}; - -#[derive(Clone)] -pub struct InstanceManageRpcService { - manager: Arc, - hooks: Arc, - remote_mutation_lock: Arc>, -} - -impl InstanceManageRpcService { - pub fn new(manager: Arc, hooks: Arc) -> Self { - let remote_mutation_lock = manager.remote_mutation_lock(); - Self { - manager, - hooks, - remote_mutation_lock, - } - } -} - -async fn is_remote_removable(control: &ConfigFileControl) -> bool { - if control.is_read_only() || !control.is_deletable() { - return false; - } - let Some(path) = control.path.as_ref() else { - return true; - }; - - !ConfigFileControl::from_path(path.clone()) - .await - .is_read_only() -} - -async fn ensure_remote_overwritable( - inst_id: uuid::Uuid, - control: &ConfigFileControl, -) -> anyhow::Result<()> { - if control.is_read_only() { - return Err(anyhow::anyhow!( - "instance {} is read-only, cannot be overwritten", - inst_id - )); - } - if !control.is_deletable() { - return Err(anyhow::anyhow!( - "instance {} is no-delete, cannot be overwritten", - inst_id - )); - } - - if let Some(path) = control.path.as_ref() { - let real_control = ConfigFileControl::from_path(path.clone()).await; - if real_control.is_read_only() { - return Err(anyhow::anyhow!( - "config file {} is read-only, cannot be overwritten", - path.display() - )); - } - } - - Ok(()) -} - -async fn ensure_overwritable( - inst_id: uuid::Uuid, - control: &ConfigFileControl, -) -> anyhow::Result<()> { - if control.is_read_only() { - return Err(anyhow::anyhow!( - "instance {} is read-only, cannot be overwritten", - inst_id - )); - } - - if let Some(path) = control.path.as_ref() { - let real_control = ConfigFileControl::from_path(path.clone()).await; - if real_control.is_read_only() { - return Err(anyhow::anyhow!( - "config file {} is read-only, cannot be overwritten", - path.display() - )); - } - } - - Ok(()) -} - -struct StartedInstanceCleanup { - manager: Arc, - started_inst_id: Option, - config_file_cleanup: Option, - restore_instance: Option<(TomlConfigLoader, ConfigFileControl)>, - armed: bool, -} - -enum ConfigFileCleanup { - Remove(std::path::PathBuf), - Restore { - path: std::path::PathBuf, - contents: Vec, - }, -} - -impl ConfigFileCleanup { - fn apply(self) { - match self { - Self::Remove(config_file) => { - let _ = std::fs::remove_file(config_file); - } - Self::Restore { path, contents } => { - if let Err(e) = std::fs::write(&path, contents) { - tracing::warn!("failed to restore config file {}: {}", path.display(), e); - } - } - } - } -} - -impl StartedInstanceCleanup { - fn new( - manager: Arc, - config_file_cleanup: Option, - restore_instance: Option<(TomlConfigLoader, ConfigFileControl)>, - ) -> Self { - Self { - manager, - started_inst_id: None, - config_file_cleanup, - restore_instance, - armed: true, - } - } - - fn mark_started(&mut self, inst_id: uuid::Uuid) { - self.started_inst_id = Some(inst_id); - } - - fn disarm(&mut self) { - self.armed = false; - } -} - -impl Drop for StartedInstanceCleanup { - fn drop(&mut self) { - if self.armed { - if let Some(inst_id) = self.started_inst_id { - let _ = self.manager.delete_network_instance(vec![inst_id]); - } - if let Some(config_file_cleanup) = self.config_file_cleanup.take() { - config_file_cleanup.apply(); - } - if let Some((cfg, control)) = self.restore_instance.take() - && let Err(e) = self.manager.run_network_instance(cfg, true, control) - { - tracing::warn!("failed to restore overwritten instance: {}", e); - } - } - } -} - -#[async_trait::async_trait] -impl WebClientService for InstanceManageRpcService { - type Controller = BaseController; - - async fn validate_config( - &self, - _: BaseController, - req: ValidateConfigRequest, - ) -> Result { - let toml_config = req.config.unwrap_or_default().gen_config()?.dump(); - Ok(ValidateConfigResponse { toml_config }) - } - - async fn run_network_instance( - &self, - _: BaseController, - req: RunNetworkInstanceRequest, - ) -> Result { - if req.config.is_none() { - return Err(anyhow::anyhow!("config is required").into()); - } - let cfg = req.config.unwrap().gen_config()?; - let mut effective_id = cfg.get_id(); - if let Some(inst_id) = req.inst_id { - effective_id = inst_id.into(); - cfg.set_id(effective_id); - } - let requested_source = ConfigSource::from_rpc(req.source); - let resp = RunNetworkInstanceResponse { - inst_id: Some(effective_id.into()), - }; - let _mutation_guard = self.remote_mutation_lock.lock().await; - let managed_remote = self.hooks.manages_remote_config_instances(); - - let mut overwrite_existing = false; - let mut restore_instance = None; - let mut control = - if let Some(control) = self.manager.get_instance_config_control(&effective_id) { - let existing_source = self - .manager - .get_instance_network_config_source(&effective_id); - let error_msg = self - .manager - .get_network_info(&effective_id) - .await - .and_then(|i| i.error_msg) - .unwrap_or_default(); - - if !req.overwrite && error_msg.is_empty() { - return Ok(resp); - } - if managed_remote { - ensure_remote_overwritable(effective_id, &control).await?; - } else { - ensure_overwritable(effective_id, &control).await?; - } - - cfg.set_network_config_source(requested_source.or(existing_source)); - overwrite_existing = true; - restore_instance = self - .manager - .get_instance_config(&effective_id) - .map(|cfg| (cfg, control.clone())); - control.clone() - } else if let Some(config_dir) = self.manager.get_config_dir() { - cfg.set_network_config_source(requested_source); - ConfigFileControl::new( - Some(config_dir.join(format!("{}.toml", effective_id))), - ConfigFilePermission::default(), - ) - } else { - cfg.set_network_config_source(requested_source); - ConfigFileControl::new(None, ConfigFilePermission::default()) - }; - - if let Err(e) = self.hooks.pre_run_network_instance(&cfg).await { - return Err(anyhow::anyhow!("pre-run hook failed: {}", e).into()); - } - - if overwrite_existing { - if managed_remote { - ensure_remote_overwritable(effective_id, &control).await?; - } else { - ensure_overwritable(effective_id, &control).await?; - } - } - - let mut config_file_cleanup = None; - if !control.is_read_only() - && let Some(config_file) = control.path.as_ref() - { - let cleanup = if config_file.exists() { - match std::fs::read(config_file) { - Ok(contents) => Some(ConfigFileCleanup::Restore { - path: config_file.clone(), - contents, - }), - Err(e) => { - return Err(anyhow::anyhow!( - "failed to backup config file {} before overwrite: {}", - config_file.display(), - e - ) - .into()); - } - } - } else { - Some(ConfigFileCleanup::Remove(config_file.clone())) - }; - match std::fs::write(config_file, cfg.dump()) { - Ok(()) => { - config_file_cleanup = cleanup; - } - Err(e) => { - tracing::warn!( - "failed to write config file {}: {}", - config_file.display(), - e - ); - control.set_read_only(true); - } - } - } - - let mut started_instance = StartedInstanceCleanup::new( - self.manager.clone(), - config_file_cleanup, - restore_instance, - ); - - if overwrite_existing { - self.manager.delete_network_instance(vec![effective_id])?; - } - - if let Err(e) = self.manager.run_network_instance(cfg, true, control) { - return Err(e.into()); - } - started_instance.mark_started(effective_id); - println!("instance {} started", effective_id); - - if let Err(e) = self.hooks.post_run_network_instance(&effective_id).await { - if managed_remote { - return Err(anyhow::anyhow!("post-run hook failed: {}", e).into()); - } - tracing::warn!("post-run hook failed: {}", e); - } - started_instance.disarm(); - - Ok(resp) - } - - async fn retain_network_instance( - &self, - _: BaseController, - req: RetainNetworkInstanceRequest, - ) -> Result { - let _mutation_guard = self.remote_mutation_lock.lock().await; - if !self.hooks.manages_remote_config_instances() { - let remain = self - .manager - .retain_network_instance(req.inst_ids.into_iter().map(Into::into).collect())?; - println!("instance {:?} retained", remain); - return Ok(RetainNetworkInstanceResponse { - remain_inst_ids: remain.iter().map(|item| (*item).into()).collect(), - }); - } - - let mut retain_id_set = req - .inst_ids - .into_iter() - .map(Into::into) - .collect::>(); - let mut removed_ids = Vec::new(); - for (instance_id, control) in self - .manager - .iter() - .map(|instance| (*instance.key(), instance.get_config_file_control().clone())) - .collect::>() - { - if retain_id_set.contains(&instance_id) { - continue; - } - if is_remote_removable(&control).await { - removed_ids.push(instance_id); - } else { - retain_id_set.insert(instance_id); - } - } - let remain = self - .manager - .retain_network_instance(retain_id_set.into_iter().collect())?; - println!("instance {:?} retained", remain); - if let Err(e) = self.hooks.post_remove_network_instances(&removed_ids).await { - return Err(anyhow::anyhow!("post-remove hook failed: {}", e).into()); - } - Ok(RetainNetworkInstanceResponse { - remain_inst_ids: remain.iter().map(|item| (*item).into()).collect(), - }) - } - - async fn collect_network_info( - &self, - _: BaseController, - req: CollectNetworkInfoRequest, - ) -> Result { - let mut ret = NetworkInstanceRunningInfoMap { - map: self - .manager - .collect_network_infos() - .await? - .into_iter() - .map(|(k, v)| (k.to_string(), v)) - .collect(), - }; - let include_inst_ids = req - .inst_ids - .iter() - .cloned() - .map(|id| id.to_string()) - .collect::>(); - if !include_inst_ids.is_empty() { - let mut to_remove = Vec::new(); - for (k, _) in ret.map.iter() { - if !include_inst_ids.contains(k) { - to_remove.push(k.clone()); - } - } - - for k in to_remove { - ret.map.remove(&k); - } - } - Ok(CollectNetworkInfoResponse { info: Some(ret) }) - } - - // rpc ListNetworkInstance(ListNetworkInstanceRequest) returns (ListNetworkInstanceResponse) {} - async fn list_network_instance( - &self, - _: BaseController, - _: ListNetworkInstanceRequest, - ) -> Result { - Ok(ListNetworkInstanceResponse { - inst_ids: self - .manager - .list_network_instance_ids() - .into_iter() - .map(Into::into) - .collect(), - }) - } - - // rpc DeleteNetworkInstance(DeleteNetworkInstanceRequest) returns (DeleteNetworkInstanceResponse) {} - async fn delete_network_instance( - &self, - _: BaseController, - req: DeleteNetworkInstanceRequest, - ) -> Result { - let _mutation_guard = self.remote_mutation_lock.lock().await; - let inst_ids: HashSet = req.inst_ids.into_iter().map(Into::into).collect(); - - if !self.hooks.manages_remote_config_instances() { - let hook_ids: Vec = inst_ids.iter().cloned().collect(); - let inst_ids = self - .manager - .iter() - .filter(|v| inst_ids.contains(v.key())) - .filter(|v| v.get_config_file_control().is_deletable()) - .map(|v| *v.key()) - .collect::>(); - let config_files = inst_ids - .iter() - .filter_map(|id| { - self.manager - .get_instance_config_control(id) - .and_then(|control| control.path) - }) - .collect::>(); - let remain_inst_ids = self.manager.delete_network_instance(inst_ids)?; - println!("instance {:?} retained", remain_inst_ids); - - if let Err(e) = self.hooks.post_remove_network_instances(&hook_ids).await { - tracing::warn!("post-remove hook failed: {}", e); - } - - for config_file in config_files { - if let Err(e) = std::fs::remove_file(&config_file) { - tracing::warn!( - "failed to remove config file {}: {}", - config_file.display(), - e - ); - } - } - return Ok(DeleteNetworkInstanceResponse { - remain_inst_ids: remain_inst_ids.into_iter().map(Into::into).collect(), - }); - } - - let mut deletable_inst_ids = Vec::new(); - for (instance_id, control) in self - .manager - .iter() - .filter(|v| inst_ids.contains(v.key())) - .map(|instance| (*instance.key(), instance.get_config_file_control().clone())) - .collect::>() - { - if is_remote_removable(&control).await { - deletable_inst_ids.push(instance_id); - } - } - let inst_ids = deletable_inst_ids; - let config_files = inst_ids - .iter() - .filter_map(|id| { - self.manager - .get_instance_config_control(id) - .and_then(|control| control.path) - }) - .collect::>(); - let hook_ids = inst_ids.clone(); - let remain_inst_ids = self.manager.delete_network_instance(inst_ids)?; - println!("instance {:?} retained", remain_inst_ids); - - if let Err(e) = self.hooks.post_remove_network_instances(&hook_ids).await { - return Err(anyhow::anyhow!("post-remove hook failed: {}", e).into()); - } - - for config_file in config_files { - if ConfigFileControl::from_path(config_file.clone()) - .await - .is_read_only() - { - continue; - } - if let Err(e) = std::fs::remove_file(&config_file) { - tracing::warn!( - "failed to remove config file {}: {}", - config_file.display(), - e - ); - } - } - Ok(DeleteNetworkInstanceResponse { - remain_inst_ids: remain_inst_ids.into_iter().map(Into::into).collect(), - }) - } - - async fn get_network_instance_config( - &self, - _: BaseController, - req: GetNetworkInstanceConfigRequest, - ) -> Result { - let inst_id: uuid::Uuid = req - .inst_id - .ok_or_else(|| anyhow::anyhow!("instance id is required"))? - .into(); - - let control = self - .manager - .get_instance_config_control(&inst_id) - .ok_or_else(|| anyhow::anyhow!("instance config control not found"))?; - - if control.is_read_only() { - return Err(anyhow::anyhow!( - "Configuration for instance {} is read-only (uses environment variables) and cannot be retrieved via API. \ - Please access the configuration file directly on the file system.", - inst_id - ) - .into()); - } - - let config = self - .manager - .get_instance_service(&inst_id) - .ok_or_else(|| anyhow::anyhow!("instance service not found"))? - .get_config_service() - .get_config(BaseController::default(), GetConfigRequest::default()) - .await? - .config; - Ok(GetNetworkInstanceConfigResponse { - config, - source: self - .manager - .get_instance_network_config_source(&inst_id) - .unwrap_or(ConfigSource::User) - .to_rpc(), - }) - } - - async fn list_network_instance_meta( - &self, - _: BaseController, - req: ListNetworkInstanceMetaRequest, - ) -> Result { - let mut metas = Vec::with_capacity(req.inst_ids.len()); - for inst_id in req.inst_ids { - let inst_id: uuid::Uuid = (inst_id).into(); - let Some(control) = self.manager.get_instance_config_control(&inst_id) else { - continue; - }; - let Some(network_name) = self.manager.get_network_name(&inst_id) else { - continue; - }; - let Some(instance_name) = self.manager.get_instance_name(&inst_id) else { - continue; - }; - let meta = NetworkMeta { - inst_id: Some(inst_id.into()), - network_name, - config_permission: control.permission.into(), - instance_name, - source: self - .manager - .get_instance_network_config_source(&inst_id) - .unwrap_or(ConfigSource::User) - .to_rpc(), - }; - metas.push(meta); - } - Ok(ListNetworkInstanceMetaResponse { metas }) - } -} - -#[cfg(test)] -mod tests { - use super::*; - use crate::common::config::TomlConfigLoader; - use crate::proto::api::manage::{NetworkConfig, NetworkingMethod}; - use crate::web_client::DefaultHooks; - use std::{path::PathBuf, sync::Mutex}; - use uuid::Uuid; - - #[derive(Default)] - struct RecordingHooks { - run_ids: Mutex>, - removed_ids: Mutex>>, - reject_pre_run: bool, - reject_post_run: bool, - reject_post_remove: bool, - readonly_on_pre_run: Mutex>, - readonly_on_post_remove: Mutex>, - } - - #[async_trait::async_trait] - impl WebClientHooks for RecordingHooks { - fn manages_remote_config_instances(&self) -> bool { - true - } - - async fn pre_run_network_instance(&self, _cfg: &TomlConfigLoader) -> Result<(), String> { - for path in self.readonly_on_pre_run.lock().unwrap().drain(..) { - set_file_readonly(&path, true); - } - if self.reject_pre_run { - Err("pre-run rejected".to_string()) - } else { - Ok(()) - } - } - - async fn post_run_network_instance(&self, _id: &Uuid) -> Result<(), String> { - if self.reject_post_run { - Err("post-run rejected".to_string()) - } else { - self.run_ids.lock().unwrap().push(*_id); - Ok(()) - } - } - - async fn post_remove_network_instances(&self, ids: &[Uuid]) -> Result<(), String> { - if self.reject_post_remove { - return Err("post-remove rejected".to_string()); - } - self.removed_ids.lock().unwrap().push(ids.to_vec()); - for path in self.readonly_on_post_remove.lock().unwrap().drain(..) { - set_file_readonly(&path, true); - } - Ok(()) - } - } - - fn temp_config_path(test_name: &str) -> PathBuf { - std::env::temp_dir().join(format!( - "easytier-instance-manage-{}-{}.toml", - test_name, - Uuid::new_v4() - )) - } - - fn set_file_readonly(path: &PathBuf, readonly: bool) { - let mut permissions = std::fs::metadata(path).unwrap().permissions(); - permissions.set_readonly(readonly); - std::fs::set_permissions(path, permissions).unwrap(); - } - - fn cleanup_temp_config(path: &PathBuf) { - if path.exists() { - set_file_readonly(path, false); - let _ = std::fs::remove_file(path); - } - } - - #[tokio::test] - async fn retain_network_instance_preserves_protected_and_reports_actual_removals() { - let manager = Arc::new(NetworkInstanceManager::new()); - let hooks = Arc::new(RecordingHooks::default()); - let service = InstanceManageRpcService::new(manager.clone(), hooks.clone()); - - let no_delete_id = manager - .run_network_instance( - TomlConfigLoader::new_from_str("listeners = []").unwrap(), - false, - ConfigFileControl::STATIC_CONFIG, - ) - .unwrap(); - let read_only_id = manager - .run_network_instance( - TomlConfigLoader::new_from_str("listeners = []").unwrap(), - false, - ConfigFileControl::new( - None, - ConfigFilePermission::default().with_flag(ConfigFilePermission::READ_ONLY), - ), - ) - .unwrap(); - let stale_readonly_path = temp_config_path("retain"); - std::fs::write(&stale_readonly_path, "listeners = []").unwrap(); - let stale_readonly_id = manager - .run_network_instance( - TomlConfigLoader::new_from_str("listeners = []").unwrap(), - false, - ConfigFileControl::new( - Some(stale_readonly_path.clone()), - ConfigFilePermission::default(), - ), - ) - .unwrap(); - set_file_readonly(&stale_readonly_path, true); - let deletable_id = manager - .run_network_instance( - TomlConfigLoader::new_from_str("listeners = []").unwrap(), - false, - ConfigFileControl::new(None, ConfigFilePermission::default()), - ) - .unwrap(); - let removed_id = manager - .run_network_instance( - TomlConfigLoader::new_from_str("listeners = []").unwrap(), - false, - ConfigFileControl::new(None, ConfigFilePermission::default()), - ) - .unwrap(); - - let response = service - .retain_network_instance( - BaseController::default(), - RetainNetworkInstanceRequest { - inst_ids: vec![deletable_id.into()], - }, - ) - .await - .unwrap(); - let remain_ids = response - .remain_inst_ids - .into_iter() - .map(Into::into) - .collect::>(); - - assert_eq!( - remain_ids, - HashSet::from([no_delete_id, read_only_id, stale_readonly_id, deletable_id]) - ); - assert_eq!( - manager - .list_network_instance_ids() - .into_iter() - .collect::>(), - HashSet::from([no_delete_id, read_only_id, stale_readonly_id, deletable_id]) - ); - assert_eq!( - hooks.removed_ids.lock().unwrap().as_slice(), - &[vec![removed_id]] - ); - cleanup_temp_config(&stale_readonly_path); - } - - #[tokio::test] - async fn retain_network_instance_with_default_hooks_keeps_existing_api_behavior() { - let manager = Arc::new(NetworkInstanceManager::new()); - let service = InstanceManageRpcService::new(manager.clone(), Arc::new(DefaultHooks)); - - let protected_id = manager - .run_network_instance( - TomlConfigLoader::new_from_str("listeners = []").unwrap(), - false, - ConfigFileControl::STATIC_CONFIG, - ) - .unwrap(); - let retained_id = manager - .run_network_instance( - TomlConfigLoader::new_from_str("listeners = []").unwrap(), - false, - ConfigFileControl::new(None, ConfigFilePermission::default()), - ) - .unwrap(); - - let response = service - .retain_network_instance( - BaseController::default(), - RetainNetworkInstanceRequest { - inst_ids: vec![retained_id.into()], - }, - ) - .await - .unwrap(); - let remain_ids = response - .remain_inst_ids - .into_iter() - .map(Into::into) - .collect::>(); - - assert_eq!(remain_ids, HashSet::from([retained_id])); - assert!(!manager.list_network_instance_ids().contains(&protected_id)); - } - - #[tokio::test] - async fn retain_network_instance_reports_post_remove_state_failures() { - let manager = Arc::new(NetworkInstanceManager::new()); - let hooks = Arc::new(RecordingHooks { - reject_post_remove: true, - ..Default::default() - }); - let service = InstanceManageRpcService::new(manager.clone(), hooks.clone()); - let _removed_id = manager - .run_network_instance( - TomlConfigLoader::new_from_str("listeners = []").unwrap(), - false, - ConfigFileControl::new(None, ConfigFilePermission::default()), - ) - .unwrap(); - - let result = service - .retain_network_instance( - BaseController::default(), - RetainNetworkInstanceRequest { inst_ids: vec![] }, - ) - .await; - - assert!(result.is_err()); - assert!(manager.list_network_instance_ids().is_empty()); - assert!(hooks.removed_ids.lock().unwrap().is_empty()); - } - - #[tokio::test] - async fn delete_network_instance_preserves_protected_and_reports_actual_removals() { - let manager = Arc::new(NetworkInstanceManager::new()); - let hooks = Arc::new(RecordingHooks::default()); - let service = InstanceManageRpcService::new(manager.clone(), hooks.clone()); - - let no_delete_id = manager - .run_network_instance( - TomlConfigLoader::new_from_str("listeners = []").unwrap(), - false, - ConfigFileControl::STATIC_CONFIG, - ) - .unwrap(); - let read_only_id = manager - .run_network_instance( - TomlConfigLoader::new_from_str("listeners = []").unwrap(), - false, - ConfigFileControl::new( - None, - ConfigFilePermission::default().with_flag(ConfigFilePermission::READ_ONLY), - ), - ) - .unwrap(); - let stale_readonly_path = temp_config_path("delete"); - std::fs::write(&stale_readonly_path, "listeners = []").unwrap(); - let stale_readonly_id = manager - .run_network_instance( - TomlConfigLoader::new_from_str("listeners = []").unwrap(), - false, - ConfigFileControl::new( - Some(stale_readonly_path.clone()), - ConfigFilePermission::default(), - ), - ) - .unwrap(); - set_file_readonly(&stale_readonly_path, true); - let removed_id = manager - .run_network_instance( - TomlConfigLoader::new_from_str("listeners = []").unwrap(), - false, - ConfigFileControl::new(None, ConfigFilePermission::default()), - ) - .unwrap(); - - let response = service - .delete_network_instance( - BaseController::default(), - DeleteNetworkInstanceRequest { - inst_ids: vec![ - no_delete_id.into(), - read_only_id.into(), - stale_readonly_id.into(), - removed_id.into(), - ], - }, - ) - .await - .unwrap(); - let remain_ids = response - .remain_inst_ids - .into_iter() - .map(Into::into) - .collect::>(); - - assert_eq!( - remain_ids, - HashSet::from([no_delete_id, read_only_id, stale_readonly_id]) - ); - assert_eq!( - manager - .list_network_instance_ids() - .into_iter() - .collect::>(), - HashSet::from([no_delete_id, read_only_id, stale_readonly_id]) - ); - assert_eq!( - hooks.removed_ids.lock().unwrap().as_slice(), - &[vec![removed_id]] - ); - cleanup_temp_config(&stale_readonly_path); - } - - #[tokio::test] - async fn delete_network_instance_with_default_hooks_keeps_existing_api_behavior() { - let manager = Arc::new(NetworkInstanceManager::new()); - let service = InstanceManageRpcService::new(manager.clone(), Arc::new(DefaultHooks)); - - let no_delete_id = manager - .run_network_instance( - TomlConfigLoader::new_from_str("listeners = []").unwrap(), - false, - ConfigFileControl::new( - None, - ConfigFilePermission::default().with_flag(ConfigFilePermission::NO_DELETE), - ), - ) - .unwrap(); - let read_only_id = manager - .run_network_instance( - TomlConfigLoader::new_from_str("listeners = []").unwrap(), - false, - ConfigFileControl::new( - None, - ConfigFilePermission::default().with_flag(ConfigFilePermission::READ_ONLY), - ), - ) - .unwrap(); - - let response = service - .delete_network_instance( - BaseController::default(), - DeleteNetworkInstanceRequest { - inst_ids: vec![no_delete_id.into(), read_only_id.into()], - }, - ) - .await - .unwrap(); - let remain_ids = response - .remain_inst_ids - .into_iter() - .map(Into::into) - .collect::>(); - - assert_eq!(remain_ids, HashSet::from([no_delete_id])); - assert!(!manager.list_network_instance_ids().contains(&read_only_id)); - } - - #[tokio::test] - async fn delete_network_instance_preserves_config_file_that_becomes_readonly_during_hook() { - let manager = Arc::new(NetworkInstanceManager::new()); - let config_path = temp_config_path("delete-hook"); - std::fs::write(&config_path, "listeners = []").unwrap(); - let hooks = Arc::new(RecordingHooks { - readonly_on_post_remove: Mutex::new(vec![config_path.clone()]), - ..Default::default() - }); - let service = InstanceManageRpcService::new(manager.clone(), hooks.clone()); - let removed_id = manager - .run_network_instance( - TomlConfigLoader::new_from_str("listeners = []").unwrap(), - false, - ConfigFileControl::new(Some(config_path.clone()), ConfigFilePermission::default()), - ) - .unwrap(); - - let response = service - .delete_network_instance( - BaseController::default(), - DeleteNetworkInstanceRequest { - inst_ids: vec![removed_id.into()], - }, - ) - .await - .unwrap(); - - assert!(response.remain_inst_ids.is_empty()); - assert!(manager.list_network_instance_ids().is_empty()); - assert_eq!( - hooks.removed_ids.lock().unwrap().as_slice(), - &[vec![removed_id]] - ); - assert!(config_path.exists()); - cleanup_temp_config(&config_path); - } - - #[tokio::test] - async fn run_network_instance_rejects_overwrite_of_no_delete_instance() { - let manager = Arc::new(NetworkInstanceManager::new()); - let hooks = Arc::new(RecordingHooks::default()); - let service = InstanceManageRpcService::new(manager.clone(), hooks.clone()); - - let protected_id = manager - .run_network_instance( - TomlConfigLoader::new_from_str("listeners = []").unwrap(), - false, - ConfigFileControl::new( - None, - ConfigFilePermission::default().with_flag(ConfigFilePermission::NO_DELETE), - ), - ) - .unwrap(); - - let result = service - .run_network_instance( - BaseController::default(), - RunNetworkInstanceRequest { - inst_id: Some(protected_id.into()), - config: Some(NetworkConfig { - networking_method: Some(NetworkingMethod::Standalone as i32), - listener_urls: Vec::new(), - ..Default::default() - }), - overwrite: true, - source: Default::default(), - }, - ) - .await; - - assert!(result.is_err()); - assert_eq!( - manager - .list_network_instance_ids() - .into_iter() - .collect::>(), - HashSet::from([protected_id]) - ); - assert!(hooks.removed_ids.lock().unwrap().is_empty()); - } - - #[tokio::test] - async fn run_network_instance_preserves_existing_when_path_becomes_readonly_during_pre_run() { - let manager = Arc::new(NetworkInstanceManager::new()); - let config_path = temp_config_path("overwrite-pre-run"); - std::fs::write(&config_path, "listeners = []").unwrap(); - let hooks = Arc::new(RecordingHooks { - readonly_on_pre_run: Mutex::new(vec![config_path.clone()]), - ..Default::default() - }); - let service = InstanceManageRpcService::new(manager.clone(), hooks.clone()); - - let existing_id = manager - .run_network_instance( - TomlConfigLoader::new_from_str("inst_name = \"existing\"\nlisteners = []").unwrap(), - false, - ConfigFileControl::new(Some(config_path.clone()), ConfigFilePermission::default()), - ) - .unwrap(); - - let result = service - .run_network_instance( - BaseController::default(), - RunNetworkInstanceRequest { - inst_id: Some(existing_id.into()), - config: Some(NetworkConfig { - network_name: Some("replacement".to_string()), - networking_method: Some(NetworkingMethod::Standalone as i32), - listener_urls: Vec::new(), - ..Default::default() - }), - overwrite: true, - source: Default::default(), - }, - ) - .await; - - assert!(result.is_err()); - assert_eq!( - manager - .list_network_instance_ids() - .into_iter() - .collect::>(), - HashSet::from([existing_id]) - ); - assert!(hooks.removed_ids.lock().unwrap().is_empty()); - cleanup_temp_config(&config_path); - } - - #[tokio::test] - async fn run_network_instance_overwrite_reports_run_without_remove() { - let manager = Arc::new(NetworkInstanceManager::new()); - let hooks = Arc::new(RecordingHooks::default()); - let service = InstanceManageRpcService::new(manager.clone(), hooks.clone()); - let existing_id = manager - .run_network_instance( - TomlConfigLoader::new_from_str("inst_name = \"existing\"\nlisteners = []").unwrap(), - false, - ConfigFileControl::new(None, ConfigFilePermission::default()), - ) - .unwrap(); - - service - .run_network_instance( - BaseController::default(), - RunNetworkInstanceRequest { - inst_id: Some(existing_id.into()), - config: Some(NetworkConfig { - network_name: Some("replacement".to_string()), - networking_method: Some(NetworkingMethod::Standalone as i32), - listener_urls: Vec::new(), - ..Default::default() - }), - overwrite: true, - source: Default::default(), - }, - ) - .await - .unwrap(); - - assert_eq!(hooks.run_ids.lock().unwrap().as_slice(), &[existing_id]); - assert!(hooks.removed_ids.lock().unwrap().is_empty()); - assert_eq!( - manager - .list_network_instance_ids() - .into_iter() - .collect::>(), - HashSet::from([existing_id]) - ); - } - - #[tokio::test] - async fn run_network_instance_overwrite_post_run_failure_does_not_report_remove() { - let manager = Arc::new(NetworkInstanceManager::new()); - let hooks = Arc::new(RecordingHooks { - reject_post_run: true, - ..Default::default() - }); - let service = InstanceManageRpcService::new(manager.clone(), hooks.clone()); - let existing_id = manager - .run_network_instance( - TomlConfigLoader::new_from_str("inst_name = \"existing\"\nlisteners = []").unwrap(), - false, - ConfigFileControl::new(None, ConfigFilePermission::default()), - ) - .unwrap(); - - let result = service - .run_network_instance( - BaseController::default(), - RunNetworkInstanceRequest { - inst_id: Some(existing_id.into()), - config: Some(NetworkConfig { - network_name: Some("replacement".to_string()), - networking_method: Some(NetworkingMethod::Standalone as i32), - listener_urls: Vec::new(), - ..Default::default() - }), - overwrite: true, - source: Default::default(), - }, - ) - .await; - - assert!(result.is_err()); - assert!(hooks.run_ids.lock().unwrap().is_empty()); - assert!(hooks.removed_ids.lock().unwrap().is_empty()); - } - - #[tokio::test] - async fn run_network_instance_overwrite_post_run_failure_keeps_existing_config_file() { - let manager = Arc::new(NetworkInstanceManager::new()); - let config_path = temp_config_path("overwrite-post-run-failure"); - let original_config = "inst_name = \"existing\"\nlisteners = []"; - std::fs::write(&config_path, original_config).unwrap(); - let hooks = Arc::new(RecordingHooks { - reject_post_run: true, - ..Default::default() - }); - let service = InstanceManageRpcService::new(manager.clone(), hooks.clone()); - let existing_id = manager - .run_network_instance( - TomlConfigLoader::new_from_str("inst_name = \"existing\"\nlisteners = []").unwrap(), - false, - ConfigFileControl::new(Some(config_path.clone()), ConfigFilePermission::default()), - ) - .unwrap(); - - let result = service - .run_network_instance( - BaseController::default(), - RunNetworkInstanceRequest { - inst_id: Some(existing_id.into()), - config: Some(NetworkConfig { - network_name: Some("replacement".to_string()), - networking_method: Some(NetworkingMethod::Standalone as i32), - listener_urls: Vec::new(), - ..Default::default() - }), - overwrite: true, - source: Default::default(), - }, - ) - .await; - - assert!(result.is_err()); - assert!(config_path.exists()); - assert_eq!( - std::fs::read_to_string(&config_path).unwrap(), - original_config - ); - assert_eq!( - manager - .list_network_instance_ids() - .into_iter() - .collect::>(), - HashSet::from([existing_id]) - ); - cleanup_temp_config(&config_path); - } - - #[tokio::test] - async fn run_network_instance_reports_post_run_state_failures() { - let manager = Arc::new(NetworkInstanceManager::new()); - let hooks = Arc::new(RecordingHooks { - reject_post_run: true, - ..Default::default() - }); - let service = InstanceManageRpcService::new(manager.clone(), hooks.clone()); - - let result = service - .run_network_instance( - BaseController::default(), - RunNetworkInstanceRequest { - config: Some(NetworkConfig { - networking_method: Some(NetworkingMethod::Standalone as i32), - listener_urls: Vec::new(), - ..Default::default() - }), - overwrite: true, - source: Default::default(), - ..Default::default() - }, - ) - .await; - - assert!(result.is_err()); - assert!(manager.list_network_instance_ids().is_empty()); - } - - #[tokio::test] - async fn run_network_instance_preserves_existing_instance_when_pre_run_rejects_overwrite() { - let manager = Arc::new(NetworkInstanceManager::new()); - let hooks = Arc::new(RecordingHooks { - reject_pre_run: true, - ..Default::default() - }); - let service = InstanceManageRpcService::new(manager.clone(), hooks.clone()); - - let existing_id = manager - .run_network_instance( - TomlConfigLoader::new_from_str("inst_name = \"existing\"\nlisteners = []").unwrap(), - false, - ConfigFileControl::new(None, ConfigFilePermission::default()), - ) - .unwrap(); - - let result = service - .run_network_instance( - BaseController::default(), - RunNetworkInstanceRequest { - inst_id: Some(existing_id.into()), - config: Some(NetworkConfig { - network_name: Some("replacement".to_string()), - networking_method: Some(NetworkingMethod::Standalone as i32), - listener_urls: Vec::new(), - ..Default::default() - }), - overwrite: true, - source: Default::default(), - }, - ) - .await; - - assert!(result.is_err()); - assert_eq!( - manager - .list_network_instance_ids() - .into_iter() - .collect::>(), - HashSet::from([existing_id]) - ); - assert!(hooks.removed_ids.lock().unwrap().is_empty()); - } -} diff --git a/easytier/src/rpc_service/json_rpc.rs b/easytier/src/rpc_service/json_rpc.rs deleted file mode 100644 index 7f63726d..00000000 --- a/easytier/src/rpc_service/json_rpc.rs +++ /dev/null @@ -1,232 +0,0 @@ -use std::sync::Arc; - -use crate::{ - instance_manager::NetworkInstanceManager, - proto::{ - api::{ - config::ConfigRpc, - instance::{ - AclManageRpc, ConnectorManageRpc, CredentialManageRpc, MappedListenerManageRpc, - PeerManageRpc, PortForwardManageRpc, StatsRpc, TcpProxyRpc, VpnPortalRpc, - }, - logger::LoggerRpc, - }, - peer_rpc::PeerCenterRpc, - rpc_types::{ - controller::BaseController, - error::{Error, Result}, - }, - }, - rpc_service::{ - acl_manage::AclManageRpcService, config::ConfigRpcService, - connector_manage::ConnectorManageRpcService, credential_manage::CredentialManageRpcService, - logger::LoggerRpcService, mapped_listener_manage::MappedListenerManageRpcService, - peer_center::PeerCenterManageRpcService, peer_manage::PeerManageRpcService, - port_forward_manage::PortForwardManageRpcService, proxy::TcpProxyRpcService, - stats::StatsRpcService, vpn_portal::VpnPortalRpcService, - }, -}; - -const INSTANCE_MANAGEMENT_SERVICE: &str = "api.manage.WebClientService"; - -fn service_not_exposed(service_name: &str) -> Error { - anyhow::anyhow!( - "service {} is not exposed through FFI/JNI generic RPC", - service_name - ) - .into() -} - -fn tcp_proxy_domain(domain_name: Option<&str>) -> Result<&'static str> { - match domain_name { - None | Some("") => Ok("tcp"), - Some("tcp") => Ok("tcp"), - Some("kcp_src") => Ok("kcp_src"), - Some("kcp_dst") => Ok("kcp_dst"), - Some("quic_src") => Ok("quic_src"), - Some("quic_dst") => Ok("quic_dst"), - Some(domain) => { - Err(anyhow::anyhow!("invalid TcpProxyRpcService domain_name: {}", domain).into()) - } - } -} - -pub async fn call_json_rpc( - instance_manager: &Arc, - service_name: &str, - method_name: &str, - domain_name: Option<&str>, - payload: serde_json::Value, -) -> Result { - let ctrl = BaseController::default(); - - match service_name { - INSTANCE_MANAGEMENT_SERVICE => Err(service_not_exposed(service_name)), - "api.instance.PeerManageRpcService" => { - PeerManageRpcService::new(instance_manager.clone()) - .json_call_method(ctrl, method_name, payload) - .await - } - "api.instance.PeerCenterManageRpcService" => { - PeerCenterManageRpcService::new(instance_manager.clone()) - .json_call_method(ctrl, method_name, payload) - .await - } - "api.instance.ConnectorManageRpcService" => { - ConnectorManageRpcService::new(instance_manager.clone()) - .json_call_method(ctrl, method_name, payload) - .await - } - "api.instance.MappedListenerManageRpcService" => { - MappedListenerManageRpcService::new(instance_manager.clone()) - .json_call_method(ctrl, method_name, payload) - .await - } - "api.instance.VpnPortalRpcService" => { - VpnPortalRpcService::new(instance_manager.clone()) - .json_call_method(ctrl, method_name, payload) - .await - } - "api.instance.TcpProxyRpcService" => { - TcpProxyRpcService::new(instance_manager.clone(), tcp_proxy_domain(domain_name)?) - .json_call_method(ctrl, method_name, payload) - .await - } - "api.instance.AclManageRpcService" => { - AclManageRpcService::new(instance_manager.clone()) - .json_call_method(ctrl, method_name, payload) - .await - } - "api.instance.PortForwardManageRpcService" => { - PortForwardManageRpcService::new(instance_manager.clone()) - .json_call_method(ctrl, method_name, payload) - .await - } - "api.instance.StatsRpcService" => { - StatsRpcService::new(instance_manager.clone()) - .json_call_method(ctrl, method_name, payload) - .await - } - "api.instance.CredentialManageRpcService" => { - CredentialManageRpcService::new(instance_manager.clone()) - .json_call_method(ctrl, method_name, payload) - .await - } - "api.logger.LoggerRpcService" => { - LoggerRpcService - .json_call_method(ctrl, method_name, payload) - .await - } - "api.config.ConfigRpcService" => { - ConfigRpcService::new(instance_manager.clone()) - .json_call_method(ctrl, method_name, payload) - .await - } - _ => Err(Error::InvalidServiceKey( - service_name.to_string(), - service_name.to_string(), - )), - } -} - -#[cfg(test)] -mod tests { - use super::*; - - fn manager() -> Arc { - Arc::new(NetworkInstanceManager::new()) - } - - #[tokio::test] - async fn logger_json_rpc_succeeds() { - let response = call_json_rpc( - &manager(), - "api.logger.LoggerRpcService", - "get_logger_config", - None, - serde_json::json!({}), - ) - .await - .unwrap(); - - assert!(response.get("level").is_some()); - } - - #[tokio::test] - async fn json_rpc_rejects_unknown_service() { - let err = call_json_rpc( - &manager(), - "api.unknown.Service", - "get_logger_config", - None, - serde_json::json!({}), - ) - .await - .unwrap_err(); - - assert!(matches!(err, Error::InvalidServiceKey(_, _))); - } - - #[tokio::test] - async fn json_rpc_rejects_instance_management_service() { - let err = call_json_rpc( - &manager(), - INSTANCE_MANAGEMENT_SERVICE, - "list_network_instance", - None, - serde_json::json!({}), - ) - .await - .unwrap_err(); - - assert!(err.to_string().contains("not exposed")); - } - - #[tokio::test] - async fn json_rpc_rejects_unknown_method() { - let err = call_json_rpc( - &manager(), - "api.logger.LoggerRpcService", - "missing_method", - None, - serde_json::json!({}), - ) - .await - .unwrap_err(); - - assert!(matches!(err, Error::InvalidMethodIndex(0, _))); - } - - #[tokio::test] - async fn json_rpc_rejects_invalid_payload() { - let err = call_json_rpc( - &manager(), - "api.logger.LoggerRpcService", - "get_logger_config", - None, - serde_json::json!([]), - ) - .await - .unwrap_err(); - - assert!(matches!(err, Error::MalformatRpcPacket(_))); - } - - #[tokio::test] - async fn json_rpc_rejects_invalid_tcp_proxy_domain() { - let err = call_json_rpc( - &manager(), - "api.instance.TcpProxyRpcService", - "list_tcp_proxy_entry", - Some("bad"), - serde_json::json!({}), - ) - .await - .unwrap_err(); - - assert!( - err.to_string() - .contains("invalid TcpProxyRpcService domain_name") - ); - } -} diff --git a/easytier/src/rpc_service/logger.rs b/easytier/src/rpc_service/logger.rs index f46e8599..a0badc89 100644 --- a/easytier/src/rpc_service/logger.rs +++ b/easytier/src/rpc_service/logger.rs @@ -1,105 +1,14 @@ -use std::sync::{Mutex, OnceLock, mpsc::Sender}; - -use crate::proto::{ - api::logger::{ - GetLoggerConfigRequest, GetLoggerConfigResponse, LogLevel, LoggerRpc, - SetLoggerConfigRequest, SetLoggerConfigResponse, - }, - rpc_types::{self, controller::BaseController}, -}; - -pub static LOGGER_LEVEL_SENDER: std::sync::OnceLock>> = OnceLock::new(); -pub static CURRENT_LOG_LEVEL: std::sync::OnceLock> = OnceLock::new(); +use easytier_core::management::LoggerControl; #[derive(Clone, Default)] -pub struct LoggerRpcService; +pub struct NativeLoggerControl; -impl LoggerRpcService { - fn log_level_to_string(level: LogLevel) -> String { - match level { - LogLevel::Disabled => "off".to_string(), - LogLevel::Error => "error".to_string(), - LogLevel::Warning => "warn".to_string(), - LogLevel::Info => "info".to_string(), - LogLevel::Debug => "debug".to_string(), - LogLevel::Trace => "trace".to_string(), - } +impl LoggerControl for NativeLoggerControl { + fn set_level(&self, level: &str) -> anyhow::Result<()> { + crate::common::log::set_file_level(level) } - pub fn string_to_log_level(level_str: &str) -> LogLevel { - match level_str.to_lowercase().as_str() { - "off" | "disabled" => LogLevel::Disabled, - "error" => LogLevel::Error, - "warn" | "warning" => LogLevel::Warning, - "info" => LogLevel::Info, - "debug" => LogLevel::Debug, - "trace" => LogLevel::Trace, - _ => LogLevel::Info, // 默认为 Info 级别 - } - } -} - -#[async_trait::async_trait] -impl LoggerRpc for LoggerRpcService { - type Controller = BaseController; - - async fn set_logger_config( - &self, - _: BaseController, - request: SetLoggerConfigRequest, - ) -> Result { - let level_str = Self::log_level_to_string(request.level()); - - // 发送新的日志级别到 logger 重载器 - if let Some(sender) = LOGGER_LEVEL_SENDER.get() { - if let Ok(sender) = sender.lock() { - if let Err(e) = sender.send(level_str) { - tracing::warn!("Failed to send new log level to reloader: {}", e); - return Err(rpc_types::error::Error::ExecutionError(anyhow::anyhow!( - "Failed to update log level: {}", - e - ))); - } - } else { - return Err(rpc_types::error::Error::ExecutionError(anyhow::anyhow!( - "Logger sender is not available" - ))); - } - } else { - return Err(rpc_types::error::Error::ExecutionError(anyhow::anyhow!( - "Logger reloader is not initialized" - ))); - } - - // 更新当前日志级别 - if let Some(current_level) = CURRENT_LOG_LEVEL.get() - && let Ok(mut level) = current_level.lock() - { - *level = Self::log_level_to_string(request.level()); - } - - Ok(SetLoggerConfigResponse {}) - } - - async fn get_logger_config( - &self, - _: BaseController, - _request: GetLoggerConfigRequest, - ) -> Result { - let current_level_str = if let Some(current_level) = CURRENT_LOG_LEVEL.get() { - if let Ok(level) = current_level.lock() { - level.clone() - } else { - "info".to_string() // 默认级别 - } - } else { - "info".to_string() // 默认级别 - }; - - let level = Self::string_to_log_level(¤t_level_str); - - Ok(GetLoggerConfigResponse { - level: level.into(), - }) + fn level(&self) -> String { + crate::common::log::file_level() } } diff --git a/easytier/src/rpc_service/mapped_listener_manage.rs b/easytier/src/rpc_service/mapped_listener_manage.rs deleted file mode 100644 index 39437634..00000000 --- a/easytier/src/rpc_service/mapped_listener_manage.rs +++ /dev/null @@ -1,38 +0,0 @@ -use std::sync::Arc; - -use crate::{ - instance_manager::NetworkInstanceManager, - proto::{ - api::instance::{ - ListMappedListenerRequest, ListMappedListenerResponse, MappedListenerManageRpc, - }, - rpc_types::controller::BaseController, - }, -}; - -#[derive(Clone)] -pub struct MappedListenerManageRpcService { - instance_manager: Arc, -} - -impl MappedListenerManageRpcService { - pub fn new(instance_manager: Arc) -> Self { - Self { instance_manager } - } -} - -#[async_trait::async_trait] -impl MappedListenerManageRpc for MappedListenerManageRpcService { - type Controller = BaseController; - - async fn list_mapped_listener( - &self, - ctrl: Self::Controller, - req: ListMappedListenerRequest, - ) -> crate::proto::rpc_types::error::Result { - super::get_instance_service(&self.instance_manager, &req.instance)? - .get_mapped_listener_manage_service() - .list_mapped_listener(ctrl, req) - .await - } -} diff --git a/easytier/src/rpc_service/mod.rs b/easytier/src/rpc_service/mod.rs index 7ac80ad3..5a79dbb9 100644 --- a/easytier/src/rpc_service/mod.rs +++ b/easytier/src/rpc_service/mod.rs @@ -1,130 +1,9 @@ -mod acl_manage; -mod config; -mod connector_manage; -mod credential_manage; -mod json_rpc; -mod mapped_listener_manage; -mod peer_center; -mod peer_manage; -mod port_forward_manage; -pub(crate) mod protected_port; -mod proxy; -mod stats; -mod vpn_portal; - pub mod api; -pub mod instance_manage; +#[cfg(feature = "management")] pub mod logger; -pub mod remote_client; +#[cfg(feature = "management")] +pub use easytier_core::management::remote_client; +#[cfg(feature = "management")] pub type ApiRpcServer = self::api::ApiRpcServer; -pub use json_rpc::call_json_rpc; - -pub trait InstanceRpcService: Sync + Send { - fn get_peer_manage_service( - &self, - ) -> &dyn crate::proto::api::instance::PeerManageRpc< - Controller = crate::proto::rpc_types::controller::BaseController, - >; - fn get_connector_manage_service( - &self, - ) -> &dyn crate::proto::api::instance::ConnectorManageRpc< - Controller = crate::proto::rpc_types::controller::BaseController, - >; - fn get_mapped_listener_manage_service( - &self, - ) -> &dyn crate::proto::api::instance::MappedListenerManageRpc< - Controller = crate::proto::rpc_types::controller::BaseController, - >; - fn get_vpn_portal_service( - &self, - ) -> &dyn crate::proto::api::instance::VpnPortalRpc< - Controller = crate::proto::rpc_types::controller::BaseController, - >; - fn get_proxy_service( - &self, - client_type: &str, - ) -> Option< - std::sync::Arc< - dyn crate::proto::api::instance::TcpProxyRpc< - Controller = crate::proto::rpc_types::controller::BaseController, - > + Send - + Sync, - >, - >; - fn get_acl_manage_service( - &self, - ) -> &dyn crate::proto::api::instance::AclManageRpc< - Controller = crate::proto::rpc_types::controller::BaseController, - >; - fn get_port_forward_manage_service( - &self, - ) -> &dyn crate::proto::api::instance::PortForwardManageRpc< - Controller = crate::proto::rpc_types::controller::BaseController, - >; - fn get_stats_service( - &self, - ) -> &dyn crate::proto::api::instance::StatsRpc< - Controller = crate::proto::rpc_types::controller::BaseController, - >; - fn get_config_service( - &self, - ) -> &dyn crate::proto::api::config::ConfigRpc< - Controller = crate::proto::rpc_types::controller::BaseController, - >; - fn get_peer_center_service( - &self, - ) -> std::sync::Arc< - dyn crate::proto::peer_rpc::PeerCenterRpc< - Controller = crate::proto::rpc_types::controller::BaseController, - > + Send - + Sync, - >; - fn get_credential_manage_service( - &self, - ) -> &dyn crate::proto::api::instance::CredentialManageRpc< - Controller = crate::proto::rpc_types::controller::BaseController, - >; -} - -fn get_instance_service( - instance_manager: &std::sync::Arc, - identifier: &Option, -) -> Result, anyhow::Error> { - use crate::proto::api; - let selector = identifier.as_ref().and_then(|s| s.selector.as_ref()); - - let id = if let Some(api::instance::instance_identifier::Selector::Id(id)) = selector { - (*id).into() - } else { - let ids = instance_manager - .iter() - .filter(|v| { - if let Some(api::instance::instance_identifier::Selector::InstanceSelector( - selector, - )) = selector - && let Some(name) = selector.name.as_ref() - && v.get_inst_name() != *name - { - return false; - } - true - }) - .map(|v| *v.key()) - .collect::>(); - match ids.len() { - 0 => return Err(anyhow::anyhow!("No instance matches the selector")), - 1 => ids[0], - _ => { - return Err(anyhow::anyhow!( - "{} instances match the selector, please specify the instance ID", - ids.len() - )); - } - } - }; - - instance_manager - .get_instance_service(&id) - .ok_or_else(|| anyhow::anyhow!("Instance not found or API service not available")) -} +pub type ReadOnlyApiRpcServer = self::api::ReadOnlyApiRpcServer; diff --git a/easytier/src/rpc_service/peer_center.rs b/easytier/src/rpc_service/peer_center.rs deleted file mode 100644 index 66be659d..00000000 --- a/easytier/src/rpc_service/peer_center.rs +++ /dev/null @@ -1,108 +0,0 @@ -use std::sync::Arc; - -use crate::{ - instance_manager::NetworkInstanceManager, - proto::{ - peer_rpc::{ - GetGlobalPeerMapRequest, GetGlobalPeerMapResponse, PeerCenterRpc, ReportPeersRequest, - ReportPeersResponse, - }, - rpc_types::controller::BaseController, - }, -}; - -#[derive(Clone)] -pub struct PeerCenterManageRpcService { - instance_manager: Arc, -} - -impl PeerCenterManageRpcService { - pub fn new(instance_manager: Arc) -> Self { - Self { instance_manager } - } -} - -#[async_trait::async_trait] -impl PeerCenterRpc for PeerCenterManageRpcService { - type Controller = BaseController; - - async fn get_global_peer_map( - &self, - ctrl: BaseController, - req: GetGlobalPeerMapRequest, - ) -> crate::proto::rpc_types::error::Result { - let instance_service = - super::get_instance_service(&self.instance_manager, &None).map_err(|e| { - let msg = e.to_string(); - if msg.contains("please specify the instance ID") { - anyhow::anyhow!( - "PeerCenter management RPC cannot select an instance automatically \ - when multiple instances are running; please use an API that allows \ - specifying an instance identifier." - ) - } else { - e - } - })?; - - instance_service - .get_peer_center_service() - .get_global_peer_map(ctrl, req) - .await - } - - async fn report_peers( - &self, - _: BaseController, - _: ReportPeersRequest, - ) -> crate::proto::rpc_types::error::Result { - Err(anyhow::anyhow!("not implemented for management API").into()) - } -} - -#[cfg(test)] -mod tests { - use super::*; - use crate::{ - instance_manager::NetworkInstanceManager, - proto::{ - peer_rpc::{GetGlobalPeerMapRequest, ReportPeersRequest}, - rpc_types::controller::BaseController, - }, - }; - - fn make_service() -> PeerCenterManageRpcService { - PeerCenterManageRpcService::new(Arc::new(NetworkInstanceManager::new())) - } - - #[tokio::test] - async fn get_global_peer_map_errors_when_no_instance() { - let svc = make_service(); - let result = svc - .get_global_peer_map( - BaseController::default(), - GetGlobalPeerMapRequest::default(), - ) - .await; - assert!(result.is_err()); - let msg = result.unwrap_err().to_string(); - assert!( - msg.contains("No instance matches the selector"), - "unexpected error: {msg}" - ); - } - - #[tokio::test] - async fn report_peers_always_returns_error() { - let svc = make_service(); - let result = svc - .report_peers(BaseController::default(), ReportPeersRequest::default()) - .await; - assert!(result.is_err()); - let msg = result.unwrap_err().to_string(); - assert!( - msg.contains("not implemented for management API"), - "unexpected error: {msg}" - ); - } -} diff --git a/easytier/src/rpc_service/peer_manage.rs b/easytier/src/rpc_service/peer_manage.rs deleted file mode 100644 index 2e8369de..00000000 --- a/easytier/src/rpc_service/peer_manage.rs +++ /dev/null @@ -1,116 +0,0 @@ -use std::sync::Arc; - -use crate::{ - instance_manager::NetworkInstanceManager, - proto::{ - api::instance::{ - self, ListPeerRequest, ListPeerResponse, ListPublicIpv6InfoRequest, - ListPublicIpv6InfoResponse, PeerManageRpc, - }, - rpc_types::controller::BaseController, - }, -}; - -#[derive(Clone)] -pub struct PeerManageRpcService { - instance_manager: Arc, -} - -impl PeerManageRpcService { - pub fn new(instance_manager: Arc) -> Self { - Self { instance_manager } - } -} - -#[async_trait::async_trait] -impl PeerManageRpc for PeerManageRpcService { - type Controller = BaseController; - - async fn list_peer( - &self, - ctrl: Self::Controller, - req: ListPeerRequest, - ) -> crate::proto::rpc_types::error::Result { - super::get_instance_service(&self.instance_manager, &req.instance)? - .get_peer_manage_service() - .list_peer(ctrl, req) - .await - } - - async fn list_public_ipv6_info( - &self, - ctrl: Self::Controller, - req: ListPublicIpv6InfoRequest, - ) -> crate::proto::rpc_types::error::Result { - super::get_instance_service(&self.instance_manager, &req.instance)? - .get_peer_manage_service() - .list_public_ipv6_info(ctrl, req) - .await - } - - async fn list_route( - &self, - ctrl: Self::Controller, - req: crate::proto::api::instance::ListRouteRequest, - ) -> crate::proto::rpc_types::error::Result { - super::get_instance_service(&self.instance_manager, &req.instance)? - .get_peer_manage_service() - .list_route(ctrl, req) - .await - } - - async fn dump_route( - &self, - ctrl: Self::Controller, - req: crate::proto::api::instance::DumpRouteRequest, - ) -> crate::proto::rpc_types::error::Result { - super::get_instance_service(&self.instance_manager, &req.instance)? - .get_peer_manage_service() - .dump_route(ctrl, req) - .await - } - - async fn list_foreign_network( - &self, - ctrl: Self::Controller, - req: crate::proto::api::instance::ListForeignNetworkRequest, - ) -> crate::proto::rpc_types::error::Result { - super::get_instance_service(&self.instance_manager, &req.instance)? - .get_peer_manage_service() - .list_foreign_network(ctrl, req) - .await - } - - async fn list_global_foreign_network( - &self, - ctrl: Self::Controller, - req: crate::proto::api::instance::ListGlobalForeignNetworkRequest, - ) -> crate::proto::rpc_types::error::Result { - super::get_instance_service(&self.instance_manager, &req.instance)? - .get_peer_manage_service() - .list_global_foreign_network(ctrl, req) - .await - } - - async fn get_foreign_network_summary( - &self, - ctrl: Self::Controller, - req: crate::proto::api::instance::GetForeignNetworkSummaryRequest, - ) -> crate::proto::rpc_types::error::Result { - super::get_instance_service(&self.instance_manager, &req.instance)? - .get_peer_manage_service() - .get_foreign_network_summary(ctrl, req) - .await - } - - async fn show_node_info( - &self, - ctrl: Self::Controller, - req: crate::proto::api::instance::ShowNodeInfoRequest, - ) -> crate::proto::rpc_types::error::Result { - super::get_instance_service(&self.instance_manager, &req.instance)? - .get_peer_manage_service() - .show_node_info(ctrl, req) - .await - } -} diff --git a/easytier/src/rpc_service/port_forward_manage.rs b/easytier/src/rpc_service/port_forward_manage.rs deleted file mode 100644 index 58726279..00000000 --- a/easytier/src/rpc_service/port_forward_manage.rs +++ /dev/null @@ -1,36 +0,0 @@ -use std::sync::Arc; - -use crate::{ - instance_manager::NetworkInstanceManager, - proto::{ - api::instance::{ListPortForwardRequest, ListPortForwardResponse, PortForwardManageRpc}, - rpc_types::controller::BaseController, - }, -}; - -#[derive(Clone)] -pub struct PortForwardManageRpcService { - instance_manager: Arc, -} - -impl PortForwardManageRpcService { - pub fn new(instance_manager: Arc) -> Self { - Self { instance_manager } - } -} - -#[async_trait::async_trait] -impl PortForwardManageRpc for PortForwardManageRpcService { - type Controller = BaseController; - - async fn list_port_forward( - &self, - ctrl: Self::Controller, - req: ListPortForwardRequest, - ) -> crate::proto::rpc_types::error::Result { - super::get_instance_service(&self.instance_manager, &req.instance)? - .get_port_forward_manage_service() - .list_port_forward(ctrl, req) - .await - } -} diff --git a/easytier/src/rpc_service/protected_port.rs b/easytier/src/rpc_service/protected_port.rs deleted file mode 100644 index b44216b5..00000000 --- a/easytier/src/rpc_service/protected_port.rs +++ /dev/null @@ -1,61 +0,0 @@ -use std::collections::HashMap; -use std::sync::Mutex; - -use once_cell::sync::Lazy; - -static PROTECTED_TCP_PORTS: Lazy>> = - Lazy::new(|| Mutex::new(HashMap::new())); - -pub fn register_protected_tcp_port(port: u16) { - let mut ports = PROTECTED_TCP_PORTS.lock().unwrap(); - *ports.entry(port).or_default() += 1; -} - -pub fn unregister_protected_tcp_port(port: u16) { - let mut ports = PROTECTED_TCP_PORTS.lock().unwrap(); - if let Some(ref_count) = ports.get_mut(&port) { - *ref_count -= 1; - if *ref_count == 0 { - ports.remove(&port); - } - } -} - -pub fn is_protected_tcp_port(port: u16) -> bool { - PROTECTED_TCP_PORTS.lock().unwrap().contains_key(&port) -} - -#[cfg(test)] -pub fn clear_protected_tcp_ports_for_test() { - PROTECTED_TCP_PORTS.lock().unwrap().clear(); -} - -#[cfg(test)] -mod tests { - use super::{ - clear_protected_tcp_ports_for_test, is_protected_tcp_port, register_protected_tcp_port, - unregister_protected_tcp_port, - }; - - #[test] - fn protected_tcp_port_registry_is_ref_counted() { - clear_protected_tcp_ports_for_test(); - - register_protected_tcp_port(15888); - register_protected_tcp_port(15888); - assert!(is_protected_tcp_port(15888)); - - unregister_protected_tcp_port(15888); - assert!(is_protected_tcp_port(15888)); - - unregister_protected_tcp_port(15888); - assert!(!is_protected_tcp_port(15888)); - } - - #[test] - fn unregistering_unknown_port_is_a_noop() { - clear_protected_tcp_ports_for_test(); - unregister_protected_tcp_port(15888); - assert!(!is_protected_tcp_port(15888)); - } -} diff --git a/easytier/src/rpc_service/proxy.rs b/easytier/src/rpc_service/proxy.rs deleted file mode 100644 index d92fb5d8..00000000 --- a/easytier/src/rpc_service/proxy.rs +++ /dev/null @@ -1,41 +0,0 @@ -use std::sync::Arc; - -use crate::{ - instance_manager::NetworkInstanceManager, - proto::{ - api::instance::{ListTcpProxyEntryRequest, ListTcpProxyEntryResponse, TcpProxyRpc}, - rpc_types::controller::BaseController, - }, -}; - -#[derive(Clone)] -pub struct TcpProxyRpcService { - instance_manager: Arc, - client_type: &'static str, -} - -impl TcpProxyRpcService { - pub fn new(instance_manager: Arc, client_type: &'static str) -> Self { - Self { - instance_manager, - client_type, - } - } -} - -#[async_trait::async_trait] -impl TcpProxyRpc for TcpProxyRpcService { - type Controller = BaseController; - - async fn list_tcp_proxy_entry( - &self, - ctrl: Self::Controller, - req: ListTcpProxyEntryRequest, - ) -> crate::proto::rpc_types::error::Result { - super::get_instance_service(&self.instance_manager, &req.instance)? - .get_proxy_service(self.client_type) - .ok_or_else(|| anyhow::anyhow!("TCP proxy service not found for {}", self.client_type))? - .list_tcp_proxy_entry(ctrl, req) - .await - } -} diff --git a/easytier/src/rpc_service/stats.rs b/easytier/src/rpc_service/stats.rs deleted file mode 100644 index da6e0065..00000000 --- a/easytier/src/rpc_service/stats.rs +++ /dev/null @@ -1,50 +0,0 @@ -use std::sync::Arc; - -use crate::{ - instance_manager::NetworkInstanceManager, - proto::{ - api::instance::{ - GetPrometheusStatsRequest, GetPrometheusStatsResponse, GetStatsRequest, - GetStatsResponse, StatsRpc, - }, - rpc_types::controller::BaseController, - }, -}; - -#[derive(Clone)] -pub struct StatsRpcService { - instance_manager: Arc, -} - -impl StatsRpcService { - pub fn new(instance_manager: Arc) -> Self { - Self { instance_manager } - } -} - -#[async_trait::async_trait] -impl StatsRpc for StatsRpcService { - type Controller = BaseController; - - async fn get_stats( - &self, - ctrl: Self::Controller, - req: GetStatsRequest, - ) -> crate::proto::rpc_types::error::Result { - super::get_instance_service(&self.instance_manager, &req.instance)? - .get_stats_service() - .get_stats(ctrl, req) - .await - } - - async fn get_prometheus_stats( - &self, - ctrl: Self::Controller, - req: GetPrometheusStatsRequest, - ) -> crate::proto::rpc_types::error::Result { - super::get_instance_service(&self.instance_manager, &req.instance)? - .get_stats_service() - .get_prometheus_stats(ctrl, req) - .await - } -} diff --git a/easytier/src/rpc_service/vpn_portal.rs b/easytier/src/rpc_service/vpn_portal.rs deleted file mode 100644 index e09694d6..00000000 --- a/easytier/src/rpc_service/vpn_portal.rs +++ /dev/null @@ -1,36 +0,0 @@ -use std::sync::Arc; - -use crate::{ - instance_manager::NetworkInstanceManager, - proto::{ - api::instance::{GetVpnPortalInfoRequest, GetVpnPortalInfoResponse, VpnPortalRpc}, - rpc_types::controller::BaseController, - }, -}; - -#[derive(Clone)] -pub struct VpnPortalRpcService { - instance_manager: Arc, -} - -impl VpnPortalRpcService { - pub fn new(instance_manager: Arc) -> Self { - Self { instance_manager } - } -} - -#[async_trait::async_trait] -impl VpnPortalRpc for VpnPortalRpcService { - type Controller = BaseController; - - async fn get_vpn_portal_info( - &self, - ctrl: Self::Controller, - req: GetVpnPortalInfoRequest, - ) -> crate::proto::rpc_types::error::Result { - super::get_instance_service(&self.instance_manager, &req.instance)? - .get_vpn_portal_service() - .get_vpn_portal_info(ctrl, req) - .await - } -} diff --git a/easytier/src/tunnel/fake_tcp/LICENSE b/easytier/src/socket/fake_tcp/LICENSE similarity index 100% rename from easytier/src/tunnel/fake_tcp/LICENSE rename to easytier/src/socket/fake_tcp/LICENSE diff --git a/easytier/src/socket/fake_tcp/mod.rs b/easytier/src/socket/fake_tcp/mod.rs new file mode 100644 index 00000000..0d2c9dd3 --- /dev/null +++ b/easytier/src/socket/fake_tcp/mod.rs @@ -0,0 +1,533 @@ +mod netfilter; +mod packet; +mod stack; + +use bytes::BytesMut; +use network_interface::NetworkInterfaceConfig; +use pnet::util::MacAddr; +use std::{ + io, + net::{IpAddr, Ipv4Addr, SocketAddr, UdpSocket}, + pin::Pin, + sync::Arc, + task::{Context as TaskContext, Poll}, +}; +use tokio::{ + io::{AsyncRead, AsyncReadExt, AsyncWrite, ReadBuf}, + net::TcpStream, +}; + +use easytier_core::{ + socket::tcp::VirtualTcpSocket, + tunnel::{IpVersion, TunnelError}, +}; + +use crate::{common::netns::NetNS, tunnel::FromUrl}; + +use self::netfilter::create_tun; + +use futures::Future; +use tokio_util::task::AbortOnDropHandle; + +use dashmap::DashMap; + +struct IpToIfNameCache { + ip_to_ifname: DashMap)>, +} + +impl IpToIfNameCache { + fn new() -> Self { + Self { + ip_to_ifname: DashMap::new(), + } + } + + fn reload_ip_to_ifname(&self) { + self.ip_to_ifname.clear(); + let Ok(interfaces) = network_interface::NetworkInterface::show() else { + tracing::warn!("failed to enumerate interfaces when reloading faketcp ip cache"); + return; + }; + for iface in interfaces { + let mac = iface.mac_addr.as_deref().and_then(|mac| { + mac.parse::().map_err(|e| { + tracing::debug!(iface = %iface.name, mac, ?e, "failed to parse interface mac") + }).ok() + }); + for ip in iface.addr.iter() { + self.ip_to_ifname.insert(ip.ip(), (iface.name.clone(), mac)); + } + } + } + + fn get_ifname(&self, ip: &IpAddr) -> Option<(String, Option)> { + if let Some(ifname) = self.ip_to_ifname.get(ip) { + Some(ifname.clone()) + } else { + self.reload_ip_to_ifname(); + self.ip_to_ifname.get(ip).map(|s| s.clone()) + } + } +} + +fn faketcp_transport_label(driver_type: &str) -> String { + format!("faketcp_{}", driver_type) +} + +async fn create_tun_off_runtime( + interface_name: String, + src_addr: Option, + dst_addr: SocketAddr, + net_ns: NetNS, +) -> Result, TunnelError> { + tokio::task::spawn_blocking(move || { + net_ns.run(|| create_tun(&interface_name, src_addr, dst_addr)) + }) + .await + .map_err(|e| TunnelError::InternalError(format!("faketcp create_tun task failed: {e}")))? + .map_err(Into::into) +} + +pub(crate) struct FakeTcpSocketListener { + addr: url::Url, + os_listener: Option, + // interface_name -> fake tcp stack + stack_map: DashMap>, + // a cache from ip addr to interface name + ip_to_ifname: IpToIfNameCache, +} + +impl std::fmt::Debug for FakeTcpSocketListener { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("FakeTcpSocketListener") + .field("addr", &self.addr) + .field("listening", &self.os_listener.is_some()) + .finish() + } +} + +impl FakeTcpSocketListener { + pub(crate) fn new(addr: url::Url) -> Self { + FakeTcpSocketListener { + addr, + os_listener: None, + stack_map: DashMap::new(), + ip_to_ifname: IpToIfNameCache::new(), + } + } + + async fn do_accept(&mut self) -> Result { + loop { + match self.os_listener.as_mut().unwrap().accept().await { + Ok((s, remote_addr)) => { + let Ok(local_addr) = s.local_addr() else { + tracing::warn!("accept fail with local_addr error"); + continue; + }; + let Some((interface_name, mac)) = + self.ip_to_ifname.get_ifname(&local_addr.ip()) + else { + tracing::warn!("accept fail with interface_name error"); + continue; + }; + return Ok(AcceptResult { + socket: s, + local_addr, + remote_addr, + interface_name, + mac, + }); + } + Err(e) => { + use std::io::ErrorKind::*; + if matches!( + e.kind(), + NotConnected | ConnectionAborted | ConnectionRefused | ConnectionReset + ) { + tracing::warn!(?e, "accept fail with retryable error: {:?}", e); + continue; + } + tracing::warn!(?e, "accept fail"); + return Err(e.into()); + } + } + } + } + + async fn get_stack( + &self, + accept_result: &AcceptResult, + ) -> Result, TunnelError> { + let local_socket_addr = accept_result.local_addr; + + let interface_name = &accept_result.interface_name; + + if let Some(entry) = self.stack_map.get(interface_name) { + let stack = entry.clone(); + drop(entry); + + if !stack.is_closed() { + return Ok(stack); + } + + tracing::warn!( + interface_name, + "fake_tcp stack reader_task finished, recreating stack" + ); + self.stack_map.remove(interface_name); + } + + let tun = create_tun_off_runtime( + interface_name.to_string(), + None, + local_socket_addr, + NetNS::new(None), + ) + .await?; + tracing::info!( + ?local_socket_addr, + "create new stack with interface_name: {:?}", + interface_name + ); + let stack = Arc::new(stack::Stack::new(tun, accept_result.mac)); + self.stack_map + .insert(interface_name.to_string(), stack.clone()); + + Ok(stack) + } +} + +fn build_os_socket_reader_task(mut socket: TcpStream) -> AbortOnDropHandle<()> { + AbortOnDropHandle::new(tokio::spawn(async move { + // read the os socket until it's closed + let mut buf = [0u8; 1024]; + while let Ok(size) = socket.read(&mut buf).await { + tracing::trace!("read {} bytes from os socket", size); + if size == 0 { + break; + } + } + tracing::info!("FakeTcpSocketListener os socket closed"); + })) +} + +type FakeTcpReadFuture = Pin> + Send + Sync + 'static>>; + +enum FakeTcpReadState { + Buffered(BytesMut), + Receiving(FakeTcpReadFuture), + Closed, +} + +pub(crate) struct FakeTcpSocket { + socket: Arc, + read_state: FakeTcpReadState, + transport_label: String, + _lifetime_guard: Box, +} + +impl FakeTcpSocket { + fn new(socket: stack::Socket, transport_label: String, lifetime_guard: T) -> Self + where + T: Send + Sync + 'static, + { + Self { + socket: Arc::new(socket), + read_state: FakeTcpReadState::Buffered(BytesMut::new()), + transport_label, + _lifetime_guard: Box::new(lifetime_guard), + } + } +} + +impl AsyncRead for FakeTcpSocket { + fn poll_read( + self: Pin<&mut Self>, + context: &mut TaskContext<'_>, + output: &mut ReadBuf<'_>, + ) -> Poll> { + let this = self.get_mut(); + loop { + let state = std::mem::replace(&mut this.read_state, FakeTcpReadState::Closed); + match state { + FakeTcpReadState::Buffered(mut buffer) if !buffer.is_empty() => { + let length = buffer.len().min(output.remaining()); + output.put_slice(&buffer.split_to(length)); + this.read_state = FakeTcpReadState::Buffered(buffer); + return Poll::Ready(Ok(())); + } + FakeTcpReadState::Buffered(_) => { + let socket = this.socket.clone(); + this.read_state = FakeTcpReadState::Receiving(Box::pin(async move { + let mut buffer = BytesMut::new(); + socket.recv(&mut buffer).await.map(|_| buffer) + })); + } + FakeTcpReadState::Receiving(mut receive) => match receive.as_mut().poll(context) { + Poll::Ready(Some(buffer)) => { + this.read_state = FakeTcpReadState::Buffered(buffer); + } + Poll::Ready(None) => { + this.read_state = FakeTcpReadState::Closed; + return Poll::Ready(Ok(())); + } + Poll::Pending => { + this.read_state = FakeTcpReadState::Receiving(receive); + return Poll::Pending; + } + }, + FakeTcpReadState::Closed => return Poll::Ready(Ok(())), + } + } + } +} + +impl AsyncWrite for FakeTcpSocket { + fn poll_write( + self: Pin<&mut Self>, + _context: &mut TaskContext<'_>, + buffer: &[u8], + ) -> Poll> { + if self.socket.try_send(buffer).is_none() { + // Preserve FakeTCP's existing lossy send behavior. A temporary + // driver lock conflict is indistinguishable from a closed stack + // here, and the former must not tear down the peer connection. + tracing::trace!( + len = buffer.len(), + "FakeTCP socket dropped an outgoing frame" + ); + } + Poll::Ready(Ok(buffer.len())) + } + + fn poll_flush(self: Pin<&mut Self>, _context: &mut TaskContext<'_>) -> Poll> { + Poll::Ready(Ok(())) + } + + fn poll_shutdown(self: Pin<&mut Self>, _context: &mut TaskContext<'_>) -> Poll> { + self.socket.close(); + Poll::Ready(Ok(())) + } +} + +impl VirtualTcpSocket for FakeTcpSocket { + fn local_addr(&self) -> io::Result { + Ok(self.socket.local_addr()) + } + + fn peer_addr(&self) -> io::Result { + Ok(self.socket.remote_addr()) + } + + fn transport_label(&self) -> Option<&str> { + Some(&self.transport_label) + } +} + +#[derive(Debug)] +struct AcceptResult { + socket: TcpStream, + local_addr: SocketAddr, + remote_addr: SocketAddr, + interface_name: String, + mac: Option, +} + +impl FakeTcpSocketListener { + pub(crate) async fn accept_socket(&mut self) -> Result { + tracing::debug!("FakeTcpSocketListener waiting for accept"); + let (res, stack, socket) = loop { + let res = self.do_accept().await?; + let stack = self.get_stack(&res).await?; + let socket = stack.try_alloc_established_socket( + res.local_addr, + res.remote_addr, + stack::State::Established, + ); + let Some(socket) = socket else { + tracing::warn!( + interface_name = res.interface_name, + "fake_tcp stack closed while accepting connection, dropping accepted socket" + ); + self.stack_map.remove(&res.interface_name); + continue; + }; + break (res, stack, socket); + }; + + tracing::info!( + ?res, + remote = socket.remote_addr().to_string(), + "FakeTcpSocketListener accepted connection" + ); + + let transport_label = faketcp_transport_label(stack.driver_type()); + Ok(FakeTcpSocket::new( + socket, + transport_label, + (build_os_socket_reader_task(res.socket), stack), + )) + } + + async fn listen_socket(&mut self) -> Result<(), TunnelError> { + let port = self.addr.port().unwrap_or(0); + let bind_addr = SocketAddr::from_url(self.addr.clone(), IpVersion::Both).await?; + let os_listener = tokio::net::TcpListener::bind(bind_addr).await?; + tracing::info!(port, "FakeTcpSocketListener listening"); + self.os_listener = Some(os_listener); + Ok(()) + } +} + +#[async_trait::async_trait] +impl easytier_core::socket::SocketListener for FakeTcpSocketListener { + type Accepted = FakeTcpSocket; + + async fn listen(&mut self) -> anyhow::Result<()> { + Ok(self.listen_socket().await?) + } + + async fn accept(&mut self) -> anyhow::Result { + Ok(self.accept_socket().await?) + } + + fn local_url(&self) -> url::Url { + self.addr.clone() + } +} + +fn get_local_ip_for_destination(destination: IpAddr) -> Option { + // 使用一个不可路由的、私有的、或回环地址创建一个临时的 socket,让内核自动选择源接口。 + // 对于 IPv4,使用 0.0.0.0; 对于 IPv6,使用 :: + let bind_addr = if destination.is_ipv4() { + IpAddr::V4(Ipv4Addr::new(0, 0, 0, 0)) + } else { + IpAddr::V6(std::net::Ipv6Addr::new(0, 0, 0, 0, 0, 0, 0, 0)) + }; + + // 绑定到一个临时端口 (0) + let socket = UdpSocket::bind((bind_addr, 0)).ok()?; + + // 尝试连接到目标地址。这不会真正发送数据包,只是让内核确定路由。 + socket.connect((destination, 80)).ok()?; // 使用一个常见的端口,例如 80 + + // 获取 socket 的本地地址信息 + socket.local_addr().map(|addr| addr.ip()).ok() +} + +async fn connect_socket_with_cache( + remote_addr: SocketAddr, + socket_mark: Option, + ip_to_if_name: &IpToIfNameCache, + net_ns: NetNS, +) -> Result { + let (local_addr, interface_name, mac, os_socket) = net_ns.run(|| { + let local_ip = get_local_ip_for_destination(remote_addr.ip()) + .ok_or(TunnelError::InternalError("Failed to get local ip".into()))?; + + let os_socket = tokio::net::TcpSocket::new_v4()?; + // SO_MARK applies only to the kernel-visible "decoy" socket below. + // The actual FakeTCP payload travels via crafted segments written + // straight to the TUN device, which the kernel doesn't tag with + // SO_MARK. Operators relying on fwmark for FakeTCP must mark the + // TUN device's traffic with a separate nftables/iptables rule. + crate::tunnel::common::apply_socket_mark(&socket2::SockRef::from(&os_socket), socket_mark)?; + os_socket.bind("0.0.0.0:0".parse().unwrap())?; + let local_addr = SocketAddr::new(local_ip, os_socket.local_addr()?.port()); + + let (interface_name, mac) = + ip_to_if_name + .get_ifname(&local_ip) + .ok_or(TunnelError::InternalError( + "Failed to get interface name".into(), + ))?; + Ok::<_, TunnelError>((local_addr, interface_name, mac, os_socket)) + })?; + + let tun = create_tun_off_runtime(interface_name, Some(remote_addr), local_addr, net_ns).await?; + let stack = stack::Stack::new(tun, mac); + let transport_label = faketcp_transport_label(stack.driver_type()); + + let socket = stack + .try_alloc_established_socket(local_addr, remote_addr, stack::State::SynSent) + .ok_or(TunnelError::InternalError( + "FakeTCP stack closed while allocating socket".into(), + ))?; + + let os_stream = os_socket.connect(remote_addr).await?; + + tracing::info!(?remote_addr, "FakeTCP socket connecting"); + + let mut buf = BytesMut::new(); + socket + .recv(&mut buf) + .await + .ok_or(TunnelError::InternalError( + "Failed to recv bytes to establish connection".into(), + ))?; + + tracing::info!(local_addr = ?socket.local_addr(), "FakeTCP socket connected"); + + Ok(FakeTcpSocket::new( + socket, + transport_label, + (build_os_socket_reader_task(os_stream), stack), + )) +} + +pub(crate) async fn connect_socket( + remote_addr: SocketAddr, + socket_mark: Option, + net_ns: NetNS, +) -> Result { + connect_socket_with_cache(remote_addr, socket_mark, &IpToIfNameCache::new(), net_ns).await +} + +#[cfg(test)] +mod tests { + use easytier_core::socket::SocketListener; + use tokio::io::{AsyncReadExt, AsyncWriteExt}; + + use super::*; + + #[tokio::test] + async fn faketcp_socket_pingpong() { + tokio::time::timeout(std::time::Duration::from_secs(5), async { + #[cfg(target_family = "unix")] + { + if unsafe { nix::libc::geteuid() } != 0 { + return; + } + } + + let mut listener = + FakeTcpSocketListener::new("faketcp://0.0.0.0:31011".parse().unwrap()); + listener.listen().await.unwrap(); + let (server_ready_tx, server_ready_rx) = tokio::sync::oneshot::channel(); + + let server = tokio::spawn(async move { + let mut socket = listener.accept().await.unwrap(); + server_ready_tx.send(()).unwrap(); + let mut request = [0; 4]; + socket.read_exact(&mut request).await.unwrap(); + assert_eq!(&request, b"ping"); + socket.write_all(b"pong").await.unwrap(); + }); + + let mut socket = + connect_socket("127.0.0.1:31011".parse().unwrap(), None, NetNS::new(None)) + .await + .unwrap(); + server_ready_rx.await.unwrap(); + socket.write_all(b"ping").await.unwrap(); + let mut response = [0; 4]; + socket.read_exact(&mut response).await.unwrap(); + assert_eq!(&response, b"pong"); + + server.await.unwrap(); + }) + .await + .expect("FakeTCP socket ping-pong timed out"); + } +} diff --git a/easytier/src/tunnel/fake_tcp/netfilter/linux_bpf.rs b/easytier/src/socket/fake_tcp/netfilter/linux_bpf.rs similarity index 99% rename from easytier/src/tunnel/fake_tcp/netfilter/linux_bpf.rs rename to easytier/src/socket/fake_tcp/netfilter/linux_bpf.rs index 0b138f91..f9922df3 100644 --- a/easytier/src/tunnel/fake_tcp/netfilter/linux_bpf.rs +++ b/easytier/src/socket/fake_tcp/netfilter/linux_bpf.rs @@ -12,7 +12,7 @@ use std::sync::atomic::{AtomicBool, Ordering as AtomicOrdering}; use std::time::{Duration, Instant}; use tokio::sync::Mutex; -use crate::tunnel::fake_tcp::stack; +use crate::socket::fake_tcp::stack; const ETH_HDR_LEN: usize = 14; const ETH_TYPE_OFFSET: u32 = 12; @@ -630,8 +630,8 @@ impl stack::Tun for LinuxBpfTun { mod tests { use super::*; - use crate::tunnel::fake_tcp::packet::build_tcp_packet; - use crate::tunnel::fake_tcp::stack::Tun; + use crate::socket::fake_tcp::packet::build_tcp_packet; + use crate::socket::fake_tcp::stack::Tun; use pnet::datalink; use pnet::packet::tcp::TcpFlags; use pnet::util::MacAddr; diff --git a/easytier/src/tunnel/fake_tcp/netfilter/macos_bpf.rs b/easytier/src/socket/fake_tcp/netfilter/macos_bpf.rs similarity index 99% rename from easytier/src/tunnel/fake_tcp/netfilter/macos_bpf.rs rename to easytier/src/socket/fake_tcp/netfilter/macos_bpf.rs index afc11927..987717d8 100644 --- a/easytier/src/tunnel/fake_tcp/netfilter/macos_bpf.rs +++ b/easytier/src/socket/fake_tcp/netfilter/macos_bpf.rs @@ -12,7 +12,7 @@ use std::sync::atomic::{AtomicBool, Ordering as AtomicOrdering}; use tokio::sync::Mutex; use tracing::{debug, info, warn}; -use crate::tunnel::fake_tcp::stack; +use crate::socket::fake_tcp::stack; const ETH_HDR_LEN: usize = 14; const ETH_TYPE_OFFSET: u32 = 12; diff --git a/easytier/src/tunnel/fake_tcp/netfilter/mod.rs b/easytier/src/socket/fake_tcp/netfilter/mod.rs similarity index 100% rename from easytier/src/tunnel/fake_tcp/netfilter/mod.rs rename to easytier/src/socket/fake_tcp/netfilter/mod.rs diff --git a/easytier/src/tunnel/fake_tcp/netfilter/pnet.rs b/easytier/src/socket/fake_tcp/netfilter/pnet.rs similarity index 86% rename from easytier/src/tunnel/fake_tcp/netfilter/pnet.rs rename to easytier/src/socket/fake_tcp/netfilter/pnet.rs index 4f6e5e66..ceea5751 100644 --- a/easytier/src/tunnel/fake_tcp/netfilter/pnet.rs +++ b/easytier/src/socket/fake_tcp/netfilter/pnet.rs @@ -14,9 +14,11 @@ use pnet::{ datalink::{self, DataLinkSender, NetworkInterface}, packet::{ethernet::EtherTypes, ip::IpNextHeaderProtocols, ipv6::Ipv6Packet}, }; +#[cfg(target_os = "linux")] +use std::os::unix::fs::MetadataExt; use tokio::sync::Mutex; -use crate::tunnel::fake_tcp::stack; +use crate::socket::fake_tcp::stack; type PacketFilter = Box bool + Send + Sync>; @@ -74,7 +76,7 @@ fn filter_tcp_packet( tracing::trace!( ?tcp, - "FakeTcpTunnelListener packet matched filter, dispatching, src_addr: {:?}, dst_addr: {:?}, packet_src_ip: {:?}, packet_dst_ip: {:?}, packet_src_port: {:?}, packet_dst_port: {:?}", + "FakeTcpSocketListener packet matched filter, dispatching, src_addr: {:?}, dst_addr: {:?}, packet_src_ip: {:?}, packet_dst_ip: {:?}, packet_src_port: {:?}, packet_dst_port: {:?}", src_addr, dst_addr, ipv4.get_source(), @@ -120,7 +122,7 @@ fn filter_tcp_packet( tracing::trace!( ?tcp, - "FakeTcpTunnelListener packet matched filter, dispatching" + "FakeTcpSocketListener packet matched filter, dispatching" ); } _ => return false, @@ -203,13 +205,46 @@ impl InterfaceWorker { } } -static INTERFACE_MANAGERS: Lazy>> = Lazy::new(DashMap::new); +#[derive(Clone, Debug, PartialEq, Eq, Hash)] +enum NetworkNamespaceId { + #[cfg(target_os = "linux")] + LinuxInode(u64), + #[cfg(not(target_os = "linux"))] + Unsupported, +} + +type InterfaceWorkerKey = (NetworkNamespaceId, String); + +static INTERFACE_MANAGERS: Lazy>> = + Lazy::new(DashMap::new); + +fn current_network_namespace_id() -> Option { + #[cfg(target_os = "linux")] + { + match std::fs::metadata("/proc/thread-self/ns/net") { + Ok(metadata) => Some(NetworkNamespaceId::LinuxInode(metadata.ino())), + Err(error) => { + tracing::warn!( + ?error, + "failed to identify current network namespace; disabling FakeTCP pnet worker sharing" + ); + None + } + } + } + + #[cfg(not(target_os = "linux"))] + Some(NetworkNamespaceId::Unsupported) +} fn get_or_create_worker(interface_name: &str) -> io::Result> { + let key = current_network_namespace_id() + .map(|namespace_id| (namespace_id, interface_name.to_owned())); // Check if we have an active worker - if let Some(worker) = INTERFACE_MANAGERS - .get(interface_name) - .and_then(|w| w.upgrade()) + if let Some(worker) = key + .as_ref() + .and_then(|key| INTERFACE_MANAGERS.get(key)) + .and_then(|worker| worker.upgrade()) { return Ok(worker); } @@ -234,7 +269,9 @@ fn get_or_create_worker(interface_name: &str) -> io::Result })?; let worker = InterfaceWorker::new(interface)?; - INTERFACE_MANAGERS.insert(interface_name.to_string(), Arc::downgrade(&worker)); + if let Some(key) = key { + INTERFACE_MANAGERS.insert(key, Arc::downgrade(&worker)); + } Ok(worker) } diff --git a/easytier/src/tunnel/fake_tcp/netfilter/windivert.rs b/easytier/src/socket/fake_tcp/netfilter/windivert.rs similarity index 99% rename from easytier/src/tunnel/fake_tcp/netfilter/windivert.rs rename to easytier/src/socket/fake_tcp/netfilter/windivert.rs index 289d4a8e..58ffc14c 100644 --- a/easytier/src/tunnel/fake_tcp/netfilter/windivert.rs +++ b/easytier/src/socket/fake_tcp/netfilter/windivert.rs @@ -11,7 +11,7 @@ use windivert::packet::WinDivertPacket; use windivert::prelude::{WinDivertFlags, WinDivertShutdownMode}; use windivert::{WinDivert, layer}; -use crate::tunnel::fake_tcp::stack; +use crate::socket::fake_tcp::stack; struct WinDivertReader { inner: UnsafeCell>, diff --git a/easytier/src/tunnel/fake_tcp/packet.rs b/easytier/src/socket/fake_tcp/packet.rs similarity index 99% rename from easytier/src/tunnel/fake_tcp/packet.rs rename to easytier/src/socket/fake_tcp/packet.rs index 8f494e10..45cfbc5c 100644 --- a/easytier/src/tunnel/fake_tcp/packet.rs +++ b/easytier/src/socket/fake_tcp/packet.rs @@ -8,8 +8,6 @@ use std::net::{IpAddr, SocketAddr}; const IPV4_HEADER_LEN: usize = 20; const IPV6_HEADER_LEN: usize = 40; const TCP_HEADER_LEN: usize = 20; -pub const MAX_PACKET_LEN: usize = 1500; - #[derive(Debug)] pub enum IPPacket<'p> { V4(ipv4::Ipv4Packet<'p>), diff --git a/easytier/src/tunnel/fake_tcp/stack.rs b/easytier/src/socket/fake_tcp/stack.rs similarity index 94% rename from easytier/src/tunnel/fake_tcp/stack.rs rename to easytier/src/socket/fake_tcp/stack.rs index a7f1b779..d3934456 100644 --- a/easytier/src/tunnel/fake_tcp/stack.rs +++ b/easytier/src/socket/fake_tcp/stack.rs @@ -44,9 +44,11 @@ use crossbeam::atomic::AtomicCell; use pnet::packet::tcp::TcpOptionNumbers; use pnet::packet::{Packet, tcp}; use pnet::util::MacAddr; -use std::collections::{HashMap, HashSet}; +use std::collections::HashMap; use std::fmt; -use std::net::{Ipv4Addr, Ipv6Addr, SocketAddr}; +#[cfg(test)] +use std::net::Ipv4Addr; +use std::net::SocketAddr; use std::sync::{ Arc, RwLock, atomic::{AtomicU32, Ordering}, @@ -57,9 +59,7 @@ use tokio_util::task::AbortOnDropHandle; use tracing::{error, info, trace, warn}; const TIMEOUT: time::Duration = time::Duration::from_secs(1); -const RETRIES: usize = 6; const MPMC_BUFFER_LEN: usize = 512; -const MAX_UNACKED_LEN: u32 = 128 * 1024 * 1024; // 128MB #[async_trait::async_trait] pub trait Tun: Send + Sync + 'static { @@ -91,7 +91,6 @@ struct StackState { struct Shared { state: RwLock, - listening: RwLock>, tun: Arc, tuples_purge: broadcast::Sender, } @@ -112,8 +111,6 @@ impl Shared { pub struct Stack { shared: Arc, - local_ip: Ipv4Addr, - local_ip6: Option, local_mac: MacAddr, reader_task: AbortOnDropHandle<()>, } @@ -122,7 +119,6 @@ pub struct Stack { pub enum State { Idle, SynSent, - SynReceived, Established, } @@ -422,17 +418,11 @@ impl Stack { /// When more than one [`Tun`](tokio_tun::Tun) object is passed in, same amount /// of reader will be spawned later. This allows user to utilize the performance /// benefit of Multiqueue Tun support on machines with SMP. - pub fn new( - tun: Arc, - local_ip: Ipv4Addr, - local_ip6: Option, - local_mac: Option, - ) -> Stack { + pub fn new(tun: Arc, local_mac: Option) -> Stack { let (tuples_purge_tx, _tuples_purge_rx) = broadcast::channel(16); let shared = Arc::new(Shared { state: RwLock::new(StackState::default()), tun: tun.clone(), - listening: RwLock::new(HashSet::new()), tuples_purge: tuples_purge_tx.clone(), }); @@ -444,8 +434,6 @@ impl Stack { Stack { shared, - local_ip, - local_ip6, local_mac: local_mac.unwrap_or(MacAddr::zero()), reader_task: AbortOnDropHandle::new(t), } @@ -460,11 +448,6 @@ impl Stack { self.shared.is_closed() || self.reader_task.is_finished() } - /// Listens for incoming connections on the given `port`. - pub fn listen(&mut self, port: u16) { - assert!(self.shared.listening.write().unwrap().insert(port)); - } - pub fn try_alloc_established_socket( &self, local_addr: SocketAddr, @@ -565,16 +548,7 @@ impl Stack { } } - if tcp_packet.get_flags() == tcp::TcpFlags::SYN - && shared - .listening - .read() - .unwrap() - .contains(&tcp_packet.get_destination()) - { - trace!(?tcp_packet, "Received SYN packet for port {}, ignoring", tcp_packet.get_destination()); - continue; - } else if (tcp_packet.get_flags() & tcp::TcpFlags::RST) != 0 { + if (tcp_packet.get_flags() & tcp::TcpFlags::RST) != 0 { info!("Unknown RST TCP packet from {}, ignoring", remote_addr); continue; } else { @@ -660,7 +634,7 @@ mod tests { #[tokio::test] async fn reader_task_closes_sockets_on_tun_recv_error() { let tun = Arc::new(FailingTun::default()); - let mut stack = Stack::new(tun.clone(), Ipv4Addr::LOCALHOST, None, None); + let mut stack = Stack::new(tun.clone(), None); let socket = stack .try_alloc_established_socket( SocketAddr::new(Ipv4Addr::LOCALHOST.into(), 10_000), diff --git a/easytier/src/socket/mod.rs b/easytier/src/socket/mod.rs new file mode 100644 index 00000000..5acb7604 --- /dev/null +++ b/easytier/src/socket/mod.rs @@ -0,0 +1,5 @@ +#[cfg(feature = "faketcp")] +pub(crate) mod fake_tcp; +pub(crate) mod tcp; +pub(crate) mod udp; +pub(crate) mod udp_src; diff --git a/easytier/src/socket/tcp.rs b/easytier/src/socket/tcp.rs new file mode 100644 index 00000000..fa27d230 --- /dev/null +++ b/easytier/src/socket/tcp.rs @@ -0,0 +1,374 @@ +use std::{ + io, + net::{IpAddr, SocketAddr}, + pin::Pin, + task::{Context, Poll}, + time::Duration, +}; + +use easytier_core::{ + socket::tcp::{ + TcpBindOptions, TcpConnectOptions, TcpListenOptions, TcpListenPurpose, TcpSocketPurpose, + VirtualTcpListener, VirtualTcpSocket, + }, + tunnel::TunnelError, +}; +use socket2::{SockRef, TcpKeepalive}; +#[cfg(unix)] +use tokio::net::UnixStream; +use tokio::{ + io::{AsyncRead, AsyncWrite, ReadBuf}, + net::{TcpListener, TcpSocket, TcpStream}, +}; + +use crate::{ + common::netns::NetNS, + tunnel::common::{BindDev, apply_socket_mark, bind}, +}; + +enum RuntimeTcpSocketInner { + Tcp(TcpStream), + #[cfg(unix)] + Unix(UnixStream), + #[cfg(feature = "faketcp")] + FakeTcp(crate::socket::fake_tcp::FakeTcpSocket), +} + +pub struct RuntimeTcpSocket { + inner: RuntimeTcpSocketInner, +} + +impl RuntimeTcpSocket { + pub(crate) fn new(stream: TcpStream) -> Self { + if let Err(error) = stream.set_nodelay(true) { + tracing::warn!(?error, "set_nodelay failed for tcp stream"); + } + Self { + inner: RuntimeTcpSocketInner::Tcp(stream), + } + } + + #[cfg(unix)] + pub(crate) fn from_unix(stream: UnixStream) -> Self { + Self { + inner: RuntimeTcpSocketInner::Unix(stream), + } + } + + #[cfg(feature = "faketcp")] + pub(crate) fn from_fake_tcp(socket: crate::socket::fake_tcp::FakeTcpSocket) -> Self { + Self { + inner: RuntimeTcpSocketInner::FakeTcp(socket), + } + } +} + +#[cfg(unix)] +pub(crate) fn url_from_unix_socket_addr(addr: tokio::net::unix::SocketAddr) -> Option { + addr.as_pathname() + .and_then(|path| path.to_str()) + .and_then(|path| format!("unix://{path}").parse().ok()) +} + +impl AsyncRead for RuntimeTcpSocket { + fn poll_read( + mut self: Pin<&mut Self>, + cx: &mut Context<'_>, + buf: &mut ReadBuf<'_>, + ) -> Poll> { + match &mut self.inner { + RuntimeTcpSocketInner::Tcp(stream) => Pin::new(stream).poll_read(cx, buf), + #[cfg(unix)] + RuntimeTcpSocketInner::Unix(stream) => Pin::new(stream).poll_read(cx, buf), + #[cfg(feature = "faketcp")] + RuntimeTcpSocketInner::FakeTcp(socket) => Pin::new(socket).poll_read(cx, buf), + } + } +} + +impl AsyncWrite for RuntimeTcpSocket { + fn poll_write( + mut self: Pin<&mut Self>, + cx: &mut Context<'_>, + buf: &[u8], + ) -> Poll> { + match &mut self.inner { + RuntimeTcpSocketInner::Tcp(stream) => Pin::new(stream).poll_write(cx, buf), + #[cfg(unix)] + RuntimeTcpSocketInner::Unix(stream) => Pin::new(stream).poll_write(cx, buf), + #[cfg(feature = "faketcp")] + RuntimeTcpSocketInner::FakeTcp(socket) => Pin::new(socket).poll_write(cx, buf), + } + } + + fn poll_flush(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + match &mut self.inner { + RuntimeTcpSocketInner::Tcp(stream) => Pin::new(stream).poll_flush(cx), + #[cfg(unix)] + RuntimeTcpSocketInner::Unix(stream) => Pin::new(stream).poll_flush(cx), + #[cfg(feature = "faketcp")] + RuntimeTcpSocketInner::FakeTcp(socket) => Pin::new(socket).poll_flush(cx), + } + } + + fn poll_shutdown(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + match &mut self.inner { + RuntimeTcpSocketInner::Tcp(stream) => Pin::new(stream).poll_shutdown(cx), + #[cfg(unix)] + RuntimeTcpSocketInner::Unix(stream) => Pin::new(stream).poll_shutdown(cx), + #[cfg(feature = "faketcp")] + RuntimeTcpSocketInner::FakeTcp(socket) => Pin::new(socket).poll_shutdown(cx), + } + } +} + +impl VirtualTcpSocket for RuntimeTcpSocket { + fn local_addr(&self) -> io::Result { + match &self.inner { + RuntimeTcpSocketInner::Tcp(stream) => stream.local_addr(), + #[cfg(unix)] + RuntimeTcpSocketInner::Unix(_) => Err(io::Error::new( + io::ErrorKind::Unsupported, + "Unix stream has no IP local address", + )), + #[cfg(feature = "faketcp")] + RuntimeTcpSocketInner::FakeTcp(socket) => socket.local_addr(), + } + } + + fn peer_addr(&self) -> io::Result { + match &self.inner { + RuntimeTcpSocketInner::Tcp(stream) => stream.peer_addr(), + #[cfg(unix)] + RuntimeTcpSocketInner::Unix(_) => Err(io::Error::new( + io::ErrorKind::Unsupported, + "Unix stream has no IP peer address", + )), + #[cfg(feature = "faketcp")] + RuntimeTcpSocketInner::FakeTcp(socket) => socket.peer_addr(), + } + } + + fn transport_label(&self) -> Option<&str> { + match &self.inner { + #[cfg(feature = "faketcp")] + RuntimeTcpSocketInner::FakeTcp(socket) => socket.transport_label(), + RuntimeTcpSocketInner::Tcp(_) => None, + #[cfg(unix)] + RuntimeTcpSocketInner::Unix(_) => None, + } + } +} + +#[derive(Debug)] +pub struct RuntimeTcpListener { + listener: TcpListener, + purpose: TcpListenPurpose, +} + +impl RuntimeTcpListener { + pub(crate) fn new(listener: TcpListener, purpose: TcpListenPurpose) -> Self { + Self { listener, purpose } + } +} + +#[async_trait::async_trait] +impl VirtualTcpListener for RuntimeTcpListener { + type Socket = RuntimeTcpSocket; + + fn local_addr(&self) -> io::Result { + self.listener.local_addr() + } + + async fn accept(&self) -> io::Result<(Self::Socket, SocketAddr)> { + let (stream, addr) = self.listener.accept().await?; + if self.purpose == TcpListenPurpose::ProxyNat { + prepare_proxy_tcp_socket(&stream)?; + } + Ok((RuntimeTcpSocket::new(stream), addr)) + } +} + +fn unspecified_bind_addr(remote_addr: SocketAddr) -> SocketAddr { + match remote_addr { + SocketAddr::V4(_) => SocketAddr::new(IpAddr::from([0, 0, 0, 0]), 0), + SocketAddr::V6(_) => SocketAddr::new(IpAddr::from([0, 0, 0, 0, 0, 0, 0, 0]), 0), + } +} + +fn bind_dev_from_options(options: &TcpBindOptions, local_addr_was_defaulted: bool) -> BindDev { + options + .bind_device + .clone() + .map(BindDev::from) + .unwrap_or_else(|| { + if local_addr_was_defaulted { + BindDev::Disabled + } else { + BindDev::Auto + } + }) +} + +fn bind_tcp_socket( + remote_addr: SocketAddr, + bind_options: TcpBindOptions, +) -> Result { + let (bind_addr, local_addr_was_defaulted) = match bind_options.local_addr { + Some(addr) => (addr, false), + None => (unspecified_bind_addr(remote_addr), true), + }; + let bind_dev = bind_dev_from_options(&bind_options, local_addr_was_defaulted); + + bind::() + .addr(bind_addr) + .dev(bind_dev) + .net_ns(NetNS::from_socket_context(&bind_options.context)) + .only_v6(bind_options.only_v6) + .reuse_addr(native_reuse_addr(&bind_options)) + .reuse_port(bind_options.reuse_port) + .maybe_socket_mark(bind_options.context.socket_mark) + .call() +} + +fn create_tcp_socket( + remote_addr: SocketAddr, + bind_options: &TcpBindOptions, +) -> Result { + // A network namespace is a thread property, but the socket retains its + // namespace after creation. Never keep the guard across connect().await. + NetNS::from_socket_context(&bind_options.context).run(|| { + let socket = if remote_addr.is_ipv4() { + TcpSocket::new_v4()? + } else { + TcpSocket::new_v6()? + }; + apply_socket_mark( + &socket2::SockRef::from(&socket), + bind_options.context.socket_mark, + )?; + Ok(socket) + }) +} + +fn must_bind_before_connect(bind_options: &TcpBindOptions) -> bool { + bind_options.local_addr.is_some() + || bind_options.bind_device.is_some() + || bind_options.reuse_port + || bind_options.only_v6 + || bind_options + .reuse_addr + .is_some_and(|reuse_addr| reuse_addr != native_reuse_addr_default()) +} + +fn native_reuse_addr_default() -> bool { + !cfg!(target_os = "windows") +} + +fn native_reuse_addr(bind_options: &TcpBindOptions) -> bool { + bind_options + .reuse_addr + .unwrap_or_else(native_reuse_addr_default) +} + +pub(crate) fn bind_tcp_listener( + options: TcpListenOptions, +) -> Result { + let purpose = options.purpose; + let bind_options = options.bind; + let net_ns = NetNS::from_socket_context(&bind_options.context); + let addr = bind_options.local_addr.ok_or_else(|| { + TunnelError::InvalidAddr("tcp listener requires a local bind address".to_owned()) + })?; + let bind_dev = if bind_options.bind_device.is_none() && purpose == TcpListenPurpose::PortLease { + BindDev::Disabled + } else { + bind_dev_from_options(&bind_options, false) + }; + let listener = bind::() + .addr(addr) + .dev(bind_dev) + .maybe_net_ns(Some(net_ns)) + .only_v6(bind_options.only_v6) + .reuse_addr(native_reuse_addr(&bind_options)) + .reuse_port(bind_options.reuse_port) + .maybe_socket_mark(bind_options.context.socket_mark) + .call()?; + Ok(RuntimeTcpListener::new(listener, purpose)) +} + +pub(crate) async fn connect_tcp( + options: TcpConnectOptions, +) -> Result { + let remote_addr = options.remote_addr; + let purpose = options.purpose; + let bind_options = options.bind; + + if !must_bind_before_connect(&bind_options) { + let socket = create_tcp_socket(remote_addr, &bind_options)?; + let stream = socket.connect(remote_addr).await?; + prepare_connected_tcp_socket(&stream, purpose)?; + return Ok(RuntimeTcpSocket::new(stream)); + } + + let socket = bind_tcp_socket(remote_addr, bind_options)?; + let stream = socket.connect(remote_addr).await?; + prepare_connected_tcp_socket(&stream, purpose)?; + Ok(RuntimeTcpSocket::new(stream)) +} + +fn prepare_connected_tcp_socket(stream: &TcpStream, purpose: TcpSocketPurpose) -> io::Result<()> { + match purpose { + TcpSocketPurpose::ProxyNat => prepare_proxy_tcp_socket(stream), + TcpSocketPurpose::StunProbe => SockRef::from(stream).set_linger(Some(Duration::ZERO)), + _ => Ok(()), + } +} + +pub(crate) fn prepare_proxy_tcp_socket(stream: &TcpStream) -> io::Result<()> { + const TCP_KEEPALIVE_TIME: std::time::Duration = std::time::Duration::from_secs(5); + const TCP_KEEPALIVE_INTERVAL: std::time::Duration = std::time::Duration::from_secs(2); + #[cfg(not(target_os = "windows"))] + const TCP_KEEPALIVE_RETRIES: u32 = 2; + + let keepalive = TcpKeepalive::new() + .with_time(TCP_KEEPALIVE_TIME) + .with_interval(TCP_KEEPALIVE_INTERVAL); + + #[cfg(not(target_os = "windows"))] + let keepalive = keepalive.with_retries(TCP_KEEPALIVE_RETRIES); + + let socket = SockRef::from(stream); + socket.set_tcp_keepalive(&keepalive)?; + if let Err(error) = socket.set_nodelay(true) { + tracing::warn!(?error, "set_nodelay failed, ignore it"); + } + + Ok(()) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn tcp_connect_binds_when_socket_option_requires_pre_connect_setup() { + assert!(must_bind_before_connect( + &TcpBindOptions::default().with_only_v6(true) + )); + assert!(must_bind_before_connect( + &TcpBindOptions::default().with_bind_device(Some("eth0".to_owned())) + )); + assert!(!must_bind_before_connect(&TcpBindOptions::default())); + assert!(!must_bind_before_connect( + &TcpBindOptions::default().with_reuse_addr(native_reuse_addr_default()) + )); + assert!(must_bind_before_connect( + &TcpBindOptions::default().with_reuse_addr(!native_reuse_addr_default()) + )); + assert_eq!( + native_reuse_addr(&TcpBindOptions::default()), + native_reuse_addr_default() + ); + } +} diff --git a/easytier/src/socket/udp.rs b/easytier/src/socket/udp.rs new file mode 100644 index 00000000..b66bb037 --- /dev/null +++ b/easytier/src/socket/udp.rs @@ -0,0 +1,343 @@ +use std::{ + net::{IpAddr, Ipv4Addr, SocketAddr, SocketAddrV4}, + sync::Arc, +}; + +use async_trait::async_trait; +#[cfg(any(feature = "wireguard", test))] +use easytier_core::socket::{ + NetNamespace, + udp::{UdpSessionAcceptKind, UdpSessionListenRequest, UdpSessionSocketListener}, +}; +use easytier_core::socket::{ + SocketContext, + udp::{ + UdpBindOptions, UdpSocketPurpose, UdpSocketRecvMeta, UdpSocketSendMeta, VirtualUdpSocket, + VirtualUdpSocketFactory, + }, +}; +use tokio::net::UdpSocket; + +#[cfg(any(feature = "wireguard", test))] +use crate::host_runtime::{NativeHostRuntime, native_host_runtime}; +use crate::{ + common::netns::NetNS, + tunnel::common::{BindDev, bind}, +}; + +use super::udp_src; + +#[cfg(any(feature = "wireguard", test))] +pub(crate) type RuntimeUdpSessionSocketListener = UdpSessionSocketListener; + +#[cfg(any(feature = "wireguard", test))] +pub(crate) fn new_runtime_udp_session_listener( + url: url::Url, + mut request: UdpSessionListenRequest, + accept_kind: UdpSessionAcceptKind, + net_ns: NetNS, +) -> RuntimeUdpSessionSocketListener { + request.bind.context.netns = net_ns.name().map(NetNamespace::new); + let runtime = native_host_runtime(); + UdpSessionSocketListener::new_with_request(url, request, accept_kind, runtime) +} + +pub struct RuntimeUdpSocket { + socket: Arc, + context: SocketContext, +} + +impl RuntimeUdpSocket { + #[cfg(test)] + fn new(socket: Arc) -> Self { + Self::new_with_context(socket, SocketContext::default()) + } + + pub(crate) fn new_with_context(socket: Arc, context: SocketContext) -> Self { + if let Err(err) = udp_src::enable_recv_pktinfo(&socket) { + tracing::debug!(?err, "enable udp pktinfo failed"); + } + Self { socket, context } + } + + #[cfg(target_os = "windows")] + pub(crate) fn socket(&self) -> Arc { + self.socket.clone() + } +} + +#[async_trait] +impl VirtualUdpSocket for RuntimeUdpSocket { + fn local_addr(&self) -> std::io::Result { + self.socket.local_addr() + } + + fn socket_context(&self) -> SocketContext { + self.context.clone() + } + + async fn send_to(&self, data: &[u8], addr: SocketAddr) -> std::io::Result { + self.socket.send_to(data, addr).await + } + + async fn recv_from(&self, buf: &mut [u8]) -> std::io::Result<(usize, SocketAddr)> { + self.socket.recv_from(buf).await + } + + async fn send_to_with_meta( + &self, + data: &[u8], + addr: SocketAddr, + meta: UdpSocketSendMeta, + ) -> std::io::Result { + if let (Some(IpAddr::V6(src)), Some(ifindex), SocketAddr::V6(dst)) = + (meta.src_ip, meta.src_ifindex, addr) + { + return udp_src::send_to_with_src_ipv6(&self.socket, src, ifindex, dst, data); + } + if let Some(src_ip) = meta.src_ip { + return udp_src::send_to_with_src_ip(&self.socket, src_ip, addr, data).await; + } + self.socket.try_send_to(data, addr) + } + + async fn recv_from_with_meta( + &self, + buf: &mut [u8], + ) -> std::io::Result<(usize, SocketAddr, UdpSocketRecvMeta)> { + let (len, addr, dst_ip) = udp_src::recv_from_with_dst_ip(&self.socket, buf).await?; + Ok((len, addr, UdpSocketRecvMeta { dst_ip })) + } +} + +#[derive(Debug, Clone, Copy, Default)] +pub(crate) struct RuntimeUdpSocketFactory; + +impl RuntimeUdpSocketFactory { + pub(crate) fn new() -> Self { + Self + } + + fn bind_device_for(&self, options: &UdpBindOptions) -> BindDev { + if let Some(bind_device) = &options.bind_device { + return BindDev::from(bind_device.as_str()); + } + + if matches!( + options.purpose, + UdpSocketPurpose::DirectConnect + | UdpSocketPurpose::PortBoundListener + | UdpSocketPurpose::PortForward + ) { + return BindDev::Auto; + } + + BindDev::Disabled + } + + fn reuse_addr_for(&self, options: &UdpBindOptions) -> bool { + options.reuse_addr + || (matches!( + options.purpose, + UdpSocketPurpose::PortBoundListener + | UdpSocketPurpose::ProxyNat + | UdpSocketPurpose::PortForward + ) && !cfg!(target_os = "windows")) + } + + fn bind_udp_socket(&self, options: UdpBindOptions) -> anyhow::Result> { + let context = options.context.clone(); + let bind_addr = options + .local_addr + .unwrap_or_else(|| SocketAddr::V4(SocketAddrV4::new(Ipv4Addr::UNSPECIFIED, 0))); + let bind_device = self.bind_device_for(&options); + let reuse_addr = self.reuse_addr_for(&options); + let socket = bind::() + .addr(bind_addr) + .dev(bind_device) + .maybe_net_ns(Some(NetNS::from_socket_context(&context))) + .only_v6(options.only_v6) + .reuse_addr(reuse_addr) + .reuse_port(options.reuse_port) + .maybe_socket_mark(context.socket_mark) + .call()?; + Ok(Arc::new(RuntimeUdpSocket::new_with_context( + Arc::new(socket), + context, + ))) + } +} + +#[async_trait] +impl VirtualUdpSocketFactory for RuntimeUdpSocketFactory { + type Socket = RuntimeUdpSocket; + + async fn bind_udp(&self, options: UdpBindOptions) -> anyhow::Result> { + self.bind_udp_socket(options) + } +} + +#[cfg(test)] +mod tests { + use easytier_core::{ + socket::SocketListener, + socket::udp::{ + UdpSessionListenRequest, send_v4_hole_punch_control_packet, + send_v6_hole_punch_control_packet, + }, + }; + + use crate::host_runtime::native_host_runtime; + + use super::*; + + #[cfg(any(target_os = "linux", target_os = "android"))] + #[tokio::test] + async fn runtime_udp_socket_reports_ipv4_destination_ip() { + let socket = Arc::new(UdpSocket::bind("0.0.0.0:0").await.unwrap()); + let runtime_socket = RuntimeUdpSocket::new(socket.clone()); + let client = UdpSocket::bind("127.0.0.1:0").await.unwrap(); + client + .send_to( + b"pktinfo", + SocketAddr::from(([127, 0, 0, 1], socket.local_addr().unwrap().port())), + ) + .await + .unwrap(); + + let mut buf = [0; 32]; + let (len, _peer, meta) = runtime_socket.recv_from_with_meta(&mut buf).await.unwrap(); + + assert_eq!(&buf[..len], b"pktinfo"); + assert_eq!(meta.dst_ip, Some(std::net::IpAddr::V4(Ipv4Addr::LOCALHOST))); + } + + #[tokio::test] + async fn runtime_v4_hole_punch_control_packet_is_forwarded() { + let local_addr = SocketAddr::from(([0, 0, 0, 0], 0)); + let mut listener = new_runtime_udp_session_listener( + "udp://0.0.0.0:0".parse().unwrap(), + UdpSessionListenRequest::new(UdpBindOptions::port_bound_listener(local_addr)), + UdpSessionAcceptKind::EasyTierMux, + NetNS::new(None), + ); + listener.listen().await.unwrap(); + + let receiver = UdpSocket::bind("127.0.0.1:0").await.unwrap(); + let runtime = native_host_runtime(); + send_v4_hole_punch_control_packet( + runtime.as_ref(), + SocketContext::default(), + listener.local_url().port().unwrap(), + match receiver.local_addr().unwrap() { + SocketAddr::V4(addr) => addr, + SocketAddr::V6(_) => unreachable!(), + }, + ) + .await + .unwrap(); + + let mut buf = [0; 128]; + tokio::time::timeout( + std::time::Duration::from_secs(2), + receiver.recv_from(&mut buf), + ) + .await + .expect("timeout waiting for v4 hole-punch packet") + .unwrap(); + } + + #[tokio::test] + async fn runtime_v6_hole_punch_control_packet_is_forwarded() { + let local_addr = "[::]:0".parse().unwrap(); + let mut listener = new_runtime_udp_session_listener( + "udp://[::]:0".parse().unwrap(), + UdpSessionListenRequest::new(UdpBindOptions::port_bound_listener(local_addr)), + UdpSessionAcceptKind::EasyTierMux, + NetNS::new(None), + ); + listener.listen().await.unwrap(); + + let receiver = UdpSocket::bind("[::]:0").await.unwrap(); + let runtime = native_host_runtime(); + send_v6_hole_punch_control_packet( + runtime.as_ref(), + SocketContext::default(), + listener.local_url().port().unwrap(), + match receiver.local_addr().unwrap() { + SocketAddr::V6(addr) => addr, + SocketAddr::V4(_) => unreachable!(), + }, + None, + ) + .await + .unwrap(); + + let mut buf = [0; 128]; + tokio::time::timeout( + std::time::Duration::from_secs(2), + receiver.recv_from(&mut buf), + ) + .await + .expect("timeout waiting for v6 hole-punch packet") + .unwrap(); + } + + #[test] + fn factory_interprets_bind_defaults_by_purpose() { + let listener_addr = SocketAddr::from(([0, 0, 0, 0], 11010)); + let factory = RuntimeUdpSocketFactory::new(); + + assert!(matches!( + factory.bind_device_for(&UdpBindOptions::port_bound_listener(listener_addr)), + BindDev::Auto + )); + assert!(matches!( + factory.bind_device_for(&UdpBindOptions::direct_connect()), + BindDev::Auto + )); + assert!(matches!( + factory.bind_device_for(&UdpBindOptions::port_forward(listener_addr)), + BindDev::Auto + )); + assert!(matches!( + factory.bind_device_for(&UdpBindOptions::port_lease(listener_addr)), + BindDev::Disabled + )); + assert!(matches!( + factory.bind_device_for(&UdpBindOptions::hole_punch_control()), + BindDev::Disabled + )); + assert_eq!( + factory.reuse_addr_for(&UdpBindOptions::port_bound_listener(listener_addr)), + !cfg!(target_os = "windows") + ); + assert!(!factory.reuse_addr_for(&UdpBindOptions::hole_punch_control())); + assert_eq!( + factory.reuse_addr_for(&UdpBindOptions::proxy_nat()), + !cfg!(target_os = "windows") + ); + assert_eq!( + factory.reuse_addr_for(&UdpBindOptions::port_forward(listener_addr)), + !cfg!(target_os = "windows") + ); + assert!(!factory.reuse_addr_for(&UdpBindOptions::port_lease(listener_addr))); + } + + #[test] + fn factory_applies_listener_bind_device_option() { + let listener_addr = SocketAddr::from(([0, 0, 0, 0], 11010)); + let factory = RuntimeUdpSocketFactory::new(); + let options = UdpBindOptions::port_bound_listener(listener_addr) + .with_bind_device(Some("eth0".to_owned())); + + match factory.bind_device_for(&options) { + BindDev::Custom(dev) => assert_eq!(dev, "eth0"), + bind_device => panic!("unexpected bind device: {bind_device:?}"), + } + assert!(matches!( + factory.bind_device_for(&UdpBindOptions::hole_punch_control()), + BindDev::Disabled + )); + } +} diff --git a/easytier/src/socket/udp_src.rs b/easytier/src/socket/udp_src.rs new file mode 100644 index 00000000..f6e2355e --- /dev/null +++ b/easytier/src/socket/udp_src.rs @@ -0,0 +1,13 @@ +#[cfg(unix)] +#[path = "udp_src/unix.rs"] +mod platform; + +#[cfg(windows)] +#[path = "udp_src/windows.rs"] +mod platform; + +#[cfg(not(any(unix, windows)))] +#[path = "udp_src/fallback.rs"] +mod platform; + +pub(crate) use platform::*; diff --git a/easytier/src/socket/udp_src/fallback.rs b/easytier/src/socket/udp_src/fallback.rs new file mode 100644 index 00000000..a5857da8 --- /dev/null +++ b/easytier/src/socket/udp_src/fallback.rs @@ -0,0 +1,81 @@ +use std::{ + io, + net::{IpAddr, Ipv4Addr, Ipv6Addr, SocketAddr, SocketAddrV4, SocketAddrV6}, +}; + +use tokio::net::UdpSocket; + +pub(crate) fn enable_recv_pktinfo(_socket: &UdpSocket) -> io::Result<()> { + Ok(()) +} + +pub(crate) async fn recv_from_with_dst_ip( + socket: &UdpSocket, + buf: &mut [u8], +) -> io::Result<(usize, SocketAddr, Option)> { + let (len, addr) = socket.recv_from(buf).await?; + Ok((len, addr, None)) +} + +pub(crate) async fn send_to_with_src_ip( + socket: &UdpSocket, + src_ip: IpAddr, + dst_addr: SocketAddr, + buf: &[u8], +) -> io::Result { + match (src_ip, dst_addr) { + (IpAddr::V4(src), SocketAddr::V4(dst)) => { + send_to_with_src_ipv4(socket, src, dst, buf).await + } + (IpAddr::V4(src), SocketAddr::V6(dst)) => { + let Some(mapped_dst) = dst.ip().to_ipv4_mapped() else { + return Err(source_address_family_mismatch(src, dst)); + }; + send_to_with_src_ipv4(socket, src, SocketAddrV4::new(mapped_dst, dst.port()), buf).await + } + (IpAddr::V6(src), SocketAddr::V6(dst)) => { + socket + .async_io(tokio::io::Interest::WRITABLE, || { + send_to_with_src_ipv6(socket, src, 0, dst, buf) + }) + .await + } + (src, dst) => Err(source_address_family_mismatch(src, dst)), + } +} + +async fn send_to_with_src_ipv4( + socket: &UdpSocket, + _src_ip: Ipv4Addr, + dst_addr: SocketAddrV4, + buf: &[u8], +) -> io::Result { + socket + .async_io(tokio::io::Interest::WRITABLE, || { + socket.try_send_to(buf, SocketAddr::V4(dst_addr)) + }) + .await +} + +fn send_to_with_src_ipv6( + _socket: &UdpSocket, + _src_ip: Ipv6Addr, + _src_ifindex: u32, + _dst_addr: SocketAddrV6, + _buf: &[u8], +) -> io::Result { + Err(io::Error::new( + io::ErrorKind::Unsupported, + "sending UDP with a selected IPv6 source is not supported on this platform", + )) +} + +fn source_address_family_mismatch( + src: impl std::fmt::Display, + dst: impl std::fmt::Display, +) -> io::Error { + io::Error::new( + io::ErrorKind::InvalidInput, + format!("source address {src} does not match destination {dst} family"), + ) +} diff --git a/easytier/src/socket/udp_src/unix.rs b/easytier/src/socket/udp_src/unix.rs new file mode 100644 index 00000000..60962662 --- /dev/null +++ b/easytier/src/socket/udp_src/unix.rs @@ -0,0 +1,570 @@ +use std::{ + io, + net::{IpAddr, Ipv4Addr, Ipv6Addr, SocketAddr, SocketAddrV4, SocketAddrV6}, +}; + +use tokio::net::UdpSocket; + +pub(crate) fn enable_recv_pktinfo(socket: &UdpSocket) -> io::Result<()> { + use std::os::fd::AsRawFd; + + use nix::libc; + + let fd = socket.as_raw_fd(); + let enabled: libc::c_int = 1; + unsafe { + #[cfg(any(target_os = "linux", target_os = "android"))] + let _ = libc::setsockopt( + fd, + libc::IPPROTO_IP, + libc::IP_PKTINFO, + &enabled as *const _ as *const libc::c_void, + std::mem::size_of_val(&enabled) as libc::socklen_t, + ); + #[cfg(any( + target_os = "freebsd", + target_os = "openbsd", + target_os = "netbsd", + target_os = "macos", + target_os = "ios" + ))] + let _ = libc::setsockopt( + fd, + libc::IPPROTO_IP, + libc::IP_RECVDSTADDR, + &enabled as *const _ as *const libc::c_void, + std::mem::size_of_val(&enabled) as libc::socklen_t, + ); + let _ = libc::setsockopt( + fd, + libc::IPPROTO_IPV6, + libc::IPV6_RECVPKTINFO, + &enabled as *const _ as *const libc::c_void, + std::mem::size_of_val(&enabled) as libc::socklen_t, + ); + } + Ok(()) +} + +#[cfg(not(any(unix, windows)))] +pub(crate) fn enable_recv_pktinfo(_socket: &UdpSocket) -> io::Result<()> { + Ok(()) +} + +pub(crate) async fn recv_from_with_dst_ip( + socket: &UdpSocket, + buf: &mut [u8], +) -> io::Result<(usize, SocketAddr, Option)> { + socket + .async_io(tokio::io::Interest::READABLE, || { + loop { + match recv_from_with_dst_ip_once(socket, buf) { + Err(err) if err.kind() == io::ErrorKind::Interrupted => continue, + ret => break ret, + } + } + }) + .await +} + +#[cfg(not(any(unix, windows)))] +pub(crate) async fn recv_from_with_dst_ip( + socket: &UdpSocket, + buf: &mut [u8], +) -> io::Result<(usize, SocketAddr, Option)> { + let (len, addr) = socket.recv_from(buf).await?; + Ok((len, addr, None)) +} + +fn recv_from_with_dst_ip_once( + socket: &UdpSocket, + buf: &mut [u8], +) -> io::Result<(usize, SocketAddr, Option)> { + use std::{mem, os::fd::AsRawFd}; + + use nix::libc; + + #[repr(align(8))] + struct ControlBuffer([u8; 256]); + + fn sockaddr_to_socket_addr( + storage: &libc::sockaddr_storage, + len: libc::socklen_t, + ) -> io::Result { + match storage.ss_family as libc::c_int { + libc::AF_INET => { + if (len as usize) < mem::size_of::() { + return Err(io::Error::new( + io::ErrorKind::InvalidData, + "short IPv4 sockaddr", + )); + } + let addr = unsafe { &*(storage as *const _ as *const libc::sockaddr_in) }; + let ip = Ipv4Addr::from(u32::from_be(addr.sin_addr.s_addr)); + let port = u16::from_be(addr.sin_port); + Ok(SocketAddr::V4(SocketAddrV4::new(ip, port))) + } + libc::AF_INET6 => { + if (len as usize) < mem::size_of::() { + return Err(io::Error::new( + io::ErrorKind::InvalidData, + "short IPv6 sockaddr", + )); + } + let addr = unsafe { &*(storage as *const _ as *const libc::sockaddr_in6) }; + let ip = Ipv6Addr::from(addr.sin6_addr.s6_addr); + let port = u16::from_be(addr.sin6_port); + Ok(SocketAddr::V6(SocketAddrV6::new( + ip, + port, + addr.sin6_flowinfo, + addr.sin6_scope_id, + ))) + } + _ => Err(io::Error::new( + io::ErrorKind::InvalidData, + "unsupported UDP sockaddr family", + )), + } + } + + let mut iov = libc::iovec { + iov_base: buf.as_mut_ptr() as *mut libc::c_void, + iov_len: buf.len(), + }; + let mut name = unsafe { mem::zeroed::() }; + let mut control = ControlBuffer([0u8; 256]); + let mut msg = unsafe { mem::zeroed::() }; + msg.msg_name = &mut name as *mut _ as *mut libc::c_void; + msg.msg_namelen = mem::size_of::() as _; + msg.msg_iov = &mut iov; + msg.msg_iovlen = 1; + msg.msg_control = control.0.as_mut_ptr() as *mut libc::c_void; + msg.msg_controllen = control.0.len() as _; + + let len = unsafe { libc::recvmsg(socket.as_raw_fd(), &mut msg, 0) }; + if len < 0 { + return Err(io::Error::last_os_error()); + } + + let remote_addr = sockaddr_to_socket_addr(&name, msg.msg_namelen)?; + let mut dst_ip = None; + unsafe { + let mut cmsg = libc::CMSG_FIRSTHDR(&msg); + while !cmsg.is_null() { + #[cfg(any(target_os = "linux", target_os = "android"))] + { + if (*cmsg).cmsg_level == libc::IPPROTO_IP && (*cmsg).cmsg_type == libc::IP_PKTINFO { + let pktinfo = &*(libc::CMSG_DATA(cmsg) as *const libc::in_pktinfo); + dst_ip = Some(IpAddr::V4(Ipv4Addr::from(u32::from_be( + pktinfo.ipi_addr.s_addr, + )))); + } + } + #[cfg(any( + target_os = "freebsd", + target_os = "openbsd", + target_os = "netbsd", + target_os = "macos", + target_os = "ios" + ))] + { + if (*cmsg).cmsg_level == libc::IPPROTO_IP + && (*cmsg).cmsg_type == libc::IP_RECVDSTADDR + { + let addr = &*(libc::CMSG_DATA(cmsg) as *const libc::in_addr); + dst_ip = Some(IpAddr::V4(Ipv4Addr::from(u32::from_be(addr.s_addr)))); + } + } + if (*cmsg).cmsg_level == libc::IPPROTO_IPV6 && (*cmsg).cmsg_type == libc::IPV6_PKTINFO { + let pktinfo = &*(libc::CMSG_DATA(cmsg) as *const libc::in6_pktinfo); + dst_ip = Some(IpAddr::V6(Ipv6Addr::from(pktinfo.ipi6_addr.s6_addr))); + } + cmsg = libc::CMSG_NXTHDR(&msg, cmsg); + } + } + + Ok((len as usize, remote_addr, dst_ip)) +} + +#[cfg(any(target_os = "linux", target_os = "android"))] +pub(crate) async fn send_to_with_src_ip( + socket: &UdpSocket, + src_ip: IpAddr, + dst_addr: SocketAddr, + buf: &[u8], +) -> io::Result { + socket + .async_io(tokio::io::Interest::WRITABLE, || { + send_to_with_src_ip_raw(socket, src_ip, dst_addr, buf) + }) + .await +} + +#[cfg(not(any(target_os = "linux", target_os = "android")))] +pub(crate) async fn send_to_with_src_ip( + socket: &UdpSocket, + src_ip: IpAddr, + dst_addr: SocketAddr, + buf: &[u8], +) -> io::Result { + match (src_ip, dst_addr) { + (IpAddr::V4(src), SocketAddr::V4(dst)) => { + socket + .async_io(tokio::io::Interest::WRITABLE, || { + send_to_with_src_ipv4_to_addr(socket, src, SocketAddr::V4(dst), buf) + }) + .await + } + (IpAddr::V4(src), SocketAddr::V6(dst)) => { + if dst.ip().to_ipv4_mapped().is_none() { + return Err(io::Error::new( + io::ErrorKind::InvalidInput, + format!("source address {src} does not match destination {dst} family"), + )); + } + socket + .async_io(tokio::io::Interest::WRITABLE, || { + send_to_with_src_ipv4_to_addr(socket, src, SocketAddr::V6(dst), buf) + }) + .await + } + (IpAddr::V6(src), SocketAddr::V6(dst)) => { + socket + .async_io(tokio::io::Interest::WRITABLE, || { + send_to_with_src_ipv6(socket, src, 0, dst, buf) + }) + .await + } + (src, dst) => Err(io::Error::new( + io::ErrorKind::InvalidInput, + format!("source address {src} does not match destination {dst} family"), + )), + } +} + +#[cfg(not(any(target_os = "linux", target_os = "android")))] +fn send_to_with_src_ipv4_to_addr( + socket: &UdpSocket, + src_ip: Ipv4Addr, + dst_addr: SocketAddr, + buf: &[u8], +) -> io::Result { + match dst_addr { + SocketAddr::V4(dst) => send_to_with_src_ipv4(socket, src_ip, dst, buf), + SocketAddr::V6(dst) => { + let Some(mapped_dst) = dst.ip().to_ipv4_mapped() else { + return Err(io::Error::new( + io::ErrorKind::InvalidInput, + format!("source address {src_ip} does not match destination {dst} family"), + )); + }; + send_to_with_src_ipv4_mapped_v6(socket, src_ip, dst, mapped_dst, buf) + } + } +} + +#[cfg(not(any(target_os = "linux", target_os = "android", windows)))] +fn send_to_with_src_ipv4_mapped_v6( + socket: &UdpSocket, + src_ip: Ipv4Addr, + dst_addr: SocketAddrV6, + mapped_dst: Ipv4Addr, + buf: &[u8], +) -> io::Result { + send_to_with_src_ipv4( + socket, + src_ip, + SocketAddrV4::new(mapped_dst, dst_addr.port()), + buf, + ) +} + +#[cfg(any(target_os = "linux", target_os = "android"))] +fn send_to_with_src_ip_raw( + socket: &UdpSocket, + src_ip: IpAddr, + dst_addr: SocketAddr, + buf: &[u8], +) -> io::Result { + match (src_ip, dst_addr) { + (IpAddr::V4(src), SocketAddr::V4(dst)) => send_to_with_src_ipv4(socket, src, dst, buf), + (IpAddr::V4(src), SocketAddr::V6(dst)) => { + let Some(mapped_dst) = dst.ip().to_ipv4_mapped() else { + return Err(io::Error::new( + io::ErrorKind::InvalidInput, + format!("source address {src} does not match destination {dst} family"), + )); + }; + send_to_with_src_ipv4(socket, src, SocketAddrV4::new(mapped_dst, dst.port()), buf) + } + (IpAddr::V6(src), SocketAddr::V6(dst)) => send_to_with_src_ipv6(socket, src, 0, dst, buf), + (src, dst) => Err(io::Error::new( + io::ErrorKind::InvalidInput, + format!("source address {src} does not match destination {dst} family"), + )), + } +} + +#[cfg(any(target_os = "linux", target_os = "android"))] +fn send_to_with_src_ipv4( + socket: &UdpSocket, + src_ip: Ipv4Addr, + dst_addr: SocketAddrV4, + buf: &[u8], +) -> io::Result { + use std::{mem, os::fd::AsRawFd, ptr}; + + use nix::libc; + + #[repr(align(8))] + struct ControlBuffer([u8; 128]); + + let pktinfo = libc::in_pktinfo { + ipi_ifindex: 0, + ipi_spec_dst: libc::in_addr { + s_addr: u32::from(src_ip).to_be(), + }, + ipi_addr: libc::in_addr { s_addr: 0 }, + }; + let mut iov = libc::iovec { + iov_base: buf.as_ptr() as *mut libc::c_void, + iov_len: buf.len(), + }; + let dst_addr = socket2::SockAddr::from(std::net::SocketAddr::V4(dst_addr)); + let control_len = + unsafe { libc::CMSG_SPACE(mem::size_of::() as libc::c_uint) as usize }; + let mut control = ControlBuffer([0u8; 128]); + if control_len > control.0.len() { + return Err(io::Error::new( + io::ErrorKind::InvalidInput, + "IPv4 packet info control buffer is too small", + )); + } + + let mut msg = unsafe { mem::zeroed::() }; + msg.msg_name = dst_addr.as_ptr() as *mut libc::c_void; + msg.msg_namelen = dst_addr.len() as _; + msg.msg_iov = &mut iov; + msg.msg_iovlen = 1; + msg.msg_control = control.0.as_mut_ptr() as *mut libc::c_void; + msg.msg_controllen = control_len as _; + + unsafe { + let cmsg = libc::CMSG_FIRSTHDR(&msg); + if cmsg.is_null() { + return Err(io::Error::new( + io::ErrorKind::InvalidInput, + "IPv4 packet info control buffer is invalid", + )); + } + (*cmsg).cmsg_level = libc::IPPROTO_IP; + (*cmsg).cmsg_type = libc::IP_PKTINFO; + (*cmsg).cmsg_len = libc::CMSG_LEN(mem::size_of::() as libc::c_uint) as _; + ptr::write(libc::CMSG_DATA(cmsg) as *mut libc::in_pktinfo, pktinfo); + + let ret = libc::sendmsg(socket.as_raw_fd(), &msg, 0); + if ret < 0 { + Err(io::Error::last_os_error()) + } else { + Ok(ret as usize) + } + } +} + +#[cfg(any( + target_os = "freebsd", + target_os = "openbsd", + target_os = "netbsd", + target_os = "macos", + target_os = "ios" +))] +fn send_to_with_src_ipv4( + socket: &UdpSocket, + src_ip: Ipv4Addr, + dst_addr: SocketAddrV4, + buf: &[u8], +) -> io::Result { + use std::{mem, os::fd::AsRawFd, ptr}; + + use nix::libc; + + #[repr(align(8))] + struct ControlBuffer([u8; 128]); + + if let Ok(SocketAddr::V4(local_addr)) = socket.local_addr() { + if !local_addr.ip().is_unspecified() { + return socket.try_send_to(buf, SocketAddr::V4(dst_addr)); + } + } + + let src_addr = libc::in_addr { + s_addr: u32::from(src_ip).to_be(), + }; + let mut iov = libc::iovec { + iov_base: buf.as_ptr() as *mut libc::c_void, + iov_len: buf.len(), + }; + let dst_addr = socket2::SockAddr::from(std::net::SocketAddr::V4(dst_addr)); + let control_len = + unsafe { libc::CMSG_SPACE(mem::size_of::() as libc::c_uint) as usize }; + let mut control = ControlBuffer([0u8; 128]); + if control_len > control.0.len() { + return Err(io::Error::new( + io::ErrorKind::InvalidInput, + "IPv4 source address control buffer is too small", + )); + } + + let mut msg = unsafe { mem::zeroed::() }; + msg.msg_name = dst_addr.as_ptr() as *mut libc::c_void; + msg.msg_namelen = dst_addr.len() as _; + msg.msg_iov = &mut iov; + msg.msg_iovlen = 1; + msg.msg_control = control.0.as_mut_ptr() as *mut libc::c_void; + msg.msg_controllen = control_len as _; + + unsafe { + let cmsg = libc::CMSG_FIRSTHDR(&msg); + if cmsg.is_null() { + return Err(io::Error::new( + io::ErrorKind::InvalidInput, + "IPv4 source address control buffer is invalid", + )); + } + (*cmsg).cmsg_level = libc::IPPROTO_IP; + (*cmsg).cmsg_type = libc::IP_RECVDSTADDR; + (*cmsg).cmsg_len = libc::CMSG_LEN(mem::size_of::() as libc::c_uint) as _; + ptr::write(libc::CMSG_DATA(cmsg) as *mut libc::in_addr, src_addr); + + let ret = libc::sendmsg(socket.as_raw_fd(), &msg, 0); + if ret < 0 { + Err(io::Error::last_os_error()) + } else { + Ok(ret as usize) + } + } +} + +#[cfg(not(any( + target_os = "linux", + target_os = "android", + target_os = "freebsd", + target_os = "openbsd", + target_os = "netbsd", + target_os = "macos", + target_os = "ios", + windows +)))] +fn send_to_with_src_ipv4( + socket: &UdpSocket, + _src_ip: Ipv4Addr, + dst_addr: SocketAddrV4, + buf: &[u8], +) -> io::Result { + socket.try_send_to(buf, SocketAddr::V4(dst_addr)) +} + +pub(crate) fn send_to_with_src_ipv6( + socket: &UdpSocket, + src_ip: Ipv6Addr, + src_ifindex: u32, + dst_addr: SocketAddrV6, + buf: &[u8], +) -> io::Result { + #[cfg(target_env = "ohos")] + { + let _ = (socket, src_ip, src_ifindex, dst_addr, buf); + return Err(io::Error::new( + io::ErrorKind::Unsupported, + "sending UDP with a selected IPv6 source is not supported on OHOS", + )); + } + + #[cfg(not(target_env = "ohos"))] + { + use std::{mem, os::fd::AsRawFd, ptr}; + + use nix::libc; + + #[repr(align(8))] + struct ControlBuffer([u8; 128]); + + #[cfg(target_os = "android")] + let ipi6_ifindex: libc::c_int = i32::try_from(src_ifindex).map_err(|_| { + io::Error::new( + io::ErrorKind::InvalidInput, + "IPv6 source interface index is out of range", + ) + })?; + #[cfg(not(target_os = "android"))] + let ipi6_ifindex: libc::c_uint = src_ifindex; + + let pktinfo = libc::in6_pktinfo { + ipi6_addr: libc::in6_addr { + s6_addr: src_ip.octets(), + }, + ipi6_ifindex, + }; + let mut iov = libc::iovec { + iov_base: buf.as_ptr() as *mut libc::c_void, + iov_len: buf.len(), + }; + let dst_addr = socket2::SockAddr::from(std::net::SocketAddr::V6(dst_addr)); + let control_len = unsafe { + libc::CMSG_SPACE(mem::size_of::() as libc::c_uint) as usize + }; + let mut control = ControlBuffer([0u8; 128]); + if control_len > control.0.len() { + return Err(io::Error::new( + io::ErrorKind::InvalidInput, + "IPv6 packet info control buffer is too small", + )); + } + + let mut msg = unsafe { mem::zeroed::() }; + msg.msg_name = dst_addr.as_ptr() as *mut libc::c_void; + msg.msg_namelen = dst_addr.len() as _; + msg.msg_iov = &mut iov; + msg.msg_iovlen = 1; + msg.msg_control = control.0.as_mut_ptr() as *mut libc::c_void; + msg.msg_controllen = control_len as _; + msg.msg_flags = 0; + + unsafe { + let cmsg = libc::CMSG_FIRSTHDR(&msg); + if cmsg.is_null() { + return Err(io::Error::new( + io::ErrorKind::InvalidInput, + "IPv6 packet info control buffer is invalid", + )); + } + (*cmsg).cmsg_level = libc::IPPROTO_IPV6; + (*cmsg).cmsg_type = libc::IPV6_PKTINFO; + (*cmsg).cmsg_len = + libc::CMSG_LEN(mem::size_of::() as libc::c_uint) as _; + ptr::write(libc::CMSG_DATA(cmsg) as *mut libc::in6_pktinfo, pktinfo); + + let ret = libc::sendmsg(socket.as_raw_fd(), &msg, 0); + if ret < 0 { + Err(io::Error::last_os_error()) + } else { + Ok(ret as usize) + } + } + } +} + +#[cfg(not(any(unix, windows)))] +pub(crate) fn send_to_with_src_ipv6( + _socket: &UdpSocket, + _src_ip: Ipv6Addr, + _src_ifindex: u32, + _dst_addr: SocketAddrV6, + _buf: &[u8], +) -> io::Result { + Err(io::Error::new( + io::ErrorKind::Unsupported, + "sending UDP with a selected IPv6 source is not supported on this platform", + )) +} diff --git a/easytier/src/socket/udp_src/windows.rs b/easytier/src/socket/udp_src/windows.rs new file mode 100644 index 00000000..813d429f --- /dev/null +++ b/easytier/src/socket/udp_src/windows.rs @@ -0,0 +1,716 @@ +use std::{ + io, + net::{IpAddr, Ipv4Addr, Ipv6Addr, SocketAddr, SocketAddrV4, SocketAddrV6}, +}; + +use std::sync::OnceLock; + +use tokio::net::UdpSocket; + +pub(crate) fn enable_recv_pktinfo(socket: &UdpSocket) -> io::Result<()> { + use std::os::windows::io::AsRawSocket; + + use windows::Win32::Networking::WinSock::{ + IP_PKTINFO, IPPROTO_IP, IPPROTO_IPV6, IPV6_PKTINFO, SOCKET, setsockopt, + }; + + let enabled = 1u32.to_ne_bytes(); + unsafe { + let _ = setsockopt( + SOCKET(socket.as_raw_socket() as usize), + IPPROTO_IP.0, + IP_PKTINFO, + Some(&enabled), + ); + let _ = setsockopt( + SOCKET(socket.as_raw_socket() as usize), + IPPROTO_IPV6.0, + IPV6_PKTINFO, + Some(&enabled), + ); + } + Ok(()) +} + +#[cfg(not(any(unix, windows)))] +pub(crate) fn enable_recv_pktinfo(_socket: &UdpSocket) -> io::Result<()> { + Ok(()) +} + +pub(crate) async fn recv_from_with_dst_ip( + socket: &UdpSocket, + buf: &mut [u8], +) -> io::Result<(usize, SocketAddr, Option)> { + socket + .async_io(tokio::io::Interest::READABLE, || { + loop { + match recv_from_with_dst_ip_once(socket, buf) { + Err(err) if err.kind() == io::ErrorKind::Interrupted => continue, + ret => break ret, + } + } + }) + .await +} + +#[cfg(not(any(unix, windows)))] +pub(crate) async fn recv_from_with_dst_ip( + socket: &UdpSocket, + buf: &mut [u8], +) -> io::Result<(usize, SocketAddr, Option)> { + let (len, addr) = socket.recv_from(buf).await?; + Ok((len, addr, None)) +} + +fn windows_cmsghdr_align(length: usize) -> usize { + use windows::Win32::Networking::WinSock::CMSGHDR; + + (length + std::mem::align_of::() - 1) & !(std::mem::align_of::() - 1) +} + +fn windows_cmsgdata_align(length: usize) -> usize { + (length + std::mem::align_of::() - 1) & !(std::mem::align_of::() - 1) +} + +fn windows_cmsg_len(length: usize) -> usize { + use windows::Win32::Networking::WinSock::CMSGHDR; + + windows_cmsgdata_align(std::mem::size_of::()) + length +} + +fn windows_cmsg_space(length: usize) -> usize { + use windows::Win32::Networking::WinSock::CMSGHDR; + + windows_cmsgdata_align(std::mem::size_of::() + windows_cmsghdr_align(length)) +} + +fn windows_cmsg_data(cmsg: *mut windows::Win32::Networking::WinSock::CMSGHDR) -> *mut u8 { + (cmsg as usize + + windows_cmsgdata_align(std::mem::size_of::< + windows::Win32::Networking::WinSock::CMSGHDR, + >())) as *mut u8 +} + +fn wsa_recvmsg_ptr() -> windows::Win32::Networking::WinSock::LPFN_WSARECVMSG { + use std::mem; + + use windows::Win32::Networking::WinSock::{ + AF_INET, INVALID_SOCKET, IPPROTO_UDP, SIO_GET_EXTENSION_FUNCTION_POINTER, SOCK_DGRAM, + WSAID_WSARECVMSG, WSAIoctl, closesocket, socket, + }; + + static WSA_RECVMSG: OnceLock = + OnceLock::new(); + + *WSA_RECVMSG.get_or_init(|| unsafe { + let Ok(socket) = socket(AF_INET.0.into(), SOCK_DGRAM, IPPROTO_UDP.0) else { + return None; + }; + if socket == INVALID_SOCKET { + return None; + } + + let guid = WSAID_WSARECVMSG; + let mut recvmsg = None; + let mut len = 0; + let ret = WSAIoctl( + socket, + SIO_GET_EXTENSION_FUNCTION_POINTER, + Some(&guid as *const _ as *const _), + mem::size_of_val(&guid) as u32, + Some(&mut recvmsg as *mut _ as *mut _), + mem::size_of_val(&recvmsg) as u32, + &mut len, + None, + None, + ); + closesocket(socket); + if ret == -1 || len as usize != mem::size_of_val(&recvmsg) { + None + } else { + recvmsg + } + }) +} + +fn sockaddr_inet_to_socket_addr( + addr: &windows::Win32::Networking::WinSock::SOCKADDR_INET, +) -> io::Result { + use windows::Win32::Networking::WinSock::{AF_INET, AF_INET6}; + + let family = unsafe { addr.si_family }; + if family == AF_INET { + let addr = unsafe { addr.Ipv4 }; + let ip = Ipv4Addr::from(u32::from_be(unsafe { addr.sin_addr.S_un.S_addr })); + let port = u16::from_be(addr.sin_port); + return Ok(SocketAddr::V4(SocketAddrV4::new(ip, port))); + } + if family == AF_INET6 { + let addr = unsafe { addr.Ipv6 }; + let ip = Ipv6Addr::from(unsafe { addr.sin6_addr.u.Byte }); + let port = u16::from_be(addr.sin6_port); + let scope_id = unsafe { addr.Anonymous.sin6_scope_id }; + return Ok(SocketAddr::V6(SocketAddrV6::new( + ip, + port, + addr.sin6_flowinfo, + scope_id, + ))); + } + Err(io::Error::new( + io::ErrorKind::InvalidData, + "unsupported UDP sockaddr family", + )) +} + +fn recv_from_with_dst_ip_once( + socket: &UdpSocket, + buf: &mut [u8], +) -> io::Result<(usize, SocketAddr, Option)> { + use std::{mem, os::windows::io::AsRawSocket, ptr}; + + use windows::{ + Win32::Networking::WinSock::{ + CMSGHDR, IN_PKTINFO, IN6_PKTINFO, IP_PKTINFO, IPPROTO_IP, IPPROTO_IPV6, IPV6_PKTINFO, + SOCKADDR_INET, SOCKET, SOCKET_ERROR, WSABUF, WSAGetLastError, WSAMSG, + }, + core::PSTR, + }; + + #[repr(align(8))] + struct ControlBuffer([u8; 256]); + + let Some(wsa_recvmsg) = wsa_recvmsg_ptr() else { + return Err(io::Error::new( + io::ErrorKind::Unsupported, + "WSARecvMsg is not supported", + )); + }; + + let mut source = unsafe { mem::zeroed::() }; + let mut data = WSABUF { + len: buf.len() as u32, + buf: PSTR(buf.as_mut_ptr()), + }; + let mut control = ControlBuffer([0u8; 256]); + let mut msg = WSAMSG { + name: &mut source as *mut _ as *mut _, + namelen: mem::size_of_val(&source) as i32, + lpBuffers: &mut data, + dwBufferCount: 1, + Control: WSABUF { + len: control.0.len() as u32, + buf: PSTR(control.0.as_mut_ptr()), + }, + dwFlags: 0, + }; + + let mut len = 0; + let ret = unsafe { + wsa_recvmsg( + SOCKET(socket.as_raw_socket() as usize), + &mut msg, + &mut len, + ptr::null_mut(), + None, + ) + }; + if ret == SOCKET_ERROR { + return Err(io::Error::from_raw_os_error(unsafe { WSAGetLastError().0 })); + } + + let remote_addr = sockaddr_inet_to_socket_addr(&source)?; + let mut dst_ip = None; + let control_start = msg.Control.buf.0 as usize; + let control_end = control_start + msg.Control.len as usize; + let mut cmsg_ptr = control_start; + while cmsg_ptr + mem::size_of::() <= control_end { + let cmsg = cmsg_ptr as *const CMSGHDR; + let cmsg_len = unsafe { (*cmsg).cmsg_len }; + if cmsg_len < mem::size_of::() || cmsg_ptr + cmsg_len > control_end { + break; + } + match (unsafe { (*cmsg).cmsg_level }, unsafe { (*cmsg).cmsg_type }) { + (level, cmsg_type) if level == IPPROTO_IP.0 && cmsg_type == IP_PKTINFO => { + let pktinfo = unsafe { &*(windows_cmsg_data(cmsg as *mut _) as *const IN_PKTINFO) }; + dst_ip = Some(IpAddr::V4(Ipv4Addr::from(u32::from_be(unsafe { + pktinfo.ipi_addr.S_un.S_addr + })))); + } + (level, cmsg_type) if level == IPPROTO_IPV6.0 && cmsg_type == IPV6_PKTINFO => { + let pktinfo = + unsafe { &*(windows_cmsg_data(cmsg as *mut _) as *const IN6_PKTINFO) }; + dst_ip = Some(IpAddr::V6(Ipv6Addr::from(unsafe { + pktinfo.ipi6_addr.u.Byte + }))); + } + _ => {} + } + cmsg_ptr += windows_cmsghdr_align(cmsg_len); + } + + Ok((len as usize, remote_addr, dst_ip)) +} + +#[cfg(any(target_os = "linux", target_os = "android"))] +pub(crate) async fn send_to_with_src_ip( + socket: &UdpSocket, + src_ip: IpAddr, + dst_addr: SocketAddr, + buf: &[u8], +) -> io::Result { + socket + .async_io(tokio::io::Interest::WRITABLE, || { + send_to_with_src_ip_raw(socket, src_ip, dst_addr, buf) + }) + .await +} + +#[cfg(not(any(target_os = "linux", target_os = "android")))] +pub(crate) async fn send_to_with_src_ip( + socket: &UdpSocket, + src_ip: IpAddr, + dst_addr: SocketAddr, + buf: &[u8], +) -> io::Result { + match (src_ip, dst_addr) { + (IpAddr::V4(src), SocketAddr::V4(dst)) => { + socket + .async_io(tokio::io::Interest::WRITABLE, || { + send_to_with_src_ipv4_to_addr(socket, src, SocketAddr::V4(dst), buf) + }) + .await + } + (IpAddr::V4(src), SocketAddr::V6(dst)) => { + if dst.ip().to_ipv4_mapped().is_none() { + return Err(io::Error::new( + io::ErrorKind::InvalidInput, + format!("source address {src} does not match destination {dst} family"), + )); + } + socket + .async_io(tokio::io::Interest::WRITABLE, || { + send_to_with_src_ipv4_to_addr(socket, src, SocketAddr::V6(dst), buf) + }) + .await + } + (IpAddr::V6(src), SocketAddr::V6(dst)) => { + socket + .async_io(tokio::io::Interest::WRITABLE, || { + send_to_with_src_ipv6(socket, src, 0, dst, buf) + }) + .await + } + (src, dst) => Err(io::Error::new( + io::ErrorKind::InvalidInput, + format!("source address {src} does not match destination {dst} family"), + )), + } +} + +#[cfg(not(any(target_os = "linux", target_os = "android")))] +fn send_to_with_src_ipv4_to_addr( + socket: &UdpSocket, + src_ip: Ipv4Addr, + dst_addr: SocketAddr, + buf: &[u8], +) -> io::Result { + match dst_addr { + SocketAddr::V4(dst) => send_to_with_src_ipv4(socket, src_ip, dst, buf), + SocketAddr::V6(dst) => { + let Some(mapped_dst) = dst.ip().to_ipv4_mapped() else { + return Err(io::Error::new( + io::ErrorKind::InvalidInput, + format!("source address {src_ip} does not match destination {dst} family"), + )); + }; + send_to_with_src_ipv4_mapped_v6(socket, src_ip, dst, mapped_dst, buf) + } + } +} + +fn send_to_with_src_ipv4_mapped_v6( + socket: &UdpSocket, + src_ip: Ipv4Addr, + dst_addr: SocketAddrV6, + _mapped_dst: Ipv4Addr, + buf: &[u8], +) -> io::Result { + send_to_with_src_ipv4_windows(socket, src_ip, SocketAddr::V6(dst_addr), buf) +} + +#[cfg(not(any(target_os = "linux", target_os = "android", windows)))] +fn send_to_with_src_ipv4_mapped_v6( + socket: &UdpSocket, + src_ip: Ipv4Addr, + dst_addr: SocketAddrV6, + mapped_dst: Ipv4Addr, + buf: &[u8], +) -> io::Result { + send_to_with_src_ipv4( + socket, + src_ip, + SocketAddrV4::new(mapped_dst, dst_addr.port()), + buf, + ) +} + +#[cfg(any(target_os = "linux", target_os = "android"))] +fn send_to_with_src_ip_raw( + socket: &UdpSocket, + src_ip: IpAddr, + dst_addr: SocketAddr, + buf: &[u8], +) -> io::Result { + match (src_ip, dst_addr) { + (IpAddr::V4(src), SocketAddr::V4(dst)) => send_to_with_src_ipv4(socket, src, dst, buf), + (IpAddr::V4(src), SocketAddr::V6(dst)) => { + let Some(mapped_dst) = dst.ip().to_ipv4_mapped() else { + return Err(io::Error::new( + io::ErrorKind::InvalidInput, + format!("source address {src} does not match destination {dst} family"), + )); + }; + send_to_with_src_ipv4(socket, src, SocketAddrV4::new(mapped_dst, dst.port()), buf) + } + (IpAddr::V6(src), SocketAddr::V6(dst)) => send_to_with_src_ipv6(socket, src, 0, dst, buf), + (src, dst) => Err(io::Error::new( + io::ErrorKind::InvalidInput, + format!("source address {src} does not match destination {dst} family"), + )), + } +} + +#[cfg(any(target_os = "linux", target_os = "android"))] +fn send_to_with_src_ipv4( + socket: &UdpSocket, + src_ip: Ipv4Addr, + dst_addr: SocketAddrV4, + buf: &[u8], +) -> io::Result { + use std::{mem, os::fd::AsRawFd, ptr}; + + use nix::libc; + + #[repr(align(8))] + struct ControlBuffer([u8; 128]); + + let pktinfo = libc::in_pktinfo { + ipi_ifindex: 0, + ipi_spec_dst: libc::in_addr { + s_addr: u32::from(src_ip).to_be(), + }, + ipi_addr: libc::in_addr { s_addr: 0 }, + }; + let mut iov = libc::iovec { + iov_base: buf.as_ptr() as *mut libc::c_void, + iov_len: buf.len(), + }; + let dst_addr = socket2::SockAddr::from(std::net::SocketAddr::V4(dst_addr)); + let control_len = + unsafe { libc::CMSG_SPACE(mem::size_of::() as libc::c_uint) as usize }; + let mut control = ControlBuffer([0u8; 128]); + if control_len > control.0.len() { + return Err(io::Error::new( + io::ErrorKind::InvalidInput, + "IPv4 packet info control buffer is too small", + )); + } + + let mut msg = unsafe { mem::zeroed::() }; + msg.msg_name = dst_addr.as_ptr() as *mut libc::c_void; + msg.msg_namelen = dst_addr.len() as _; + msg.msg_iov = &mut iov; + msg.msg_iovlen = 1; + msg.msg_control = control.0.as_mut_ptr() as *mut libc::c_void; + msg.msg_controllen = control_len as _; + + unsafe { + let cmsg = libc::CMSG_FIRSTHDR(&msg); + if cmsg.is_null() { + return Err(io::Error::new( + io::ErrorKind::InvalidInput, + "IPv4 packet info control buffer is invalid", + )); + } + (*cmsg).cmsg_level = libc::IPPROTO_IP; + (*cmsg).cmsg_type = libc::IP_PKTINFO; + (*cmsg).cmsg_len = libc::CMSG_LEN(mem::size_of::() as libc::c_uint) as _; + ptr::write(libc::CMSG_DATA(cmsg) as *mut libc::in_pktinfo, pktinfo); + + let ret = libc::sendmsg(socket.as_raw_fd(), &msg, 0); + if ret < 0 { + Err(io::Error::last_os_error()) + } else { + Ok(ret as usize) + } + } +} + +#[cfg(any( + target_os = "freebsd", + target_os = "openbsd", + target_os = "netbsd", + target_os = "macos", + target_os = "ios" +))] +fn send_to_with_src_ipv4( + socket: &UdpSocket, + src_ip: Ipv4Addr, + dst_addr: SocketAddrV4, + buf: &[u8], +) -> io::Result { + use std::{mem, os::fd::AsRawFd, ptr}; + + use nix::libc; + + #[repr(align(8))] + struct ControlBuffer([u8; 128]); + + if let Ok(SocketAddr::V4(local_addr)) = socket.local_addr() { + if !local_addr.ip().is_unspecified() { + return socket.try_send_to(buf, SocketAddr::V4(dst_addr)); + } + } + + let src_addr = libc::in_addr { + s_addr: u32::from(src_ip).to_be(), + }; + let mut iov = libc::iovec { + iov_base: buf.as_ptr() as *mut libc::c_void, + iov_len: buf.len(), + }; + let dst_addr = socket2::SockAddr::from(std::net::SocketAddr::V4(dst_addr)); + let control_len = + unsafe { libc::CMSG_SPACE(mem::size_of::() as libc::c_uint) as usize }; + let mut control = ControlBuffer([0u8; 128]); + if control_len > control.0.len() { + return Err(io::Error::new( + io::ErrorKind::InvalidInput, + "IPv4 source address control buffer is too small", + )); + } + + let mut msg = unsafe { mem::zeroed::() }; + msg.msg_name = dst_addr.as_ptr() as *mut libc::c_void; + msg.msg_namelen = dst_addr.len() as _; + msg.msg_iov = &mut iov; + msg.msg_iovlen = 1; + msg.msg_control = control.0.as_mut_ptr() as *mut libc::c_void; + msg.msg_controllen = control_len as _; + + unsafe { + let cmsg = libc::CMSG_FIRSTHDR(&msg); + if cmsg.is_null() { + return Err(io::Error::new( + io::ErrorKind::InvalidInput, + "IPv4 source address control buffer is invalid", + )); + } + (*cmsg).cmsg_level = libc::IPPROTO_IP; + (*cmsg).cmsg_type = libc::IP_RECVDSTADDR; + (*cmsg).cmsg_len = libc::CMSG_LEN(mem::size_of::() as libc::c_uint) as _; + ptr::write(libc::CMSG_DATA(cmsg) as *mut libc::in_addr, src_addr); + + let ret = libc::sendmsg(socket.as_raw_fd(), &msg, 0); + if ret < 0 { + Err(io::Error::last_os_error()) + } else { + Ok(ret as usize) + } + } +} + +fn send_to_with_src_ipv4( + socket: &UdpSocket, + src_ip: Ipv4Addr, + dst_addr: SocketAddrV4, + buf: &[u8], +) -> io::Result { + send_to_with_src_ipv4_windows(socket, src_ip, SocketAddr::V4(dst_addr), buf) +} + +fn send_to_with_src_ipv4_windows( + socket: &UdpSocket, + src_ip: Ipv4Addr, + dst_addr: SocketAddr, + buf: &[u8], +) -> io::Result { + use std::{mem, os::windows::io::AsRawSocket, ptr}; + + use windows::{ + Win32::Networking::WinSock::{ + CMSGHDR, IN_ADDR, IN_ADDR_0, IN_PKTINFO, IP_PKTINFO, IPPROTO_IP, SOCKET, SOCKET_ERROR, + WSABUF, WSAGetLastError, WSAMSG, WSASendMsg, + }, + core::PSTR, + }; + + #[repr(align(8))] + struct ControlBuffer([u8; 128]); + + let dst = socket2::SockAddr::from(dst_addr); + let mut data = WSABUF { + len: buf.len() as u32, + buf: PSTR(buf.as_ptr() as *mut u8), + }; + let control_len = windows_cmsg_space(mem::size_of::()); + let mut control = ControlBuffer([0u8; 128]); + if control_len > control.0.len() { + return Err(io::Error::new( + io::ErrorKind::InvalidInput, + "IPv4 packet info control buffer is too small", + )); + } + let msg = WSAMSG { + name: dst.as_ptr() as *mut _, + namelen: dst.len(), + lpBuffers: &mut data, + dwBufferCount: 1, + Control: WSABUF { + len: control_len as u32, + buf: PSTR(control.0.as_mut_ptr()), + }, + dwFlags: 0, + }; + + let pktinfo = IN_PKTINFO { + ipi_addr: IN_ADDR { + S_un: IN_ADDR_0 { + S_addr: u32::from(src_ip).to_be(), + }, + }, + ipi_ifindex: 0, + }; + + unsafe { + let cmsg = control.0.as_mut_ptr() as *mut CMSGHDR; + (*cmsg).cmsg_level = IPPROTO_IP.0; + (*cmsg).cmsg_type = IP_PKTINFO; + (*cmsg).cmsg_len = windows_cmsg_len(mem::size_of::()); + ptr::write(windows_cmsg_data(cmsg) as *mut IN_PKTINFO, pktinfo); + + let mut sent = 0; + let ret = WSASendMsg( + SOCKET(socket.as_raw_socket() as usize), + &msg, + 0, + Some(&mut sent), + None, + None, + ); + if ret == SOCKET_ERROR { + return Err(io::Error::from_raw_os_error(WSAGetLastError().0)); + } + Ok(sent as usize) + } +} + +#[cfg(not(any( + target_os = "linux", + target_os = "android", + target_os = "freebsd", + target_os = "openbsd", + target_os = "netbsd", + target_os = "macos", + target_os = "ios", + windows +)))] +fn send_to_with_src_ipv4( + socket: &UdpSocket, + _src_ip: Ipv4Addr, + dst_addr: SocketAddrV4, + buf: &[u8], +) -> io::Result { + socket.try_send_to(buf, SocketAddr::V4(dst_addr)) +} + +pub(crate) fn send_to_with_src_ipv6( + socket: &UdpSocket, + src_ip: Ipv6Addr, + src_ifindex: u32, + dst_addr: SocketAddrV6, + buf: &[u8], +) -> io::Result { + use std::{mem, os::windows::io::AsRawSocket, ptr}; + + use windows::{ + Win32::Networking::WinSock::{ + CMSGHDR, IN6_ADDR, IN6_ADDR_0, IN6_PKTINFO, IPPROTO_IPV6, IPV6_PKTINFO, SOCKET, + SOCKET_ERROR, WSABUF, WSAGetLastError, WSAMSG, WSASendMsg, + }, + core::PSTR, + }; + + #[repr(align(8))] + struct ControlBuffer([u8; 128]); + + let dst = socket2::SockAddr::from(std::net::SocketAddr::V6(dst_addr)); + let mut data = WSABUF { + len: buf.len() as u32, + buf: PSTR(buf.as_ptr() as *mut u8), + }; + let control_len = windows_cmsg_space(mem::size_of::()); + let mut control = ControlBuffer([0u8; 128]); + if control_len > control.0.len() { + return Err(io::Error::new( + io::ErrorKind::InvalidInput, + "IPv6 packet info control buffer is too small", + )); + } + let msg = WSAMSG { + name: dst.as_ptr() as *mut _, + namelen: dst.len(), + lpBuffers: &mut data, + dwBufferCount: 1, + Control: WSABUF { + len: control_len as u32, + buf: PSTR(control.0.as_mut_ptr()), + }, + dwFlags: 0, + }; + + let pktinfo = IN6_PKTINFO { + ipi6_addr: IN6_ADDR { + u: IN6_ADDR_0 { + Byte: src_ip.octets(), + }, + }, + ipi6_ifindex: src_ifindex, + }; + + unsafe { + let cmsg = control.0.as_mut_ptr() as *mut CMSGHDR; + (*cmsg).cmsg_level = IPPROTO_IPV6.0; + (*cmsg).cmsg_type = IPV6_PKTINFO; + (*cmsg).cmsg_len = windows_cmsg_len(mem::size_of::()); + ptr::write(windows_cmsg_data(cmsg) as *mut IN6_PKTINFO, pktinfo); + let mut sent = 0; + let ret = WSASendMsg( + SOCKET(socket.as_raw_socket() as usize), + &msg, + 0, + Some(&mut sent), + None, + None, + ); + if ret == SOCKET_ERROR { + return Err(io::Error::from_raw_os_error(WSAGetLastError().0)); + } + Ok(sent as usize) + } +} + +#[cfg(not(any(unix, windows)))] +pub(crate) fn send_to_with_src_ipv6( + _socket: &UdpSocket, + _src_ip: Ipv6Addr, + _src_ifindex: u32, + _dst_addr: SocketAddrV6, + _buf: &[u8], +) -> io::Result { + Err(io::Error::new( + io::ErrorKind::Unsupported, + "sending UDP with a selected IPv6 source is not supported on this platform", + )) +} diff --git a/easytier/src/tests/credential_tests.rs b/easytier/src/tests/credential_tests.rs index 0907a322..f373e216 100644 --- a/easytier/src/tests/credential_tests.rs +++ b/easytier/src/tests/credential_tests.rs @@ -5,19 +5,24 @@ //! - Credential nodes use X25519 keypairs to authenticate without network_secret //! - Credentials can be revoked and propagate across the network -use std::time::Duration; +use std::{sync::Arc, time::Duration}; + +use easytier_core::peers::credential_manager::CredentialCreateOptions; +use easytier_core::process_runtime::CoreProcessRuntime; use crate::{ common::{ config::{ConfigLoader, NetworkIdentity, TomlConfigLoader}, global_ctx::GlobalCtxEvent, }, - instance::instance::Instance, + instance::test_instance::TestInstance as Instance, tests::three_node::{generate_secure_mode_config, generate_secure_mode_config_with_key}, - tunnel::{common::tests::wait_for_condition, tcp::TcpTunnelConnector, udp::UdpTunnelConnector}, + tunnel::common::tests::wait_for_condition, }; -use super::{add_ns_to_bridge, create_netns, del_netns, drop_insts, ping_test}; +use super::{ + InstanceTestExt as _, add_ns_to_bridge, create_netns, del_netns, drop_insts, ping_test, +}; use rstest::rstest; @@ -25,6 +30,56 @@ const PUBLIC_SERVER_NETWORK_NAME: &str = "__public_server__"; const PUBLIC_SERVER_SHARED_SECRET: &str = "public-server-shared-secret"; const NEED_P2P_ADMIN_NETWORK_NAME: &str = "need_p2p_credential_test_network"; +fn generate_credential( + admin: &Instance, + groups: Vec, + allow_relay: bool, + allowed_proxy_cidrs: Vec, + ttl: Duration, +) -> (String, String) { + generate_credential_with_options( + admin, + groups, + allow_relay, + allowed_proxy_cidrs, + ttl, + None, + true, + ) +} + +fn generate_credential_with_options( + admin: &Instance, + groups: Vec, + allow_relay: bool, + allowed_proxy_cidrs: Vec, + ttl: Duration, + credential_id: Option, + reusable: bool, +) -> (String, String) { + let generated = admin + .get_core_instance() + .generate_credential(CredentialCreateOptions { + groups, + allow_relay, + allowed_proxy_cidrs, + ttl, + credential_id, + reusable, + }) + .unwrap(); + (generated.credential_id, generated.secret) +} + +async fn set_avoid_relay_data(inst: &Instance, avoid_relay_data: bool) { + let mut config = crate::instance::config::test_runtime_instance_config(&inst.get_global_ctx()); + Arc::make_mut(&mut config.peer).avoid_relay_data_preference = avoid_relay_data; + inst.get_core_instance() + .update_runtime_config(config) + .await + .unwrap(); +} + /// Prepare network namespaces for credential tests /// Topology: /// br_a (10.1.1.0/24): ns_adm (10.1.1.1), ns_c1 (10.1.1.2), ns_c2 (10.1.1.3), ns_c3 (10.1.1.4), ns_c4 (10.1.1.5) @@ -128,10 +183,8 @@ async fn create_credential_config( ipv4: &str, ipv6: &str, ) -> TomlConfigLoader { - let (_cred_id, cred_secret) = admin_inst - .get_global_ctx() - .get_credential_manager() - .generate_credential(vec![], false, vec![], Duration::from_secs(3600)); + let (_cred_id, cred_secret) = + generate_credential(admin_inst, vec![], false, vec![], Duration::from_secs(3600)); build_credential_config( admin_inst @@ -318,7 +371,7 @@ fn create_public_server_credential_config( async fn wait_direct_peer(inst: &Instance, peer_id: u32, timeout: Duration, label: &str) { wait_for_condition( || async { - let peers = inst.get_peer_manager().get_peer_map().list_peers(); + let peers = inst.get_core_instance().connected_peers().await; let connected = peers.contains(&peer_id); println!("{label}: direct peers={:?}, target={}", peers, peer_id); connected @@ -331,7 +384,7 @@ async fn wait_direct_peer(inst: &Instance, peer_id: u32, timeout: Duration, labe async fn wait_running_listener(inst: &Instance, scheme: &str, timeout: Duration, label: &str) { wait_for_condition( || async { - let listeners = inst.get_global_ctx().get_running_listeners(); + let listeners = inst.get_core_instance().running_listeners(); let matched = listeners.iter().any(|listener| { listener.scheme() == scheme && listener.port().is_some_and(|p| p != 0) }); @@ -346,7 +399,7 @@ async fn wait_running_listener(inst: &Instance, scheme: &str, timeout: Duration, async fn wait_route_cost(inst: &Instance, peer_id: u32, cost: i32, timeout: Duration, label: &str) { wait_for_condition( || async { - let routes = inst.get_peer_manager().list_routes().await; + let routes = inst.get_core_instance().route_snapshots().await; let matched = routes .iter() .any(|route| route.peer_id == peer_id && route.cost == cost); @@ -370,11 +423,9 @@ async fn wait_foreign_network_count(inst: &Instance, expected: usize, timeout: D wait_for_condition( || async { let foreign_networks = inst - .get_peer_manager() - .get_foreign_network_manager() - .list_foreign_networks() - .await - .foreign_networks; + .get_core_instance() + .foreign_network_snapshots(false) + .await; println!("foreign networks: {:?}", foreign_networks); foreign_networks.len() == expected }, @@ -400,11 +451,16 @@ async fn credential_peers_p2p_to_need_p2p_admin_through_public_server( #[case] admin_listener_scheme: &str, ) { prepare_credential_network(); + let process_runtime = CoreProcessRuntime::new(); - let mut public_server_inst = Instance::new(create_public_server_config()); + let mut public_server_inst = + Instance::new_with_process_runtime(create_public_server_config(), process_runtime.clone()); public_server_inst.run().await.unwrap(); - let mut admin_inst = Instance::new(create_need_p2p_admin_config(admin_listener_scheme)); + let mut admin_inst = Instance::new_with_process_runtime( + create_need_p2p_admin_config(admin_listener_scheme), + process_runtime.clone(), + ); admin_inst.run().await.unwrap(); wait_running_listener( &admin_inst, @@ -413,77 +469,67 @@ async fn credential_peers_p2p_to_need_p2p_admin_through_public_server( "admin ephemeral listener", ) .await; - admin_inst - .get_conn_manager() - .add_connector(UdpTunnelConnector::new( - "udp://10.1.1.1:11010".parse().unwrap(), - )); + admin_inst.add_connector_url("udp://10.1.1.1:11010".parse().unwrap()); wait_foreign_network_count(&public_server_inst, 1, Duration::from_secs(10)).await; - let (_credential_a_id, credential_a_secret) = admin_inst - .get_global_ctx() - .get_credential_manager() - .generate_credential_with_options( - vec![], - false, - vec!["10.1.0.0/24".to_string()], - Duration::from_secs(3600), - Some("credential-peer-a".to_string()), - false, - ); - let (_credential_b_id, credential_b_secret) = admin_inst - .get_global_ctx() - .get_credential_manager() - .generate_credential_with_options( - vec![], - false, - vec![], - Duration::from_secs(3600), - Some("credential-peer-b".to_string()), - false, - ); + let (_credential_a_id, credential_a_secret) = generate_credential_with_options( + &admin_inst, + vec![], + false, + vec!["10.1.0.0/24".to_string()], + Duration::from_secs(3600), + Some("credential-peer-a".to_string()), + false, + ); + let (_credential_b_id, credential_b_secret) = generate_credential_with_options( + &admin_inst, + vec![], + false, + vec![], + Duration::from_secs(3600), + Some("credential-peer-b".to_string()), + false, + ); admin_inst .get_global_ctx() .issue_event(GlobalCtxEvent::CredentialChanged); wait_foreign_network_count(&public_server_inst, 1, Duration::from_secs(10)).await; - let mut credential_a_inst = Instance::new(create_public_server_credential_config( - &credential_a_secret, - "credential-peer-a", - "credential-a", - "ns_c1", - "10.154.0.1", - "fd00::1/64", - 11030, - 11031, - &["10.1.0.0/24"], - )); - let mut credential_b_inst = Instance::new(create_public_server_credential_config( - &credential_b_secret, - "credential-peer-b", - "credential-b", - "ns_c2", - "10.154.0.2", - "fd00::2/64", - 11040, - 11041, - &[], - )); + let mut credential_a_inst = Instance::new_with_process_runtime( + create_public_server_credential_config( + &credential_a_secret, + "credential-peer-a", + "credential-a", + "ns_c1", + "10.154.0.1", + "fd00::1/64", + 11030, + 11031, + &["10.1.0.0/24"], + ), + process_runtime.clone(), + ); + let mut credential_b_inst = Instance::new_with_process_runtime( + create_public_server_credential_config( + &credential_b_secret, + "credential-peer-b", + "credential-b", + "ns_c2", + "10.154.0.2", + "fd00::2/64", + 11040, + 11041, + &[], + ), + process_runtime.clone(), + ); credential_a_inst.run().await.unwrap(); credential_b_inst.run().await.unwrap(); - credential_a_inst - .get_conn_manager() - .add_connector(UdpTunnelConnector::new( - "udp://10.1.1.1:11010".parse().unwrap(), - )); - credential_b_inst - .get_conn_manager() - .add_connector(UdpTunnelConnector::new( - "udp://10.1.1.1:11010".parse().unwrap(), - )); + credential_a_inst.add_connector_url("udp://10.1.1.1:11010".parse().unwrap()); + credential_b_inst.add_connector_url("udp://10.1.1.1:11010".parse().unwrap()); let admin_peer_id = admin_inst.peer_id(); let credential_a_peer_id = credential_a_inst.peer_id(); @@ -554,10 +600,8 @@ fn create_generated_credential_config( ipv4: &str, ipv6: &str, ) -> (TomlConfigLoader, String) { - let (cred_id, cred_secret) = admin_inst - .get_global_ctx() - .get_credential_manager() - .generate_credential(vec![], false, vec![], Duration::from_secs(3600)); + let (cred_id, cred_secret) = + generate_credential(admin_inst, vec![], false, vec![], Duration::from_secs(3600)); let config = build_credential_config( admin_inst .get_global_ctx() @@ -594,8 +638,8 @@ async fn wait_route_presence_on_admins( ) { wait_for_condition( || async { - let admin_a_routes = admin_a_inst.get_peer_manager().list_routes().await; - let admin_c_routes = admin_c_inst.get_peer_manager().list_routes().await; + let admin_a_routes = admin_a_inst.get_core_instance().route_snapshots().await; + let admin_c_routes = admin_c_inst.get_core_instance().route_snapshots().await; let admin_a_has = admin_a_routes.iter().any(|r| r.peer_id == peer_id); let admin_c_has = admin_c_routes.iter().any(|r| r.peer_id == peer_id); if should_exist { @@ -618,8 +662,8 @@ async fn assert_shared_visibility_stable( label: &str, ) { for _ in 0..5 { - let admin_a_routes = admin_a_inst.get_peer_manager().list_routes().await; - let admin_c_routes = admin_c_inst.get_peer_manager().list_routes().await; + let admin_a_routes = admin_a_inst.get_core_instance().route_snapshots().await; + let admin_c_routes = admin_c_inst.get_core_instance().route_snapshots().await; let admin_a_has = admin_a_routes.iter().any(|r| r.peer_id == peer_id); let admin_c_has = admin_c_routes.iter().any(|r| r.peer_id == peer_id); if should_exist { @@ -666,8 +710,8 @@ async fn wait_stable_single_visible_peer_on_admins( let mut stable_samples = 0; loop { - let admin_a_routes = admin_a_inst.get_peer_manager().list_routes().await; - let admin_c_routes = admin_c_inst.get_peer_manager().list_routes().await; + let admin_a_routes = admin_a_inst.get_core_instance().route_snapshots().await; + let admin_c_routes = admin_c_inst.get_core_instance().route_snapshots().await; let admin_a_has_a = admin_a_routes.iter().any(|r| r.peer_id == peer_a_id); let admin_a_has_b = admin_a_routes.iter().any(|r| r.peer_id == peer_b_id); @@ -723,10 +767,11 @@ async fn wait_stable_single_visible_peer_on_admins( #[serial_test::serial] async fn credential_basic_connectivity() { prepare_credential_network(); + let process_runtime = CoreProcessRuntime::new(); // Create admin node let admin_config = create_admin_config("admin", Some("ns_adm"), "10.144.144.1", "fd00::1/64"); - let mut admin_inst = Instance::new(admin_config); + let mut admin_inst = Instance::new_with_process_runtime(admin_config, process_runtime.clone()); admin_inst.run().await.unwrap(); // Create credential node @@ -738,15 +783,11 @@ async fn credential_basic_connectivity() { "fd00::2/64", ) .await; - let mut cred_inst = Instance::new(cred_config); + let mut cred_inst = Instance::new_with_process_runtime(cred_config, process_runtime.clone()); cred_inst.run().await.unwrap(); // Credential connects to admin - cred_inst - .get_conn_manager() - .add_connector(TcpTunnelConnector::new( - "tcp://10.1.1.1:11010".parse().unwrap(), - )); + cred_inst.add_connector_url("tcp://10.1.1.1:11010".parse().unwrap()); let cred_peer_id = cred_inst.peer_id(); let admin_peer_id = admin_inst.peer_id(); @@ -759,18 +800,18 @@ async fn credential_basic_connectivity() { tokio::time::sleep(Duration::from_secs(2)).await; // Check peers and connections - let admin_peers = admin_inst.get_peer_manager().get_peer_map().list_peers(); - let cred_peers = cred_inst.get_peer_manager().get_peer_map().list_peers(); + let admin_peers = admin_inst.get_core_instance().connected_peers().await; + let cred_peers = cred_inst.get_core_instance().connected_peers().await; println!("Admin peers: {:?}", admin_peers); println!("Credential peers: {:?}", cred_peers); // Wait for credential to appear in admin's route table wait_for_condition( || async { - let routes = admin_inst.get_peer_manager().list_routes().await; - let cred_routes = cred_inst.get_peer_manager().list_routes().await; - let admin_peers = admin_inst.get_peer_manager().get_peer_map().list_peers(); - let cred_peers = cred_inst.get_peer_manager().get_peer_map().list_peers(); + let routes = admin_inst.get_core_instance().route_snapshots().await; + let cred_routes = cred_inst.get_core_instance().route_snapshots().await; + let admin_peers = admin_inst.get_core_instance().connected_peers().await; + let cred_peers = cred_inst.get_core_instance().connected_peers().await; println!( "Admin peers: {:?}, routes: {:?}", admin_peers, @@ -820,37 +861,43 @@ async fn credential_basic_connectivity() { #[tokio::test] #[serial_test::serial] async fn credential_relay_capability(#[case] allow_relay: bool) { - use crate::peers::route_trait::NextHopPolicy; - prepare_credential_network(); + let process_runtime = CoreProcessRuntime::new(); // Create admin node let admin_config = create_admin_config("admin", Some("ns_adm"), "10.144.144.1", "fd00::1/64"); - let mut admin_inst = Instance::new(admin_config); + let mut admin_inst = Instance::new_with_process_runtime(admin_config, process_runtime.clone()); // if cred c allow relay, we set admin inst avoid relay (if other same-cost path available, admin will not relay data) - admin_inst - .get_global_ctx() - .set_avoid_relay_data_preference(allow_relay); + set_avoid_relay_data(&admin_inst, allow_relay).await; admin_inst.run().await.unwrap(); let admin_peer_id = admin_inst.peer_id(); // Generate credentials for A, B, C // C has configurable allow_relay - let (_cred_a_id, cred_a_secret) = admin_inst - .get_global_ctx() - .get_credential_manager() - .generate_credential(vec![], false, vec![], Duration::from_secs(3600)); + let (_cred_a_id, cred_a_secret) = generate_credential( + &admin_inst, + vec![], + false, + vec![], + Duration::from_secs(3600), + ); - let (_cred_b_id, cred_b_secret) = admin_inst - .get_global_ctx() - .get_credential_manager() - .generate_credential(vec![], false, vec![], Duration::from_secs(3600)); + let (_cred_b_id, cred_b_secret) = generate_credential( + &admin_inst, + vec![], + false, + vec![], + Duration::from_secs(3600), + ); - let (_cred_c_id, cred_c_secret) = admin_inst - .get_global_ctx() - .get_credential_manager() - .generate_credential(vec![], allow_relay, vec![], Duration::from_secs(3600)); + let (_cred_c_id, cred_c_secret) = generate_credential( + &admin_inst, + vec![], + allow_relay, + vec![], + Duration::from_secs(3600), + ); // Create credential A on ns_c1 let cred_a_config = { @@ -876,7 +923,8 @@ async fn credential_relay_capability(#[case] allow_relay: bool) { config.set_secure_mode(Some(generate_secure_mode_config_with_key(&private))); config }; - let mut cred_a_inst = Instance::new(cred_a_config); + let mut cred_a_inst = + Instance::new_with_process_runtime(cred_a_config, process_runtime.clone()); cred_a_inst.run().await.unwrap(); // Create credential B on ns_c2 @@ -903,7 +951,8 @@ async fn credential_relay_capability(#[case] allow_relay: bool) { config.set_secure_mode(Some(generate_secure_mode_config_with_key(&private))); config }; - let mut cred_b_inst = Instance::new(cred_b_config); + let mut cred_b_inst = + Instance::new_with_process_runtime(cred_b_config, process_runtime.clone()); cred_b_inst.run().await.unwrap(); // Create credential C on ns_c3 WITH listener (so A and B can connect to it) @@ -932,7 +981,8 @@ async fn credential_relay_capability(#[case] allow_relay: bool) { config.set_secure_mode(Some(generate_secure_mode_config_with_key(&private))); config }; - let mut cred_c_inst = Instance::new(cred_c_config); + let mut cred_c_inst = + Instance::new_with_process_runtime(cred_c_config, process_runtime.clone()); cred_c_inst.run().await.unwrap(); let cred_a_peer_id = cred_a_inst.peer_id(); @@ -940,34 +990,14 @@ async fn credential_relay_capability(#[case] allow_relay: bool) { let cred_c_peer_id = cred_c_inst.peer_id(); // All credentials connect to admin - cred_a_inst - .get_conn_manager() - .add_connector(TcpTunnelConnector::new( - "tcp://10.1.1.1:11010".parse().unwrap(), - )); - cred_b_inst - .get_conn_manager() - .add_connector(TcpTunnelConnector::new( - "tcp://10.1.1.1:11010".parse().unwrap(), - )); - cred_c_inst - .get_conn_manager() - .add_connector(TcpTunnelConnector::new( - "tcp://10.1.1.1:11010".parse().unwrap(), - )); + cred_a_inst.add_connector_url("tcp://10.1.1.1:11010".parse().unwrap()); + cred_b_inst.add_connector_url("tcp://10.1.1.1:11010".parse().unwrap()); + cred_c_inst.add_connector_url("tcp://10.1.1.1:11010".parse().unwrap()); // A and B also connect to C (simulating P2P discovery and connection) // C is on ns_c3 with IP 10.1.1.4, listener on port 11020 - cred_a_inst - .get_conn_manager() - .add_connector(TcpTunnelConnector::new( - "tcp://10.1.1.4:11020".parse().unwrap(), - )); - cred_b_inst - .get_conn_manager() - .add_connector(TcpTunnelConnector::new( - "tcp://10.1.1.4:11020".parse().unwrap(), - )); + cred_a_inst.add_connector_url("tcp://10.1.1.4:11020".parse().unwrap()); + cred_b_inst.add_connector_url("tcp://10.1.1.4:11020".parse().unwrap()); // print all peer ids println!("Admin peer id: {:?}", admin_peer_id); println!("Cred A peer id: {:?}", cred_a_peer_id); @@ -977,7 +1007,7 @@ async fn credential_relay_capability(#[case] allow_relay: bool) { // Wait for all nodes to appear in admin's route table wait_for_condition( || async { - let routes = admin_inst.get_peer_manager().list_routes().await; + let routes = admin_inst.get_core_instance().route_snapshots().await; let has_a = routes.iter().any(|r| r.peer_id == cred_a_peer_id); let has_b = routes.iter().any(|r| r.peer_id == cred_b_peer_id); let has_c = routes.iter().any(|r| r.peer_id == cred_c_peer_id); @@ -991,9 +1021,9 @@ async fn credential_relay_capability(#[case] allow_relay: bool) { // Wait for P2P connections to establish wait_for_condition( || async { - let peers_a = cred_a_inst.get_peer_manager().get_peer_map().list_peers(); - let peers_b = cred_b_inst.get_peer_manager().get_peer_map().list_peers(); - let peers_c = cred_c_inst.get_peer_manager().get_peer_map().list_peers(); + let peers_a = cred_a_inst.get_core_instance().connected_peers().await; + let peers_b = cred_b_inst.get_core_instance().connected_peers().await; + let peers_c = cred_c_inst.get_core_instance().connected_peers().await; let a_connected_c = peers_a.contains(&cred_c_peer_id); let b_connected_c = peers_b.contains(&cred_c_peer_id); @@ -1018,7 +1048,7 @@ async fn credential_relay_capability(#[case] allow_relay: bool) { // Wait for routes to propagate wait_for_condition( || async { - let routes_a = cred_a_inst.get_peer_manager().list_routes().await; + let routes_a = cred_a_inst.get_core_instance().route_snapshots().await; let a_sees_b = routes_a.iter().any(|r| r.peer_id == cred_b_peer_id); let cost_a_to_b = routes_a .iter() @@ -1035,10 +1065,12 @@ async fn credential_relay_capability(#[case] allow_relay: bool) { wait_for_condition( || async { let next_hop_a_to_b = cred_a_inst - .get_peer_manager() - .get_route() - .get_next_hop_with_policy(cred_b_peer_id, NextHopPolicy::LeastCost) - .await; + .get_core_instance() + .route_snapshots() + .await + .into_iter() + .find(|route| route.peer_id == cred_b_peer_id) + .and_then(|route| route.next_hop_peer_id_latency_first); println!( "Next hop convergence A->B={:?} (admin={}, c={}), allow_relay={}", next_hop_a_to_b, admin_peer_id, cred_c_peer_id, allow_relay @@ -1058,10 +1090,12 @@ async fn credential_relay_capability(#[case] allow_relay: bool) { // Verify next hop from A to B based on allow_relay flag let next_hop_a_to_b = cred_a_inst - .get_peer_manager() - .get_route() - .get_next_hop_with_policy(cred_b_peer_id, NextHopPolicy::LeastCost) - .await; + .get_core_instance() + .route_snapshots() + .await + .into_iter() + .find(|route| route.peer_id == cred_b_peer_id) + .and_then(|route| route.next_hop_peer_id_latency_first); println!( "Next hop A->B={:?} (admin={}, c={}), allow_relay={}", @@ -1095,10 +1129,11 @@ async fn credential_relay_capability(#[case] allow_relay: bool) { #[serial_test::serial] async fn credential_two_credentials_communicate_tcp() { prepare_credential_network(); + let process_runtime = CoreProcessRuntime::new(); // Create admin node let admin_config = create_admin_config("admin", Some("ns_adm"), "10.144.144.1", "fd00::1/64"); - let mut admin_inst = Instance::new(admin_config); + let mut admin_inst = Instance::new_with_process_runtime(admin_config, process_runtime.clone()); admin_inst.run().await.unwrap(); // Create credential1 on ns_c1 @@ -1110,7 +1145,7 @@ async fn credential_two_credentials_communicate_tcp() { "fd00::2/64", ) .await; - let mut cred1_inst = Instance::new(cred1_config); + let mut cred1_inst = Instance::new_with_process_runtime(cred1_config, process_runtime.clone()); cred1_inst.run().await.unwrap(); // Create credential2 on ns_c2 @@ -1122,20 +1157,12 @@ async fn credential_two_credentials_communicate_tcp() { "fd00::3/64", ) .await; - let mut cred2_inst = Instance::new(cred2_config); + let mut cred2_inst = Instance::new_with_process_runtime(cred2_config, process_runtime.clone()); cred2_inst.run().await.unwrap(); // Both credentials connect to admin - cred1_inst - .get_conn_manager() - .add_connector(TcpTunnelConnector::new( - "tcp://10.1.1.1:11010".parse().unwrap(), - )); - cred2_inst - .get_conn_manager() - .add_connector(TcpTunnelConnector::new( - "tcp://10.1.1.1:11010".parse().unwrap(), - )); + cred1_inst.add_connector_url("tcp://10.1.1.1:11010".parse().unwrap()); + cred2_inst.add_connector_url("tcp://10.1.1.1:11010".parse().unwrap()); let cred1_peer_id = cred1_inst.peer_id(); let cred2_peer_id = cred2_inst.peer_id(); @@ -1143,7 +1170,7 @@ async fn credential_two_credentials_communicate_tcp() { // Wait for both credentials to appear in admin's route table wait_for_condition( || async { - let routes = admin_inst.get_peer_manager().list_routes().await; + let routes = admin_inst.get_core_instance().route_snapshots().await; routes.iter().any(|r| r.peer_id == cred1_peer_id) && routes.iter().any(|r| r.peer_id == cred2_peer_id) }, @@ -1174,17 +1201,21 @@ async fn credential_two_credentials_communicate_tcp() { #[serial_test::serial] async fn credential_revocation_propagates() { prepare_credential_network(); + let process_runtime = CoreProcessRuntime::new(); // Create admin on ns_adm (10.1.1.1) let admin_config = create_admin_config("admin", Some("ns_adm"), "10.144.144.1", "fd00::1/64"); - let mut admin_inst = Instance::new(admin_config); + let mut admin_inst = Instance::new_with_process_runtime(admin_config, process_runtime.clone()); admin_inst.run().await.unwrap(); // Generate credential on admin - let (cred_id, cred_secret) = admin_inst - .get_global_ctx() - .get_credential_manager() - .generate_credential(vec![], false, vec![], Duration::from_secs(3600)); + let (cred_id, cred_secret) = generate_credential( + &admin_inst, + vec![], + false, + vec![], + Duration::from_secs(3600), + ); // Create credential node let cred_config = { @@ -1213,15 +1244,11 @@ async fn credential_revocation_propagates() { config }; - let mut cred_inst = Instance::new(cred_config); + let mut cred_inst = Instance::new_with_process_runtime(cred_config, process_runtime.clone()); cred_inst.run().await.unwrap(); // Credential connects to admin - cred_inst - .get_conn_manager() - .add_connector(TcpTunnelConnector::new( - "tcp://10.1.1.1:11010".parse().unwrap(), - )); + cred_inst.add_connector_url("tcp://10.1.1.1:11010".parse().unwrap()); let cred_peer_id = cred_inst.peer_id(); @@ -1229,8 +1256,8 @@ async fn credential_revocation_propagates() { wait_for_condition( || async { admin_inst - .get_peer_manager() - .list_routes() + .get_core_instance() + .route_snapshots() .await .iter() .any(|r| r.peer_id == cred_peer_id) @@ -1249,9 +1276,9 @@ async fn credential_revocation_propagates() { // Revoke the credential assert!( admin_inst - .get_global_ctx() - .get_credential_manager() - .revoke_credential(&cred_id), + .get_core_instance() + .revoke_credential(&cred_id) + .unwrap(), "Credential should be revoked successfully" ); @@ -1264,8 +1291,8 @@ async fn credential_revocation_propagates() { wait_for_condition( || async { !admin_inst - .get_peer_manager() - .list_routes() + .get_core_instance() + .route_snapshots() .await .iter() .any(|r| r.peer_id == cred_peer_id) @@ -1294,22 +1321,21 @@ async fn credential_revocation_propagates() { #[serial_test::serial] async fn credential_non_reusable_allows_only_one_peer() { prepare_credential_network(); + let process_runtime = CoreProcessRuntime::new(); let admin_config = create_admin_config("admin", Some("ns_adm"), "10.144.144.1", "fd00::1/64"); - let mut admin_inst = Instance::new(admin_config); + let mut admin_inst = Instance::new_with_process_runtime(admin_config, process_runtime.clone()); admin_inst.run().await.unwrap(); - let (_cred_id, cred_secret) = admin_inst - .get_global_ctx() - .get_credential_manager() - .generate_credential_with_options( - vec![], - false, - vec![], - Duration::from_secs(3600), - None, - false, - ); + let (_cred_id, cred_secret) = generate_credential_with_options( + &admin_inst, + vec![], + false, + vec![], + Duration::from_secs(3600), + None, + false, + ); let network_name = admin_inst .get_global_ctx() @@ -1333,22 +1359,22 @@ async fn credential_non_reusable_allows_only_one_peer() { "fd00::3/64", ); - let mut cred1_inst = Some(Instance::new(cred1_config)); + let mut cred1_inst = Some(Instance::new_with_process_runtime( + cred1_config, + process_runtime.clone(), + )); cred1_inst.as_mut().unwrap().run().await.unwrap(); cred1_inst .as_ref() .unwrap() - .get_conn_manager() - .add_connector(TcpTunnelConnector::new( - "tcp://10.1.1.1:11010".parse().unwrap(), - )); + .add_connector_url("tcp://10.1.1.1:11010".parse().unwrap()); let cred1_peer_id = cred1_inst.as_ref().unwrap().peer_id(); wait_for_condition( || async { admin_inst - .get_peer_manager() - .list_routes() + .get_core_instance() + .route_snapshots() .await .iter() .any(|r| r.peer_id == cred1_peer_id) @@ -1358,22 +1384,22 @@ async fn credential_non_reusable_allows_only_one_peer() { .await; wait_ping_reachability("ns_adm", "10.144.144.2", true, Duration::from_secs(10)).await; - let mut cred2_inst = Some(Instance::new(cred2_config)); + let mut cred2_inst = Some(Instance::new_with_process_runtime( + cred2_config, + process_runtime.clone(), + )); cred2_inst.as_mut().unwrap().run().await.unwrap(); cred2_inst .as_ref() .unwrap() - .get_conn_manager() - .add_connector(TcpTunnelConnector::new( - "tcp://10.1.1.1:11010".parse().unwrap(), - )); + .add_connector_url("tcp://10.1.1.1:11010".parse().unwrap()); let cred2_peer_id = cred2_inst.as_ref().unwrap().peer_id(); tokio::time::sleep(Duration::from_secs(3)).await; // The non-reusable credential owner is elected by lowest peer_id, so either cred1 or cred2 // may win. Determine the winner and loser dynamically. - let admin_routes = admin_inst.get_peer_manager().list_routes().await; + let admin_routes = admin_inst.get_core_instance().route_snapshots().await; let (winner_peer_id, winner_ip, winner_inst, loser_peer_id, loser_ip, loser_inst) = if admin_routes.iter().any(|r| r.peer_id == cred1_peer_id) { ( @@ -1401,7 +1427,7 @@ async fn credential_non_reusable_allows_only_one_peer() { }; for _ in 0..5 { - let admin_routes = admin_inst.get_peer_manager().list_routes().await; + let admin_routes = admin_inst.get_core_instance().route_snapshots().await; assert!( admin_routes.iter().any(|r| r.peer_id == winner_peer_id), "winning credential peer should remain present: {:?}", @@ -1415,7 +1441,7 @@ async fn credential_non_reusable_allows_only_one_peer() { } for _ in 0..5 { - let admin_routes = admin_inst.get_peer_manager().list_routes().await; + let admin_routes = admin_inst.get_core_instance().route_snapshots().await; assert!( !admin_routes.iter().any(|r| r.peer_id == loser_peer_id), "losing credential peer should not appear in routes: {:?}", @@ -1432,7 +1458,7 @@ async fn credential_non_reusable_allows_only_one_peer() { wait_for_condition( || async { - let routes = admin_inst.get_peer_manager().list_routes().await; + let routes = admin_inst.get_core_instance().route_snapshots().await; !routes.iter().any(|r| r.peer_id == winner_peer_id) && routes.iter().any(|r| r.peer_id == loser_peer_id) }, @@ -1451,10 +1477,11 @@ async fn credential_non_reusable_allows_only_one_peer() { #[serial_test::serial] async fn credential_unknown_rejected() { prepare_credential_network(); + let process_runtime = CoreProcessRuntime::new(); // Create admin node let admin_config = create_admin_config("admin", Some("ns_adm"), "10.144.144.1", "fd00::1/64"); - let mut admin_inst = Instance::new(admin_config); + let mut admin_inst = Instance::new_with_process_runtime(admin_config, process_runtime.clone()); admin_inst.run().await.unwrap(); // Create credential node with random key (not generated by admin) @@ -1477,15 +1504,11 @@ async fn credential_unknown_rejected() { config }; - let mut cred_inst = Instance::new(cred_config); + let mut cred_inst = Instance::new_with_process_runtime(cred_config, process_runtime.clone()); cred_inst.run().await.unwrap(); // Attempt to connect to admin - cred_inst - .get_conn_manager() - .add_connector(TcpTunnelConnector::new( - "tcp://10.1.1.1:11010".parse().unwrap(), - )); + cred_inst.add_connector_url("tcp://10.1.1.1:11010".parse().unwrap()); let cred_peer_id = cred_inst.peer_id(); @@ -1493,7 +1516,7 @@ async fn credential_unknown_rejected() { tokio::time::sleep(Duration::from_secs(5)).await; // Verify credential does NOT appear in admin's route table - let routes = admin_inst.get_peer_manager().list_routes().await; + let routes = admin_inst.get_core_instance().route_snapshots().await; assert!( !routes.iter().any(|r| r.peer_id == cred_peer_id), "Unknown credential node should NOT appear in admin's route table" @@ -1517,38 +1540,34 @@ async fn credential_unknown_rejected() { #[serial_test::serial] async fn credential_unknown_via_shared_rejected(#[values(true, false)] test_revoke: bool) { prepare_credential_network(); + let process_runtime = CoreProcessRuntime::new(); let admin_a_config = create_admin_config("admin_a", Some("ns_adm"), "10.144.144.1", "fd00::1/64"); - let mut admin_a_inst = Instance::new(admin_a_config); + let mut admin_a_inst = + Instance::new_with_process_runtime(admin_a_config, process_runtime.clone()); admin_a_inst.run().await.unwrap(); let shared_b_config = create_shared_config("shared_b", Some("ns_c1"), "10.144.144.2", "fd00::2/64"); - let mut shared_b_inst = Instance::new(shared_b_config); + let mut shared_b_inst = + Instance::new_with_process_runtime(shared_b_config, process_runtime.clone()); shared_b_inst.run().await.unwrap(); let admin_c_config = create_admin_config("admin_c", Some("ns_c3"), "10.144.144.4", "fd00::4/64"); - let mut admin_c_inst = Instance::new(admin_c_config); + let mut admin_c_inst = + Instance::new_with_process_runtime(admin_c_config, process_runtime.clone()); admin_c_inst.run().await.unwrap(); - admin_a_inst - .get_conn_manager() - .add_connector(TcpTunnelConnector::new( - "tcp://10.1.1.2:11010".parse().unwrap(), - )); - admin_c_inst - .get_conn_manager() - .add_connector(TcpTunnelConnector::new( - "tcp://10.1.1.2:11010".parse().unwrap(), - )); + admin_a_inst.add_connector_url("tcp://10.1.1.2:11010".parse().unwrap()); + admin_c_inst.add_connector_url("tcp://10.1.1.2:11010".parse().unwrap()); let admin_c_peer_id = admin_c_inst.peer_id(); wait_for_condition( || async { - let a_routes = admin_a_inst.get_peer_manager().list_routes().await; - let c_routes = admin_c_inst.get_peer_manager().list_routes().await; + let a_routes = admin_a_inst.get_core_instance().route_snapshots().await; + let c_routes = admin_c_inst.get_core_instance().route_snapshots().await; a_routes.iter().any(|r| r.peer_id == admin_c_peer_id) || c_routes.iter().any(|r| r.peer_id == admin_a_inst.peer_id()) }, @@ -1581,14 +1600,11 @@ async fn credential_unknown_via_shared_rejected(#[values(true, false)] test_revo None, ) }; - let mut unknown_inst = Instance::new(credential_config); + let mut unknown_inst = + Instance::new_with_process_runtime(credential_config, process_runtime.clone()); unknown_inst.run().await.unwrap(); - unknown_inst - .get_conn_manager() - .add_connector(TcpTunnelConnector::new( - "tcp://10.1.1.2:11010".parse().unwrap(), - )); + unknown_inst.add_connector_url("tcp://10.1.1.2:11010".parse().unwrap()); let unknown_peer_id = unknown_inst.peer_id(); @@ -1608,9 +1624,9 @@ async fn credential_unknown_via_shared_rejected(#[values(true, false)] test_revo assert!( admin_a_inst - .get_global_ctx() - .get_credential_manager() - .revoke_credential(credential_id.as_ref().unwrap()), + .get_core_instance() + .revoke_credential(credential_id.as_ref().unwrap()) + .unwrap(), "credential should be revoked successfully" ); admin_a_inst @@ -1628,11 +1644,7 @@ async fn credential_unknown_via_shared_rejected(#[values(true, false)] test_revo wait_ping_reachability("ns_adm", "10.144.144.5", false, Duration::from_secs(5)).await; wait_ping_reachability("ns_c3", "10.144.144.5", false, Duration::from_secs(5)).await; - unknown_inst - .get_conn_manager() - .add_connector(TcpTunnelConnector::new( - "tcp://10.1.1.2:11010".parse().unwrap(), - )); + unknown_inst.add_connector_url("tcp://10.1.1.2:11010".parse().unwrap()); assert_shared_visibility_stable( &admin_a_inst, @@ -1673,35 +1685,31 @@ async fn credential_admin_shared_admin_credential_connectivity( #[values(true, false)] connect_to_admin: bool, ) { prepare_credential_network(); + let process_runtime = CoreProcessRuntime::new(); // 10.1.1.1 let admin_a_config = create_admin_config("admin_a", Some("ns_adm"), "10.144.144.1", "fd00::1/64"); - let mut admin_a_inst = Instance::new(admin_a_config); + let mut admin_a_inst = + Instance::new_with_process_runtime(admin_a_config, process_runtime.clone()); admin_a_inst.run().await.unwrap(); // 10.1.1.2 let shared_b_config = create_shared_config("shared_b", Some("ns_c1"), "10.144.144.2", "fd00::2/64"); - let mut shared_b_inst = Instance::new(shared_b_config); + let mut shared_b_inst = + Instance::new_with_process_runtime(shared_b_config, process_runtime.clone()); shared_b_inst.run().await.unwrap(); // 10.1.1.4 let admin_c_config = create_admin_config("admin_c", Some("ns_c3"), "10.144.144.4", "fd00::4/64"); - let mut admin_c_inst = Instance::new(admin_c_config); + let mut admin_c_inst = + Instance::new_with_process_runtime(admin_c_config, process_runtime.clone()); admin_c_inst.run().await.unwrap(); - admin_a_inst - .get_conn_manager() - .add_connector(TcpTunnelConnector::new( - "tcp://10.1.1.2:11010".parse().unwrap(), - )); - admin_c_inst - .get_conn_manager() - .add_connector(TcpTunnelConnector::new( - "tcp://10.1.1.2:11010".parse().unwrap(), - )); + admin_a_inst.add_connector_url("tcp://10.1.1.2:11010".parse().unwrap()); + admin_c_inst.add_connector_url("tcp://10.1.1.2:11010".parse().unwrap()); // print all peer ids println!("admin_a_peer_id: {:?}", admin_a_inst.peer_id()); @@ -1711,8 +1719,8 @@ async fn credential_admin_shared_admin_credential_connectivity( let admin_c_peer_id = admin_c_inst.peer_id(); wait_for_condition( || async { - let a_routes = admin_a_inst.get_peer_manager().list_routes().await; - let c_routes = admin_c_inst.get_peer_manager().list_routes().await; + let a_routes = admin_a_inst.get_core_instance().route_snapshots().await; + let c_routes = admin_c_inst.get_core_instance().route_snapshots().await; println!( "bootstrap routes: a={:?} c={:?}", a_routes.iter().map(|r| r.peer_id).collect::>(), @@ -1737,27 +1745,26 @@ async fn credential_admin_shared_admin_credential_connectivity( .get_global_ctx() .issue_event(GlobalCtxEvent::CredentialChanged); - let mut cred_d_inst = Instance::new(cred_d_config); + let mut cred_d_inst = + Instance::new_with_process_runtime(cred_d_config, process_runtime.clone()); cred_d_inst.run().await.unwrap(); let cred_d_peer_id = cred_d_inst.peer_id(); - cred_d_inst - .get_conn_manager() - .add_connector(TcpTunnelConnector::new(if !connect_to_admin { - // connect to shared node - "tcp://10.1.1.2:11010".parse().unwrap() - } else { - // connect to admin node - "tcp://10.1.1.4:11010".parse().unwrap() - })); + cred_d_inst.add_connector_url(if !connect_to_admin { + // connect to shared node + "tcp://10.1.1.2:11010".parse().unwrap() + } else { + // connect to admin node + "tcp://10.1.1.4:11010".parse().unwrap() + }); // print all peer ids println!("cred_d_peer_id: {:?}", cred_d_peer_id); wait_for_condition( || async { admin_c_inst - .get_peer_manager() - .list_routes() + .get_core_instance() + .route_snapshots() .await .iter() .any(|r| r.peer_id == cred_d_peer_id) @@ -1791,38 +1798,34 @@ async fn credential_admin_shared_admin_credential_connectivity( #[serial_test::serial] async fn credential_non_reusable_across_two_admins_allows_only_one_peer() { prepare_credential_network(); + let process_runtime = CoreProcessRuntime::new(); let admin_a_config = create_admin_config("admin_a", Some("ns_adm"), "10.144.144.1", "fd00::1/64"); - let mut admin_a_inst = Instance::new(admin_a_config); + let mut admin_a_inst = + Instance::new_with_process_runtime(admin_a_config, process_runtime.clone()); admin_a_inst.run().await.unwrap(); let shared_b_config = create_shared_config("shared_b", Some("ns_c1"), "10.144.144.2", "fd00::2/64"); - let mut shared_b_inst = Instance::new(shared_b_config); + let mut shared_b_inst = + Instance::new_with_process_runtime(shared_b_config, process_runtime.clone()); shared_b_inst.run().await.unwrap(); let admin_c_config = create_admin_config("admin_c", Some("ns_c3"), "10.144.144.4", "fd00::4/64"); - let mut admin_c_inst = Instance::new(admin_c_config); + let mut admin_c_inst = + Instance::new_with_process_runtime(admin_c_config, process_runtime.clone()); admin_c_inst.run().await.unwrap(); - admin_a_inst - .get_conn_manager() - .add_connector(TcpTunnelConnector::new( - "tcp://10.1.1.2:11010".parse().unwrap(), - )); - admin_c_inst - .get_conn_manager() - .add_connector(TcpTunnelConnector::new( - "tcp://10.1.1.2:11010".parse().unwrap(), - )); + admin_a_inst.add_connector_url("tcp://10.1.1.2:11010".parse().unwrap()); + admin_c_inst.add_connector_url("tcp://10.1.1.2:11010".parse().unwrap()); let admin_c_peer_id = admin_c_inst.peer_id(); wait_for_condition( || async { - let a_routes = admin_a_inst.get_peer_manager().list_routes().await; - let c_routes = admin_c_inst.get_peer_manager().list_routes().await; + let a_routes = admin_a_inst.get_core_instance().route_snapshots().await; + let c_routes = admin_c_inst.get_core_instance().route_snapshots().await; a_routes.iter().any(|r| r.peer_id == admin_c_peer_id) || c_routes.iter().any(|r| r.peer_id == admin_a_inst.peer_id()) }, @@ -1830,17 +1833,15 @@ async fn credential_non_reusable_across_two_admins_allows_only_one_peer() { ) .await; - let (_cred_id, cred_secret) = admin_a_inst - .get_global_ctx() - .get_credential_manager() - .generate_credential_with_options( - vec![], - false, - vec![], - Duration::from_secs(3600), - None, - false, - ); + let (_cred_id, cred_secret) = generate_credential_with_options( + &admin_a_inst, + vec![], + false, + vec![], + Duration::from_secs(3600), + None, + false, + ); admin_a_inst .get_global_ctx() .issue_event(GlobalCtxEvent::CredentialChanged); @@ -1867,25 +1868,25 @@ async fn credential_non_reusable_across_two_admins_allows_only_one_peer() { "fd00::6/64", ); - let mut cred_left_inst = Some(Instance::new(cred_left_config)); + let mut cred_left_inst = Some(Instance::new_with_process_runtime( + cred_left_config, + process_runtime.clone(), + )); cred_left_inst.as_mut().unwrap().run().await.unwrap(); cred_left_inst .as_ref() .unwrap() - .get_conn_manager() - .add_connector(TcpTunnelConnector::new( - "tcp://10.1.1.1:11010".parse().unwrap(), - )); + .add_connector_url("tcp://10.1.1.1:11010".parse().unwrap()); - let mut cred_right_inst = Some(Instance::new(cred_right_config)); + let mut cred_right_inst = Some(Instance::new_with_process_runtime( + cred_right_config, + process_runtime.clone(), + )); cred_right_inst.as_mut().unwrap().run().await.unwrap(); cred_right_inst .as_ref() .unwrap() - .get_conn_manager() - .add_connector(TcpTunnelConnector::new( - "tcp://10.1.1.4:11010".parse().unwrap(), - )); + .add_connector_url("tcp://10.1.1.4:11010".parse().unwrap()); let cred_left_peer_id = cred_left_inst.as_ref().unwrap().peer_id(); let cred_right_peer_id = cred_right_inst.as_ref().unwrap().peer_id(); @@ -1947,8 +1948,8 @@ async fn credential_non_reusable_across_two_admins_allows_only_one_peer() { wait_for_condition( || async { - let admin_a_routes = admin_a_inst.get_peer_manager().list_routes().await; - let admin_c_routes = admin_c_inst.get_peer_manager().list_routes().await; + let admin_a_routes = admin_a_inst.get_core_instance().route_snapshots().await; + let admin_c_routes = admin_c_inst.get_core_instance().route_snapshots().await; admin_a_routes.iter().any(|r| r.peer_id == loser_peer_id) && admin_c_routes.iter().any(|r| r.peer_id == loser_peer_id) && !admin_a_routes.iter().any(|r| r.peer_id == winner_peer_id) diff --git a/easytier/src/tests/ipv6_test.rs b/easytier/src/tests/ipv6_test.rs index e953a521..c21a7ddd 100644 --- a/easytier/src/tests/ipv6_test.rs +++ b/easytier/src/tests/ipv6_test.rs @@ -1,10 +1,6 @@ -use std::net::Ipv6Addr; - use crate::{ common::config::{ConfigLoader, TomlConfigLoader}, common::global_ctx::tests::get_mock_global_ctx, - peers::peer_manager::RouteAlgoType, - proto::peer_rpc::RoutePeerInfo, }; #[tokio::test] @@ -30,36 +26,16 @@ async fn test_global_ctx_ipv6() { } #[tokio::test] -async fn test_route_peer_info_ipv6() { +async fn native_peer_config_normalizes_ipv6_route() { let global_ctx = get_mock_global_ctx(); // Set IPv6 address in global context let ipv6_cidr = "fd00::1/64".parse().unwrap(); global_ctx.set_ipv6(Some(ipv6_cidr)); - // Create RoutePeerInfo with IPv6 support - let updated_info = RoutePeerInfo::new_updated_self(123, 456, &global_ctx, None); + let config = crate::instance::config::test_core_instance_config(&global_ctx).peer; + let ipv6 = config.snapshot.runtime.core.routes.ipv6.unwrap(); - // Verify IPv6 address is included - assert!(updated_info.ipv6_addr.is_some()); - let ipv6_addr: Ipv6Addr = updated_info.ipv6_addr.unwrap().address.unwrap().into(); - assert_eq!(ipv6_addr, ipv6_cidr.address()); -} - -#[tokio::test] -async fn test_peer_manager_ipv6() { - let global_ctx = get_mock_global_ctx(); - let (packet_sender, _packet_receiver) = tokio::sync::mpsc::channel(100); - let peer_mgr = crate::peers::peer_manager::PeerManager::new( - RouteAlgoType::Ospf, - global_ctx.clone(), - packet_sender, - ); - - // Test IPv6 address lookup for unknown address - let ipv6_addr = Ipv6Addr::new(0xfd00, 0, 0, 0, 0, 0, 0, 2); - let (peers, _is_self) = peer_mgr.get_msg_dst_peer_ipv6(&ipv6_addr).await; - - // Should return empty peers list for unknown IPv6 - assert!(peers.is_empty()); + assert_eq!(ipv6.address, std::net::IpAddr::V6(ipv6_cidr.address())); + assert_eq!(ipv6.prefix_len, ipv6_cidr.network_length()); } diff --git a/easytier/src/tests/mod.rs b/easytier/src/tests/mod.rs index 3d5385b5..623ca559 100644 --- a/easytier/src/tests/mod.rs +++ b/easytier/src/tests/mod.rs @@ -7,10 +7,39 @@ mod ipv6_test; mod credential_tests; #[cfg(target_os = "linux")] +#[cfg(feature = "upnp")] mod upnp_test; -use crate::common::PeerId; -use crate::peers::peer_manager::PeerManager; +use crate::instance::test_instance::TestInstance as Instance; +use easytier_core::config::PeerId; + +trait InstanceTestExt { + fn add_connector_url(&self, url: url::Url); + + fn peer_id(&self) -> PeerId; + + fn ring_listener_url(&self) -> url::Url; +} + +impl InstanceTestExt for Instance { + fn add_connector_url(&self, url: url::Url) { + self.get_core_instance() + .add_connector(url) + .expect("test connector URL should be supported"); + } + + fn peer_id(&self) -> PeerId { + self.get_core_instance().peer_id() + } + + fn ring_listener_url(&self) -> url::Url { + self.get_core_instance() + .running_listeners() + .into_iter() + .find(|url| url.scheme() == "ring") + .expect("test instance has no running Ring listener") + } +} pub fn set_env_var, V: AsRef>(key: K, value: V) { unsafe { std::env::set_var(key, value) } @@ -108,58 +137,6 @@ pub fn create_netns(name: &str, ipv4: &str, ipv6: &str) { } } -pub struct TestNetnsGuard { - name: String, - host_ipv4: Option, -} - -impl TestNetnsGuard { - fn run_ip(args: &[&str]) { - let status = std::process::Command::new("ip") - .args(args) - .status() - .unwrap(); - assert!(status.success(), "ip command failed: {:?}", args); - } - - pub fn new(name: &str, guest_ipv4: &str, guest_ipv6: &str) -> Self { - del_netns(name); - create_netns(name, guest_ipv4, guest_ipv6); - Self { - name: name.to_string(), - host_ipv4: None, - } - } - - pub fn set_host_ipv4(&mut self, host_ipv4: &str) { - Self::run_ip(&[ - "addr", - "add", - host_ipv4, - "dev", - get_host_veth_name(&self.name), - ]); - self.host_ipv4 = Some(host_ipv4.to_string()); - } -} - -impl Drop for TestNetnsGuard { - fn drop(&mut self) { - if let Some(host_ipv4) = self.host_ipv4.as_deref() { - let _ = std::process::Command::new("ip") - .args([ - "addr", - "del", - host_ipv4, - "dev", - get_host_veth_name(&self.name), - ]) - .status(); - } - del_netns(&self.name); - } -} - pub fn prepare_bridge(name: &str) { // del bridge with brctl let _ = std::process::Command::new("brctl") @@ -186,7 +163,11 @@ pub fn add_ns_to_bridge(br_name: &str, ns_name: &str) { .unwrap(); } -fn check_route(ipv4: &str, dst_peer_id: PeerId, routes: Vec) { +fn check_route( + ipv4: &str, + dst_peer_id: PeerId, + routes: Vec, +) { let mut found = false; for r in routes.iter() { if r.ipv4_addr == Some(ipv4.parse().unwrap()) { @@ -202,9 +183,9 @@ fn check_route(ipv4: &str, dst_peer_id: PeerId, routes: Vec, + routes: Vec, peer_id: PeerId, - checker: impl Fn(&crate::proto::api::instance::Route) -> bool, + checker: impl Fn(&easytier_proto::core_peer::peer::Route) -> bool, ) { let mut found = false; for r in routes.iter() { @@ -217,14 +198,14 @@ fn check_route_ex( } async fn wait_proxy_route_appear( - mgr: &std::sync::Arc, + core: &std::sync::Arc, ipv4: &str, dst_peer_id: PeerId, proxy_cidr: &str, ) { let now = std::time::Instant::now(); loop { - for r in mgr.list_routes().await.iter() { + for r in core.route_snapshots().await.iter() { if r.proxy_cidrs.contains(&proxy_cidr.to_owned()) { assert_eq!(r.peer_id, dst_peer_id); assert_eq!(r.ipv4_addr, Some(ipv4.parse().unwrap())); @@ -255,18 +236,18 @@ fn set_link_status(net_ns: &str, up: bool) { tracing::info!("set link status: {:?}, net_ns: {}, up: {}", ret, net_ns, up); } -pub async fn drop_insts(insts: Vec) { +pub async fn drop_insts(insts: Vec) { let mut set = tokio::task::JoinSet::new(); for mut inst in insts { set.spawn(async move { inst.clear_resources().await; - let pm = std::sync::Arc::downgrade(&inst.get_peer_manager()); + let core = std::sync::Arc::downgrade(&inst.get_core_instance()); drop(inst); let now = std::time::Instant::now(); - while now.elapsed().as_secs() < 5 && pm.strong_count() > 0 { + while now.elapsed().as_secs() < 5 && core.strong_count() > 0 { tokio::time::sleep(std::time::Duration::from_millis(50)).await; } - assert_eq!(pm.strong_count(), 0, "PeerManager should be dropped"); + assert_eq!(core.strong_count(), 0, "CoreInstance should be dropped"); }); } while set.join_next().await.is_some() {} diff --git a/easytier/src/tests/three_node.rs b/easytier/src/tests/three_node.rs index 969188d2..cfc6d275 100644 --- a/easytier/src/tests/three_node.rs +++ b/easytier/src/tests/three_node.rs @@ -7,6 +7,13 @@ use std::{ time::Duration, }; +use easytier_core::{ + connectivity::protocol::raw::TunnelDialer, + foundation::stats::{LabelSet, LabelType, MetricName, MetricSnapshot}, + process_runtime::CoreProcessRuntime, + socket::SocketListener, + tunnel::Tunnel, +}; use rand::{Rng, rngs::OsRng}; use tokio::{net::UdpSocket, task::JoinSet}; use x25519_dalek::StaticSecret; @@ -19,29 +26,80 @@ use crate::{ common::{ config::{ConfigLoader, NetworkIdentity, PortForwardConfig, TomlConfigLoader}, netns::{NetNS, ROOT_NETNS_NAME}, - stats_manager::{LabelSet, LabelType, MetricName}, }, - instance::instance::Instance, + instance::config::test_runtime_instance_config, + instance::test_instance::TestInstance as Instance, proto::{ api::instance::TcpProxyEntryTransportType, common::{CompressionAlgoPb, SecureModeConfig}, - }, - tunnel::{ - common::tests::{ - _tunnel_bench_netns, _tunnel_pingpong_netns_with_timeout, wait_for_condition, + rpc::standalone::{ + RuntimeRpcDialer, RuntimeRpcListener, runtime_rpc_dialer, runtime_rpc_listener, + runtime_udp_tunnel_dialer, runtime_udp_tunnel_listener, }, - ring::RingTunnelConnector, - tcp::{TcpTunnelConnector, TcpTunnelListener}, - udp::UdpTunnelConnector, + }, + tunnel::common::tests::{ + _tunnel_bench_netns, _tunnel_pingpong_netns_with_timeout, wait_for_condition, }, }; +fn metric_value(metrics: &[MetricSnapshot], name: MetricName, labels: &LabelSet) -> Option { + metrics + .iter() + .find(|metric| metric.name == name && metric.labels == *labels) + .map(|metric| metric.value) +} + +fn core_tcp_listener(url: url::Url) -> RuntimeRpcListener { + let addr = url + .socket_addrs(|| Some(11010)) + .expect("test TCP listener URL should resolve") + .into_iter() + .next() + .expect("test TCP listener URL should have an address"); + runtime_rpc_listener(addr) +} + +fn core_tcp_dialer(url: url::Url) -> RuntimeRpcDialer { + runtime_rpc_dialer(url) +} + +fn core_udp_listener(url: url::Url) -> impl SocketListener> + Sync { + let addr = url + .socket_addrs(|| Some(11010)) + .expect("test UDP listener URL should resolve") + .into_iter() + .next() + .expect("test UDP listener URL should have an address"); + runtime_udp_tunnel_listener(url, addr) +} + +fn core_udp_dialer(url: url::Url) -> impl TunnelDialer { + runtime_udp_tunnel_dialer(url) +} + +async fn reload_instance_acl(inst: &Instance, acl: Option<&crate::proto::acl::Acl>) { + let mut config = test_runtime_instance_config(&inst.get_global_ctx()); + config.services.acl = easytier_core::config::peers::AclRuleConfig { + acl: acl.cloned(), + ..Default::default() + }; + inst.get_core_instance() + .update_runtime_config(config) + .await + .unwrap(); +} + +async fn set_foreign_network_refresh_interval(inst: &Instance, seconds: u64) { + let mut config = test_runtime_instance_config(&inst.get_global_ctx()); + Arc::make_mut(&mut config.peer).ospf_update_my_foreign_network_interval_sec = seconds; + inst.get_core_instance() + .update_runtime_config(config) + .await + .unwrap(); +} + #[cfg(feature = "wireguard")] -use crate::{ - common::config::VpnPortalConfig, - tunnel::wireguard::{WgConfig, WgTunnelConnector}, - vpn_portal::wireguard::get_wg_config_for_portal, -}; +use crate::{common::config::VpnPortalConfig, vpn_portal::wireguard::get_wg_config_for_portal}; pub fn prepare_linux_namespaces() { del_netns("net_a"); @@ -89,7 +147,23 @@ pub fn get_inst_config( } pub async fn init_three_node(proto: &str) -> Vec { - init_three_node_ex(proto, |cfg| cfg, false).await + init_three_node_with_process_runtime(proto, CoreProcessRuntime::new()).await +} + +async fn init_three_node_with_process_runtime( + proto: &str, + process_runtime: Arc, +) -> Vec { + init_three_node_ex_with_inst3( + proto, + |cfg| cfg, + false, + "net_c", + "10.144.144.3", + "fd00::3/64", + process_runtime, + ) + .await } async fn init_three_node_ex_with_inst3 TomlConfigLoader>( @@ -99,93 +173,69 @@ async fn init_three_node_ex_with_inst3 TomlConfigLoad inst3_ns: &str, inst3_ipv4: &str, inst3_ipv6: &str, + process_runtime: Arc, ) -> Vec { prepare_linux_namespaces(); - let mut inst1 = Instance::new(cfg_cb(get_inst_config( - "inst1", - Some("net_a"), - "10.144.144.1", - "fd00::1/64", - ))); - let mut inst2 = Instance::new(cfg_cb(get_inst_config( - "inst2", - Some("net_b"), - "10.144.144.2", - "fd00::2/64", - ))); - let mut inst3 = Instance::new(cfg_cb(get_inst_config( - "inst3", - Some(inst3_ns), - inst3_ipv4, - inst3_ipv6, - ))); + let mut inst1 = Instance::new_with_process_runtime( + cfg_cb(get_inst_config( + "inst1", + Some("net_a"), + "10.144.144.1", + "fd00::1/64", + )), + process_runtime.clone(), + ); + let mut inst2 = Instance::new_with_process_runtime( + cfg_cb(get_inst_config( + "inst2", + Some("net_b"), + "10.144.144.2", + "fd00::2/64", + )), + process_runtime.clone(), + ); + let mut inst3 = Instance::new_with_process_runtime( + cfg_cb(get_inst_config( + "inst3", + Some(inst3_ns), + inst3_ipv4, + inst3_ipv6, + )), + process_runtime.clone(), + ); inst1.run().await.unwrap(); inst2.run().await.unwrap(); inst3.run().await.unwrap(); if proto == "tcp" { - inst1 - .get_conn_manager() - .add_connector(TcpTunnelConnector::new( - "tcp://10.1.1.2:11010".parse().unwrap(), - )); + inst1.add_connector_url("tcp://10.1.1.2:11010".parse().unwrap()); } else if proto == "udp" { - inst1 - .get_conn_manager() - .add_connector(UdpTunnelConnector::new( - "udp://10.1.1.2:11010".parse().unwrap(), - )); + inst1.add_connector_url("udp://10.1.1.2:11010".parse().unwrap()); } else if proto == "wg" { #[cfg(feature = "wireguard")] - inst1 - .get_conn_manager() - .add_connector(WgTunnelConnector::new( - "wg://10.1.1.2:11011".parse().unwrap(), - WgConfig::new_from_network_identity( - &inst2.get_global_ctx().get_network_identity().network_name, - &inst2 - .get_global_ctx() - .get_network_identity() - .network_secret - .unwrap_or_default(), - ), - )); + inst1.add_connector_url("wg://10.1.1.2:11011".parse().unwrap()); } else if proto == "ws" { #[cfg(feature = "websocket")] - inst1 - .get_conn_manager() - .add_connector(crate::tunnel::websocket::WsTunnelConnector::new( - "ws://10.1.1.2:11011".parse().unwrap(), - )); + inst1.add_connector_url("ws://10.1.1.2:11011".parse().unwrap()); } else if proto == "wss" { #[cfg(feature = "websocket")] - inst1 - .get_conn_manager() - .add_connector(crate::tunnel::websocket::WsTunnelConnector::new( - "wss://10.1.1.2:11012".parse().unwrap(), - )); + inst1.add_connector_url("wss://10.1.1.2:11012".parse().unwrap()); } - inst3 - .get_conn_manager() - .add_connector(RingTunnelConnector::new( - format!("ring://{}", inst2.id()).parse().unwrap(), - )); + inst3.add_connector_url(inst2.ring_listener_url()); // wait inst2 have two route. wait_for_condition( || async { if !use_public_server { - inst2.get_peer_manager().list_routes().await.len() == 2 + inst2.get_core_instance().route_snapshots().await.len() == 2 } else { inst2 - .get_peer_manager() - .get_foreign_network_manager() - .list_foreign_networks() + .get_core_instance() + .foreign_network_snapshots(false) .await - .foreign_networks .len() == 1 } @@ -196,7 +246,7 @@ async fn init_three_node_ex_with_inst3 TomlConfigLoad wait_for_condition( || async { - let routes = inst1.get_peer_manager().list_routes().await; + let routes = inst1.get_core_instance().route_snapshots().await; println!("routes: {:?}", routes); routes.len() == 2 }, @@ -206,7 +256,7 @@ async fn init_three_node_ex_with_inst3 TomlConfigLoad wait_for_condition( || async { - let routes = inst3.get_peer_manager().list_routes().await; + let routes = inst3.get_core_instance().route_snapshots().await; println!("routes: {:?}", routes); routes.len() == 2 }, @@ -229,6 +279,7 @@ pub async fn init_three_node_ex TomlConfigLoader>( "net_c", "10.144.144.3", "fd00::3/64", + CoreProcessRuntime::new(), ) .await } @@ -237,7 +288,16 @@ async fn init_lazy_p2p_three_node_ex TomlConfigLoader proto: &str, cfg_cb: F, ) -> Vec { - init_three_node_ex_with_inst3(proto, cfg_cb, false, "net_e", "10.144.144.3", "fd00::3/64").await + init_three_node_ex_with_inst3( + proto, + cfg_cb, + false, + "net_e", + "10.144.144.3", + "fd00::3/64", + CoreProcessRuntime::new(), + ) + .await } pub async fn drop_insts(insts: Vec) { @@ -245,115 +305,19 @@ pub async fn drop_insts(insts: Vec) { for mut inst in insts { set.spawn(async move { inst.clear_resources().await; - let pm = Arc::downgrade(&inst.get_peer_manager()); + let core = Arc::downgrade(&inst.get_core_instance()); drop(inst); let now = std::time::Instant::now(); - while now.elapsed().as_secs() < 5 && pm.strong_count() > 0 { + while now.elapsed().as_secs() < 5 && core.strong_count() > 0 { tokio::time::sleep(std::time::Duration::from_millis(50)).await; } - debug_assert_eq!(pm.strong_count(), 0, "PeerManager should be dropped"); + debug_assert_eq!(core.strong_count(), 0, "CoreInstance should be dropped"); }); } while set.join_next().await.is_some() {} } -mod direct_connector_mapped_listener_tests { - use std::sync::Arc; - - use crate::{ - common::{ - config::{ConfigLoader, TomlConfigLoader}, - global_ctx::GlobalCtx, - stun::MockStunInfoCollector, - }, - connector::direct::DirectConnectorManager, - instance::listeners::ListenerManager, - peers::{ - create_packet_recv_chan, - peer_manager::{PeerManager, RouteAlgoType}, - tests::{ - connect_peer_manager, create_mock_peer_manager, wait_route_appear, - wait_route_appear_with_cost, - }, - }, - proto::{common::NatType, peer_rpc::GetIpListResponse}, - tests::TestNetnsGuard, - }; - - async fn create_mock_peer_manager_in_netns(netns: &str) -> Arc { - let (s, _r) = create_packet_recv_chan(); - let config = TomlConfigLoader::default(); - config.set_netns(Some(netns.to_owned())); - let global_ctx = Arc::new(GlobalCtx::new(config)); - global_ctx.replace_stun_info_collector(Box::new(MockStunInfoCollector { - udp_nat_type: NatType::Unknown, - })); - - let peer_mgr = Arc::new(PeerManager::new(RouteAlgoType::Ospf, global_ctx, s)); - peer_mgr.run().await.unwrap(); - peer_mgr - } - - async fn run_direct_connector_mapped_listener_without_port_test( - mapped_listener: &str, - listener: &str, - ) { - let ns_name = "dmlp"; - let mut _ns = TestNetnsGuard::new(ns_name, "10.199.0.2/24", "fd99::2/64"); - _ns.set_host_ipv4("10.199.0.1/24"); - - let p_a = create_mock_peer_manager().await; - let p_b = create_mock_peer_manager().await; - let p_c = create_mock_peer_manager_in_netns(ns_name).await; - connect_peer_manager(p_a.clone(), p_b.clone()).await; - connect_peer_manager(p_b.clone(), p_c.clone()).await; - - wait_route_appear(p_a.clone(), p_c.clone()).await.unwrap(); - - let mut f = p_a.get_global_ctx().get_flags(); - f.bind_device = false; - p_a.get_global_ctx().set_flags(f); - - p_c.get_global_ctx() - .config - .set_mapped_listeners(Some(vec![mapped_listener.parse().unwrap()])); - - p_c.get_global_ctx() - .config - .set_listeners(vec![listener.parse().unwrap()]); - let mut lis_c = ListenerManager::new(p_c.get_global_ctx(), p_c.clone()); - lis_c.prepare_listeners().await.unwrap(); - lis_c.run().await.unwrap(); - - tokio::time::sleep(std::time::Duration::from_secs(1)).await; - let dm_a = DirectConnectorManager::new(p_a.get_global_ctx(), p_a.clone()); - let mut ip_list = GetIpListResponse::default(); - ip_list.listeners.push(mapped_listener.parse().unwrap()); - dm_a.try_direct_connect_with_ip_list(p_c.my_peer_id(), ip_list) - .await - .unwrap(); - - wait_route_appear_with_cost(p_a.clone(), p_c.my_peer_id(), Some(1)) - .await - .unwrap(); - } - - #[rstest::rstest] - #[tokio::test] - #[serial_test::serial] - async fn direct_connector_mapped_listener_without_port( - #[values( - ("tcp://10.199.0.2", "tcp://0.0.0.0:11010"), - ("ws://10.199.0.2", "ws://0.0.0.0:80"), - ("wss://10.199.0.2", "wss://0.0.0.0:443") - )] - case: (&str, &str), - ) { - run_direct_connector_mapped_listener_without_port_test(case.0, case.1).await; - } -} - async fn ping_test(from_netns: &str, target_ip: &str, payload_size: Option) -> bool { let _g = NetNS::new(Some(ROOT_NETNS_NAME.to_owned())).guard(); let code = tokio::process::Command::new("ip") @@ -710,7 +674,7 @@ fn get_public_ipv6_config( async fn init_public_ipv6_two_node( client_inst_id: uuid::Uuid, -) -> (PublicIpv6Lab, Instance, Instance) { +) -> (PublicIpv6Lab, Arc, Instance, Instance) { init_public_ipv6_two_node_with_topology(client_inst_id, PublicIpv6LabTopology::DelegatedPrefix) .await } @@ -718,8 +682,9 @@ async fn init_public_ipv6_two_node( async fn init_public_ipv6_two_node_with_topology( client_inst_id: uuid::Uuid, topology: PublicIpv6LabTopology, -) -> (PublicIpv6Lab, Instance, Instance) { +) -> (PublicIpv6Lab, Arc, Instance, Instance) { let lab = PublicIpv6Lab::setup_with_topology(topology); + let process_runtime = CoreProcessRuntime::new(); let provider_cfg = get_public_ipv6_config( "provider_public_ipv6", @@ -739,43 +704,41 @@ async fn init_public_ipv6_two_node_with_topology( ); client_cfg.set_ipv6_public_addr_auto(true); - let mut provider = Instance::new(provider_cfg); - let mut client = Instance::new(client_cfg); + let mut provider = Instance::new_with_process_runtime(provider_cfg, process_runtime.clone()); + let mut client = Instance::new_with_process_runtime(client_cfg, process_runtime.clone()); provider.run().await.unwrap(); client.run().await.unwrap(); - provider - .get_conn_manager() - .add_connector(TcpTunnelConnector::new( - "tcp://10.1.1.2:11010".parse().unwrap(), - )); + provider.add_connector_url("tcp://10.1.1.2:11010".parse().unwrap()); wait_for_condition( || async { - provider.get_peer_manager().list_routes().await.len() == 1 - && client.get_peer_manager().list_routes().await.len() == 1 + provider.get_core_instance().route_snapshots().await.len() == 1 + && client.get_core_instance().route_snapshots().await.len() == 1 }, Duration::from_secs(8), ) .await; - (lab, provider, client) + (lab, process_runtime, provider, client) } async fn wait_for_public_ipv6_addr(inst: &Instance) -> cidr::Ipv6Inet { wait_for_condition( || async { - inst.get_peer_manager() - .get_my_public_ipv6_addr() + inst.get_core_instance() + .packet_plane() + .public_ipv6_addr() .await .is_some() }, Duration::from_secs(10), ) .await; - inst.get_peer_manager() - .get_my_public_ipv6_addr() + inst.get_core_instance() + .packet_plane() + .public_ipv6_addr() .await .unwrap() } @@ -783,8 +746,9 @@ async fn wait_for_public_ipv6_addr(inst: &Instance) -> cidr::Ipv6Inet { async fn wait_for_public_ipv6_route(inst: &Instance, target: cidr::Ipv6Inet) { wait_for_condition( || async { - inst.get_peer_manager() - .list_public_ipv6_routes() + inst.get_core_instance() + .packet_plane() + .public_ipv6_routes() .await .contains(&target) }, @@ -814,13 +778,15 @@ fn ndp_proxy_exists_in_ns(ns: &str, dev: &str, addr: std::net::Ipv6Addr) -> bool #[serial_test::serial] pub async fn public_ipv6_auto_addr_end_to_end() { let client_id = uuid::Uuid::parse_str("22222222-2222-2222-2222-222222222222").unwrap(); - let (_lab, provider, client) = init_public_ipv6_two_node(client_id).await; + let (_lab, _process_runtime, provider, client) = init_public_ipv6_two_node(client_id).await; wait_for_condition( || async { provider - .get_global_ctx() - .get_advertised_ipv6_public_addr_prefix() + .get_core_instance() + .node_snapshot() + .await + .ipv6_public_addr_prefix == Some(PublicIpv6Lab::PROVIDER_PREFIX.parse().unwrap()) }, Duration::from_secs(10), @@ -839,8 +805,10 @@ pub async fn public_ipv6_auto_addr_end_to_end() { ); assert_eq!( provider - .get_global_ctx() - .get_advertised_ipv6_public_addr_prefix(), + .get_core_instance() + .node_snapshot() + .await + .ipv6_public_addr_prefix, Some(PublicIpv6Lab::PROVIDER_PREFIX.parse().unwrap()) ); let provider_prefix = PublicIpv6Lab::PROVIDER_PREFIX @@ -848,8 +816,8 @@ pub async fn public_ipv6_auto_addr_end_to_end() { .unwrap(); assert_eq!( provider - .get_peer_manager() - .get_my_info() + .get_core_instance() + .node_snapshot() .await .ipv6_public_addr_prefix, Some( @@ -858,14 +826,10 @@ pub async fn public_ipv6_auto_addr_end_to_end() { provider_prefix.network_length() ) .unwrap() - .into() ) ); - let provider_info = provider - .get_peer_manager() - .get_local_public_ipv6_info() - .await; - let client_peer_id = client.get_peer_manager().get_my_info().await.peer_id; + let provider_info = provider.get_core_instance().local_public_ipv6_info().await; + let client_peer_id = client.get_core_instance().node_snapshot().await.peer_id; assert_eq!( provider_info.provider_prefix, Some( @@ -936,15 +900,17 @@ pub async fn public_ipv6_auto_addr_end_to_end() { #[serial_test::serial] pub async fn public_ipv6_auto_addr_on_link_ndp_proxy_end_to_end() { let client_id = uuid::Uuid::parse_str("44444444-4444-4444-4444-444444444444").unwrap(); - let (_lab, provider, client) = + let (_lab, _process_runtime, provider, client) = init_public_ipv6_two_node_with_topology(client_id, PublicIpv6LabTopology::OnLinkPrefix) .await; wait_for_condition( || async { provider - .get_global_ctx() - .get_advertised_ipv6_public_addr_prefix() + .get_core_instance() + .node_snapshot() + .await + .ipv6_public_addr_prefix == Some(PublicIpv6Lab::PROVIDER_PREFIX.parse().unwrap()) }, Duration::from_secs(10), @@ -997,7 +963,7 @@ pub async fn public_ipv6_auto_addr_on_link_ndp_proxy_end_to_end() { #[serial_test::serial] pub async fn public_ipv6_auto_addr_reconnect_reuses_same_address() { let client_id = uuid::Uuid::parse_str("33333333-3333-3333-3333-333333333333").unwrap(); - let (_lab, provider, client) = init_public_ipv6_two_node(client_id).await; + let (_lab, process_runtime, provider, client) = init_public_ipv6_two_node(client_id).await; let first = wait_for_public_ipv6_addr(&client).await; drop_insts(vec![client]).await; @@ -1010,18 +976,14 @@ pub async fn public_ipv6_auto_addr_reconnect_reuses_same_address() { client_id, ); client_cfg.set_ipv6_public_addr_auto(true); - let mut client = Instance::new(client_cfg); + let mut client = Instance::new_with_process_runtime(client_cfg, process_runtime.clone()); client.run().await.unwrap(); - provider - .get_conn_manager() - .add_connector(TcpTunnelConnector::new( - "tcp://10.1.1.2:11010".parse().unwrap(), - )); + provider.add_connector_url("tcp://10.1.1.2:11010".parse().unwrap()); wait_for_condition( || async { - provider.get_peer_manager().list_routes().await.len() == 1 - && client.get_peer_manager().list_routes().await.len() == 1 + provider.get_core_instance().route_snapshots().await.len() == 1 + && client.get_core_instance().route_snapshots().await.len() == 1 }, Duration::from_secs(8), ) @@ -1080,13 +1042,13 @@ pub async fn basic_three_node_test( check_route( "10.144.144.2/24", insts[1].peer_id(), - insts[0].get_peer_manager().list_routes().await, + insts[0].get_core_instance().route_snapshots().await, ); check_route( "10.144.144.3/24", insts[2].peer_id(), - insts[0].get_peer_manager().list_routes().await, + insts[0].get_core_instance().route_snapshots().await, ); // Test IPv4 connectivity @@ -1157,7 +1119,7 @@ pub async fn subnet_proxy_loop_prevention_test() { // 等待代理路由出现 - inst1 应该看到 inst2 的代理路由 wait_proxy_route_appear( - &insts[0].get_peer_manager(), + &insts[0].get_core_instance(), "10.144.144.2/24", insts[1].peer_id(), "10.1.2.0/24", @@ -1166,7 +1128,7 @@ pub async fn subnet_proxy_loop_prevention_test() { // 等待代理路由出现 - inst2 应该看到 inst1 的代理路由 wait_proxy_route_appear( - &insts[1].get_peer_manager(), + &insts[1].get_core_instance(), "10.144.144.1/24", insts[0].peer_id(), "10.1.2.0/24", @@ -1183,20 +1145,13 @@ pub async fn subnet_proxy_loop_prevention_test() { println!( "inst0 metrics: {:?}", - insts[0] - .get_global_ctx() - .stats_manager() - .export_prometheus() + insts[0].get_core_instance().prometheus_metrics() ); - let all_metrics = insts[0].get_global_ctx().stats_manager().get_all_metrics(); + let all_metrics = insts[0].get_core_instance().metric_snapshots(); for metric in all_metrics { if metric.name == MetricName::TrafficPacketsSelfTx { - let counter = insts[0] - .get_global_ctx() - .stats_manager() - .get_counter(metric.name, metric.labels.clone()); - assert!(counter.get() < 40); + assert!(metric.value < 40); } } @@ -1204,15 +1159,11 @@ pub async fn subnet_proxy_loop_prevention_test() { } async fn subnet_proxy_test_udp(listen_ip: &str, target_ip: &str, timeout: Duration) { - use crate::tunnel::{ - common::tests::_tunnel_pingpong_netns_with_timeout, udp::UdpTunnelListener, - }; + use crate::tunnel::common::tests::_tunnel_pingpong_netns_with_timeout; use rand::Rng; - let udp_listener = - UdpTunnelListener::new(format!("udp://{}:22233", listen_ip).parse().unwrap()); - let udp_connector = - UdpTunnelConnector::new(format!("udp://{}:22233", target_ip).parse().unwrap()); + let udp_listener = core_udp_listener(format!("udp://{}:22233", listen_ip).parse().unwrap()); + let udp_connector = core_udp_dialer(format!("udp://{}:22233", target_ip).parse().unwrap()); // NOTE: this should not excced udp tunnel max buffer size let mut buf = vec![0; 7 * 1024]; @@ -1236,10 +1187,8 @@ async fn subnet_proxy_test_udp(listen_ip: &str, target_ip: &str, timeout: Durati assert!(result.is_ok(), "{}", result.unwrap_err()); // no fragment - let udp_listener = - UdpTunnelListener::new(format!("udp://{}:22233", listen_ip).parse().unwrap()); - let udp_connector = - UdpTunnelConnector::new(format!("udp://{}:22233", target_ip).parse().unwrap()); + let udp_listener = core_udp_listener(format!("udp://{}:22233", listen_ip).parse().unwrap()); + let udp_connector = core_udp_dialer(format!("udp://{}:22233", target_ip).parse().unwrap()); let mut buf = vec![0; 1024]; rand::thread_rng().fill(&mut buf[..]); @@ -1257,14 +1206,11 @@ async fn subnet_proxy_test_udp(listen_ip: &str, target_ip: &str, timeout: Durati } async fn subnet_proxy_test_tcp(listen_ip: &str, connect_ip: &str, timeout: Duration) { - use crate::tunnel::{ - common::tests::_tunnel_pingpong_netns_with_timeout, tcp::TcpTunnelListener, - }; + use crate::tunnel::common::tests::_tunnel_pingpong_netns_with_timeout; use rand::Rng; - let tcp_listener = TcpTunnelListener::new(format!("tcp://{listen_ip}:22223").parse().unwrap()); - let tcp_connector = - TcpTunnelConnector::new(format!("tcp://{}:22223", connect_ip).parse().unwrap()); + let tcp_listener = core_tcp_listener(format!("tcp://{listen_ip}:22223").parse().unwrap()); + let tcp_connector = core_tcp_dialer(format!("tcp://{}:22223", connect_ip).parse().unwrap()); let mut buf = vec![0; 32]; rand::thread_rng().fill(&mut buf[..]); @@ -1323,7 +1269,7 @@ pub async fn quic_proxy() { assert_eq!(insts[2].get_global_ctx().config.get_proxy_cidrs().len(), 1); wait_proxy_route_appear( - &insts[0].get_peer_manager(), + &insts[0].get_core_instance(), "10.144.144.3/24", insts[2].peer_id(), "10.1.2.0/24", @@ -1338,9 +1284,11 @@ pub async fn quic_proxy() { subnet_proxy_test_tcp("0.0.0.0", "10.144.144.3", Duration::from_secs(5)).await; let metrics = insts[0] - .get_global_ctx() - .stats_manager() - .get_metrics_by_prefix(&MetricName::TcpProxyConnect.to_string()); + .get_core_instance() + .metric_snapshots() + .into_iter() + .filter(|metric| metric.name == MetricName::TcpProxyConnect) + .collect::>(); assert_eq!(metrics.len(), 2); assert_eq!(1, metrics[0].value); assert_eq!(1, metrics[1].value); @@ -1408,14 +1356,14 @@ pub async fn subnet_proxy_three_node_test( assert_eq!(insts[2].get_global_ctx().config.get_proxy_cidrs().len(), 2); wait_proxy_route_appear( - &insts[0].get_peer_manager(), + &insts[0].get_core_instance(), "10.144.144.3/24", insts[2].peer_id(), "10.1.2.0/24", ) .await; wait_proxy_route_appear( - &insts[0].get_peer_manager(), + &insts[0].get_core_instance(), "10.144.144.3/24", insts[2].peer_id(), "10.1.3.0/24", @@ -1434,9 +1382,11 @@ pub async fn subnet_proxy_three_node_test( } if enable_quic_proxy && !disable_quic_input { let metrics = insts[0] - .get_global_ctx() - .stats_manager() - .get_metrics_by_prefix(&MetricName::TcpProxyConnect.to_string()); + .get_core_instance() + .metric_snapshots() + .into_iter() + .filter(|metric| metric.name == MetricName::TcpProxyConnect) + .collect::>(); assert_eq!(metrics.len(), 3); for metric in metrics { assert_eq!(1, metric.value); @@ -1448,9 +1398,11 @@ pub async fn subnet_proxy_three_node_test( } } else if enable_kcp_proxy && !disable_kcp_input { let metrics = insts[0] - .get_global_ctx() - .stats_manager() - .get_metrics_by_prefix(&MetricName::TcpProxyConnect.to_string()); + .get_core_instance() + .metric_snapshots() + .into_iter() + .filter(|metric| metric.name == MetricName::TcpProxyConnect) + .collect::>(); assert_eq!(metrics.len(), 3); for metric in metrics { assert_eq!(1, metric.value); @@ -1463,9 +1415,11 @@ pub async fn subnet_proxy_three_node_test( } else { // tcp subnet proxy let metrics = insts[2] - .get_global_ctx() - .stats_manager() - .get_metrics_by_prefix(&MetricName::TcpProxyConnect.to_string()); + .get_core_instance() + .metric_snapshots() + .into_iter() + .filter(|metric| metric.name == MetricName::TcpProxyConnect) + .collect::>(); if no_tun { assert_eq!(metrics.len(), 3); } else { @@ -1532,36 +1486,18 @@ pub async fn data_compress( #[tokio::test] #[serial_test::serial] pub async fn proxy_three_node_disconnect_test(#[values("tcp", "wg")] proto: &str) { - use crate::tunnel::wireguard::{WgConfig, WgTunnelConnector}; use tokio_util::task::AbortOnDropHandle; - let insts = init_three_node(proto).await; - let mut inst4 = Instance::new(get_inst_config( - "inst4", - Some("net_d"), - "10.144.144.4", - "fd00::4/64", - )); + let process_runtime = CoreProcessRuntime::new(); + let insts = init_three_node_with_process_runtime(proto, process_runtime.clone()).await; + let mut inst4 = Instance::new_with_process_runtime( + get_inst_config("inst4", Some("net_d"), "10.144.144.4", "fd00::4/64"), + process_runtime.clone(), + ); if proto == "tcp" { - inst4 - .get_conn_manager() - .add_connector(TcpTunnelConnector::new( - "tcp://10.1.2.3:11010".parse().unwrap(), - )); + inst4.add_connector_url("tcp://10.1.2.3:11010".parse().unwrap()); } else if proto == "wg" { - inst4 - .get_conn_manager() - .add_connector(WgTunnelConnector::new( - "wg://10.1.2.3:11011".parse().unwrap(), - WgConfig::new_from_network_identity( - &inst4.get_global_ctx().get_network_identity().network_name, - &inst4 - .get_global_ctx() - .get_network_identity() - .network_secret - .unwrap_or_default(), - ), - )); + inst4.add_connector_url("wg://10.1.2.3:11011".parse().unwrap()); } else { unreachable!("not support"); } @@ -1578,8 +1514,8 @@ pub async fn proxy_three_node_disconnect_test(#[values("tcp", "wg")] proto: &str wait_for_condition( || async { insts[0] - .get_peer_manager() - .list_routes() + .get_core_instance() + .route_snapshots() .await .iter() .any(|r| r.peer_id == inst4.peer_id()) @@ -1598,9 +1534,8 @@ pub async fn proxy_three_node_disconnect_test(#[values("tcp", "wg")] proto: &str wait_for_condition( || async { !insts[2] - .get_peer_manager() - .get_peer_map() - .list_peers_with_conn() + .get_core_instance() + .connected_peers() .await .iter() .any(|r| *r == inst4.peer_id()) @@ -1615,8 +1550,8 @@ pub async fn proxy_three_node_disconnect_test(#[values("tcp", "wg")] proto: &str wait_for_condition( || async { !insts[0] - .get_peer_manager() - .list_routes() + .get_core_instance() + .route_snapshots() .await .iter() .any(|r| r.peer_id == inst4.peer_id()) @@ -1689,48 +1624,38 @@ pub async fn udp_broadcast_test() { #[serial_test::serial] pub async fn foreign_network_forward_nic_data() { prepare_linux_namespaces(); + let process_runtime = CoreProcessRuntime::new(); let center_node_config = get_inst_config("inst1", Some("net_a"), "10.144.144.1", "fd00::1/64"); center_node_config .set_network_identity(NetworkIdentity::new("center".to_string(), "".to_string())); - let mut center_inst = Instance::new(center_node_config); + let mut center_inst = + Instance::new_with_process_runtime(center_node_config, process_runtime.clone()); - let mut inst1 = Instance::new(get_inst_config( - "inst1", - Some("net_b"), - "10.144.145.1", - "fd00:1::1/64", - )); - let mut inst2 = Instance::new(get_inst_config( - "inst2", - Some("net_c"), - "10.144.145.2", - "fd00:1::2/64", - )); + let mut inst1 = Instance::new_with_process_runtime( + get_inst_config("inst1", Some("net_b"), "10.144.145.1", "fd00:1::1/64"), + process_runtime.clone(), + ); + let mut inst2 = Instance::new_with_process_runtime( + get_inst_config("inst2", Some("net_c"), "10.144.145.2", "fd00:1::2/64"), + process_runtime, + ); center_inst.run().await.unwrap(); inst1.run().await.unwrap(); inst2.run().await.unwrap(); - assert_ne!(inst1.id(), center_inst.id()); - assert_ne!(inst2.id(), center_inst.id()); + assert_ne!(inst1.ring_listener_url(), center_inst.ring_listener_url()); + assert_ne!(inst2.ring_listener_url(), center_inst.ring_listener_url()); - inst1 - .get_conn_manager() - .add_connector(RingTunnelConnector::new( - format!("ring://{}", center_inst.id()).parse().unwrap(), - )); + inst1.add_connector_url(center_inst.ring_listener_url()); - inst2 - .get_conn_manager() - .add_connector(RingTunnelConnector::new( - format!("ring://{}", center_inst.id()).parse().unwrap(), - )); + inst2.add_connector_url(center_inst.ring_listener_url()); wait_for_condition( || async { - inst1.get_peer_manager().list_routes().await.len() == 2 - && inst2.get_peer_manager().list_routes().await.len() == 2 + inst1.get_core_instance().route_snapshots().await.len() == 2 + && inst2.get_core_instance().route_snapshots().await.len() == 2 }, Duration::from_secs(5), ) @@ -1802,7 +1727,20 @@ fn run_wireguard_client( #[tokio::test] #[serial_test::serial] pub async fn wireguard_vpn_portal(#[values(true, false)] test_v6: bool) { - let mut insts = init_three_node("tcp").await; + let insts = init_three_node_ex( + "tcp", + |config| { + if config.get_inst_name() == "inst3" { + config.set_vpn_portal_config(VpnPortalConfig { + wireguard_listen: "0.0.0.0:22121".parse().unwrap(), + client_cidr: "10.14.14.0/24".parse().unwrap(), + }); + } + config + }, + false, + ) + .await; if test_v6 { ping6_test("net_d", "fd12::3", None).await; @@ -1810,17 +1748,6 @@ pub async fn wireguard_vpn_portal(#[values(true, false)] test_v6: bool) { ping_test("net_d", "10.1.2.3", None).await; } - let net_ns = NetNS::new(Some("net_d".into())); - let _g = net_ns.guard(); - insts[2] - .get_global_ctx() - .config - .set_vpn_portal_config(VpnPortalConfig { - wireguard_listen: "0.0.0.0:22121".parse().unwrap(), - client_cidr: "10.14.14.0/24".parse().unwrap(), - }); - insts[2].run_vpn_portal().await.unwrap(); - let dst_socket_addr = if test_v6 { "[fd12::3]:22121".parse().unwrap() } else { @@ -1948,56 +1875,49 @@ pub async fn socks5_vpn_portal( #[tokio::test] #[serial_test::serial] pub async fn foreign_network_functional_cluster() { - crate::set_global_var!(OSPF_UPDATE_MY_GLOBAL_FOREIGN_NETWORK_INTERVAL_SEC, 1); prepare_linux_namespaces(); + let process_runtime = CoreProcessRuntime::new(); let center_node_config1 = get_inst_config("inst1", Some("net_a"), "10.144.144.1", "fd00::1/64"); center_node_config1 .set_network_identity(NetworkIdentity::new("center".to_string(), "".to_string())); - let mut center_inst1 = Instance::new(center_node_config1); + let mut center_inst1 = + Instance::new_with_process_runtime(center_node_config1, process_runtime.clone()); let center_node_config2 = get_inst_config("inst2", Some("net_b"), "10.144.144.2", "fd00::2/64"); center_node_config2 .set_network_identity(NetworkIdentity::new("center".to_string(), "".to_string())); - let mut center_inst2 = Instance::new(center_node_config2); + let mut center_inst2 = + Instance::new_with_process_runtime(center_node_config2, process_runtime.clone()); let inst1_config = get_inst_config("inst1", Some("net_c"), "10.144.145.1", "fd00:2::1/64"); inst1_config.set_listeners(vec![]); - let mut inst1 = Instance::new(inst1_config); + let mut inst1 = Instance::new_with_process_runtime(inst1_config, process_runtime.clone()); - let mut inst2 = Instance::new(get_inst_config( - "inst2", - Some("net_d"), - "10.144.145.2", - "fd00:2::2/64", - )); + let mut inst2 = Instance::new_with_process_runtime( + get_inst_config("inst2", Some("net_d"), "10.144.145.2", "fd00:2::2/64"), + process_runtime, + ); center_inst1.run().await.unwrap(); center_inst2.run().await.unwrap(); inst1.run().await.unwrap(); inst2.run().await.unwrap(); - center_inst1 - .get_conn_manager() - .add_connector(RingTunnelConnector::new( - format!("ring://{}", center_inst2.id()).parse().unwrap(), - )); + for instance in [¢er_inst1, ¢er_inst2, &inst1, &inst2] { + set_foreign_network_refresh_interval(instance, 1).await; + } - inst1 - .get_conn_manager() - .add_connector(RingTunnelConnector::new( - format!("ring://{}", center_inst1.id()).parse().unwrap(), - )); + center_inst1.add_connector_url(center_inst2.ring_listener_url()); - inst2 - .get_conn_manager() - .add_connector(RingTunnelConnector::new( - format!("ring://{}", center_inst2.id()).parse().unwrap(), - )); + inst1.add_connector_url(center_inst1.ring_listener_url()); - let peer_map_inst1 = inst1.get_peer_manager(); - println!("inst1 peer map: {:?}", peer_map_inst1.list_routes().await); - drop(peer_map_inst1); + inst2.add_connector_url(center_inst2.ring_listener_url()); + + println!( + "inst1 peer map: {:?}", + inst1.get_core_instance().route_snapshots().await + ); wait_for_condition( || async { ping_test("net_c", "10.144.145.2", None).await }, @@ -2006,11 +1926,7 @@ pub async fn foreign_network_functional_cluster() { .await; // connect to two centers, ping should work - inst1 - .get_conn_manager() - .add_connector(RingTunnelConnector::new( - format!("ring://{}", center_inst2.id()).parse().unwrap(), - )); + inst1.add_connector_url(center_inst2.ring_listener_url()); tokio::time::sleep(tokio::time::Duration::from_secs(5)).await; wait_for_condition( || async { ping_test("net_c", "10.144.145.2", None).await }, @@ -2026,67 +1942,60 @@ pub async fn foreign_network_functional_cluster() { #[serial_test::serial] pub async fn manual_reconnector(#[values(true, false)] is_foreign: bool) { prepare_linux_namespaces(); + let process_runtime = CoreProcessRuntime::new(); let center_node_config = get_inst_config("inst1", Some("net_a"), "10.144.144.1", "fd00::1/64"); if is_foreign { center_node_config .set_network_identity(NetworkIdentity::new("center".to_string(), "".to_string())); } - let mut center_inst = Instance::new(center_node_config); + let mut center_inst = + Instance::new_with_process_runtime(center_node_config, process_runtime.clone()); let inst1_config = get_inst_config("inst1", Some("net_b"), "10.144.145.1", "fd00:1::1/64"); inst1_config.set_listeners(vec![]); - let mut inst1 = Instance::new(inst1_config); + let mut inst1 = Instance::new_with_process_runtime(inst1_config, process_runtime.clone()); - let mut inst2 = Instance::new(get_inst_config( - "inst2", - Some("net_c"), - "10.144.145.2", - "fd00:1::2/64", - )); + let mut inst2 = Instance::new_with_process_runtime( + get_inst_config("inst2", Some("net_c"), "10.144.145.2", "fd00:1::2/64"), + process_runtime, + ); center_inst.run().await.unwrap(); inst1.run().await.unwrap(); inst2.run().await.unwrap(); - assert_ne!(inst1.id(), center_inst.id()); - assert_ne!(inst2.id(), center_inst.id()); + assert_ne!(inst1.ring_listener_url(), center_inst.ring_listener_url()); + assert_ne!(inst2.ring_listener_url(), center_inst.ring_listener_url()); - inst1 - .get_conn_manager() - .add_connector(RingTunnelConnector::new( - format!("ring://{}", center_inst.id()).parse().unwrap(), - )); + inst1.add_connector_url(center_inst.ring_listener_url()); - inst2 - .get_conn_manager() - .add_connector(RingTunnelConnector::new( - format!("ring://{}", center_inst.id()).parse().unwrap(), - )); + inst2.add_connector_url(center_inst.ring_listener_url()); tokio::time::sleep(tokio::time::Duration::from_secs(5)).await; - let peer_map = if !is_foreign { - inst1.get_peer_manager().get_peer_map() - } else { - inst1 - .get_peer_manager() - .get_foreign_network_client() - .get_peer_map() - }; let center_inst_peer_id = if !is_foreign { center_inst.peer_id() } else { + let network_name = inst1.get_global_ctx().get_network_identity().network_name; center_inst - .get_peer_manager() - .get_foreign_network_manager() - .get_network_peer_id(&inst1.get_global_ctx().get_network_identity().network_name) + .get_core_instance() + .foreign_network_snapshots(false) + .await + .get(&network_name) + .map(|network| network.my_peer_id_for_this_network) .unwrap() }; - let conns = peer_map.list_peer_conns(center_inst_peer_id).await.unwrap(); + let conns_len = inst1 + .get_core_instance() + .peer_snapshots() + .await + .into_iter() + .find(|peer| peer.peer_id == center_inst_peer_id) + .map_or(0, |peer| peer.conns.len()); - assert!(!conns.is_empty()); + assert!(conns_len > 0); wait_for_condition( || async { ping_test("net_b", "10.144.145.2", None).await }, @@ -2094,7 +2003,6 @@ pub async fn manual_reconnector(#[values(true, false)] is_foreign: bool) { ) .await; - drop(peer_map); drop_insts(vec![center_inst, inst1, inst2]).await; } @@ -2163,10 +2071,8 @@ pub async fn port_forward_test( ) .await; - use crate::tunnel::{tcp::TcpTunnelListener, udp::UdpTunnelConnector, udp::UdpTunnelListener}; - - let tcp_listener = TcpTunnelListener::new("tcp://0.0.0.0:23456".parse().unwrap()); - let tcp_connector = TcpTunnelConnector::new("tcp://127.0.0.1:23456".parse().unwrap()); + let tcp_listener = core_tcp_listener("tcp://0.0.0.0:23456".parse().unwrap()); + let tcp_connector = core_tcp_dialer("tcp://127.0.0.1:23456".parse().unwrap()); let mut buf = vec![0; buf_size as usize]; rand::thread_rng().fill(&mut buf[..]); @@ -2182,8 +2088,8 @@ pub async fn port_forward_test( .await .unwrap(); - let tcp_listener = TcpTunnelListener::new("tcp://0.0.0.0:23457".parse().unwrap()); - let tcp_connector = TcpTunnelConnector::new("tcp://127.0.0.1:23457".parse().unwrap()); + let tcp_listener = core_tcp_listener("tcp://0.0.0.0:23457".parse().unwrap()); + let tcp_connector = core_tcp_dialer("tcp://127.0.0.1:23457".parse().unwrap()); let mut buf = vec![0; buf_size as usize]; rand::thread_rng().fill(&mut buf[..]); @@ -2199,8 +2105,8 @@ pub async fn port_forward_test( .await .unwrap(); - let udp_listener = UdpTunnelListener::new("udp://0.0.0.0:23458".parse().unwrap()); - let udp_connector = UdpTunnelConnector::new("udp://127.0.0.1:23458".parse().unwrap()); + let udp_listener = core_udp_listener("udp://0.0.0.0:23458".parse().unwrap()); + let udp_connector = core_udp_dialer("udp://127.0.0.1:23458".parse().unwrap()); let mut buf = vec![0; buf_size as usize]; rand::thread_rng().fill(&mut buf[..]); @@ -2216,8 +2122,8 @@ pub async fn port_forward_test( .await .unwrap(); - let udp_listener = UdpTunnelListener::new("udp://0.0.0.0:23459".parse().unwrap()); - let udp_connector = UdpTunnelConnector::new("udp://127.0.0.1:23459".parse().unwrap()); + let udp_listener = core_udp_listener("udp://0.0.0.0:23459".parse().unwrap()); + let udp_connector = core_udp_dialer("udp://127.0.0.1:23459".parse().unwrap()); let mut buf = vec![0; buf_size as usize]; rand::thread_rng().fill(&mut buf[..]); @@ -2317,10 +2223,9 @@ pub async fn port_forward_with_inbound_default_drop_acl_test( } for (bind_port, server_ns) in [(23456, "net_c"), (23457, "net_d")] { - let tcp_listener = - TcpTunnelListener::new(format!("tcp://0.0.0.0:{bind_port}").parse().unwrap()); + let tcp_listener = core_tcp_listener(format!("tcp://0.0.0.0:{bind_port}").parse().unwrap()); let tcp_connector = - TcpTunnelConnector::new(format!("tcp://127.0.0.1:{bind_port}").parse().unwrap()); + core_tcp_dialer(format!("tcp://127.0.0.1:{bind_port}").parse().unwrap()); let mut buf = vec![0; 64]; rand::thread_rng().fill(&mut buf[..]); @@ -2335,7 +2240,7 @@ pub async fn port_forward_with_inbound_default_drop_acl_test( ) .await; - let stats = insts[0].get_global_ctx().get_acl_filter().get_stats(); + let stats = insts[0].get_core_instance().acl_stats(); println!( "port forward source bind_port={} dhcp={} enable_quic_proxy={} ACL stats: {}", bind_port, dhcp, enable_quic_proxy, stats @@ -2377,8 +2282,8 @@ pub async fn relay_bps_limit_test(#[values(100, 200, 400, 800)] bps_limit: u64) .await; // connect to virtual ip (no tun mode) - let tcp_listener = TcpTunnelListener::new("tcp://0.0.0.0:22223".parse().unwrap()); - let tcp_connector = TcpTunnelConnector::new("tcp://10.144.144.3:22223".parse().unwrap()); + let tcp_listener = core_tcp_listener("tcp://0.0.0.0:22223".parse().unwrap()); + let tcp_connector = core_tcp_dialer("tcp://10.144.144.3:22223".parse().unwrap()); let bps = _tunnel_bench_netns( tcp_listener, @@ -2420,8 +2325,8 @@ pub async fn instance_recv_bps_limit_test(#[values(100, 800)] bps_limit: u64) { ) .await; - let tcp_listener = TcpTunnelListener::new("tcp://0.0.0.0:22223".parse().unwrap()); - let tcp_connector = TcpTunnelConnector::new("tcp://10.144.144.3:22223".parse().unwrap()); + let tcp_listener = core_tcp_listener("tcp://0.0.0.0:22223".parse().unwrap()); + let tcp_connector = core_tcp_dialer("tcp://10.144.144.3:22223".parse().unwrap()); let bps = _tunnel_bench_netns( tcp_listener, @@ -2444,17 +2349,50 @@ pub async fn instance_recv_bps_limit_test(#[values(100, 800)] bps_limit: u64) { drop_insts(insts).await; } -async fn assert_try_direct_connect_err(inst: &Instance, connector: C) -where - C: crate::tunnel::TunnelConnector + std::fmt::Debug, -{ - let ret = tokio::time::timeout( - Duration::from_millis(100), - inst.get_peer_manager().try_direct_connect(connector), - ) - .await; - - assert!(matches!(ret, Err(_) | Ok(Err(_)))); +async fn assert_peer_admission_blocked(inst: &Instance, url: url::Url) { + let ip = url + .host_str() + .expect("test URL should have a host") + .parse() + .expect("test URL should have a literal IP"); + let target = std::net::SocketAddr::new(ip, url.port().expect("test URL should have a port")); + let host = crate::instance::host::native_instance_host(inst.get_global_ctx()); + let protocol = crate::tunnel::protocol::runtime_client_protocol_upgrader(inst.get_global_ctx()); + let core = inst.get_core_instance(); + let connect = async { + let connected = match url.scheme() { + "tcp" => easytier_core::connectivity::transport::ConnectedTransport::Tcp( + easytier_core::socket::tcp::VirtualTcpSocketFactory::connect_tcp( + host.as_ref(), + easytier_core::socket::tcp::TcpConnectOptions::direct_connect(target), + ) + .await?, + ), + "udp" => easytier_core::connectivity::transport::ConnectedTransport::Udp( + easytier_core::connectivity::transport::connect_udp( + host, + target, + Vec::new(), + easytier_core::socket::udp::UdpBindOptions::direct_connect(), + easytier_core::connectivity::transport::UdpSessionMode::EasyTierMux, + ) + .await?, + ), + scheme => panic!("unsupported test scheme: {scheme}"), + }; + let tunnel = easytier_core::connectivity::protocol::ClientProtocolUpgrader::upgrade_client( + protocol.as_ref(), + connected, + url, + ) + .await?; + core.admit_client_tunnel_for_test(tunnel, true) + .await + .map(|_| ()) + .map_err(anyhow::Error::from) + }; + let result = tokio::time::timeout(Duration::from_millis(100), connect).await; + assert!(matches!(result, Err(_) | Ok(Err(_)))); } use std::fs; @@ -2525,29 +2463,13 @@ async fn avoid_tunnel_loop_back_to_virtual_network( ) .await; - assert_try_direct_connect_err( - &insts[0], - TcpTunnelConnector::new("tcp://10.144.144.2:11010".parse().unwrap()), - ) - .await; + assert_peer_admission_blocked(&insts[0], "tcp://10.144.144.2:11010".parse().unwrap()).await; - assert_try_direct_connect_err( - &insts[0], - UdpTunnelConnector::new("udp://10.144.144.3:11010".parse().unwrap()), - ) - .await; + assert_peer_admission_blocked(&insts[0], "udp://10.144.144.3:11010".parse().unwrap()).await; - assert_try_direct_connect_err( - &insts[0], - TcpTunnelConnector::new("tcp://10.1.2.3:11010".parse().unwrap()), - ) - .await; + assert_peer_admission_blocked(&insts[0], "tcp://10.1.2.3:11010".parse().unwrap()).await; - assert_try_direct_connect_err( - &insts[0], - UdpTunnelConnector::new("udp://10.1.2.3:11010".parse().unwrap()), - ) - .await; + assert_peer_admission_blocked(&insts[0], "udp://10.1.2.3:11010".parse().unwrap()).await; drop_insts(insts).await; @@ -2561,11 +2483,7 @@ pub async fn acl_rule_test_inbound( #[values(true, false)] enable_kcp_proxy: bool, #[values(true, false)] enable_quic_proxy: bool, ) { - use crate::tunnel::{ - common::tests::_tunnel_pingpong_netns_with_timeout, - tcp::{TcpTunnelConnector, TcpTunnelListener}, - udp::{UdpTunnelConnector, UdpTunnelListener}, - }; + use crate::tunnel::common::tests::_tunnel_pingpong_netns_with_timeout; use rand::Rng; let insts = init_three_node_ex( "udp", @@ -2637,25 +2555,22 @@ pub async fn acl_rule_test_inbound( let acl_toml = toml::to_string(&acl).unwrap(); println!("ACL TOML: {}", acl_toml); - insts[2] - .get_global_ctx() - .get_acl_filter() - .reload_rules(Some(&acl)); + reload_instance_acl(&insts[2], Some(&acl)).await; // TCP 测试部分 { // 2. 在 inst2 上监听 8080 和 8081 - let listener_8080 = TcpTunnelListener::new("tcp://0.0.0.0:8080".parse().unwrap()); - let listener_8081 = TcpTunnelListener::new("tcp://0.0.0.0:8081".parse().unwrap()); - let listener_8082 = TcpTunnelListener::new("tcp://0.0.0.0:8082".parse().unwrap()); + let listener_8080 = core_tcp_listener("tcp://0.0.0.0:8080".parse().unwrap()); + let listener_8081 = core_tcp_listener("tcp://0.0.0.0:8081".parse().unwrap()); + let listener_8082 = core_tcp_listener("tcp://0.0.0.0:8082".parse().unwrap()); // 3. inst1 作为客户端,尝试连接 inst2 的 8080(应被拒绝)和 8081(应被允许) let connector_8080 = - TcpTunnelConnector::new(format!("tcp://{}:8080", "10.144.144.3").parse().unwrap()); + core_tcp_dialer(format!("tcp://{}:8080", "10.144.144.3").parse().unwrap()); let connector_8081 = - TcpTunnelConnector::new(format!("tcp://{}:8081", "10.144.144.3").parse().unwrap()); + core_tcp_dialer(format!("tcp://{}:8081", "10.144.144.3").parse().unwrap()); let connector_8082 = - TcpTunnelConnector::new(format!("tcp://{}:8082", "10.144.144.3").parse().unwrap()); + core_tcp_dialer(format!("tcp://{}:8082", "10.144.144.3").parse().unwrap()); // 4. 构造测试数据 let mut buf = vec![0; 32]; @@ -2699,21 +2614,21 @@ pub async fn acl_rule_test_inbound( assert!(result.is_err(), "TCP 连接 8082 应被 ACL 拦截,不能成功"); - let stats = insts[2].get_global_ctx().get_acl_filter().get_stats(); + let stats = insts[2].get_core_instance().acl_stats(); println!("stats: {:?}", stats); } // UDP 测试部分 { // 1. 在 inst2 上监听 UDP 8080 和 8081 - let listener_8080 = UdpTunnelListener::new("udp://0.0.0.0:8080".parse().unwrap()); - let listener_8081 = UdpTunnelListener::new("udp://0.0.0.0:8081".parse().unwrap()); + let listener_8080 = core_udp_listener("udp://0.0.0.0:8080".parse().unwrap()); + let listener_8081 = core_udp_listener("udp://0.0.0.0:8081".parse().unwrap()); // 2. inst1 作为客户端,尝试连接 inst2 的 8080(应被拒绝)和 8081(应被允许) let connector_8080 = - UdpTunnelConnector::new(format!("udp://{}:8080", "10.144.144.3").parse().unwrap()); + core_udp_dialer(format!("udp://{}:8080", "10.144.144.3").parse().unwrap()); let connector_8081 = - UdpTunnelConnector::new(format!("udp://{}:8081", "10.144.144.3").parse().unwrap()); + core_udp_dialer(format!("udp://{}:8081", "10.144.144.3").parse().unwrap()); // 3. 构造测试数据 let mut buf = vec![0; 32]; @@ -2744,15 +2659,12 @@ pub async fn acl_rule_test_inbound( assert!(result.is_err(), "UDP 连接 8080 应被 ACL 拦截,不能成功"); - let stats = insts[2].get_global_ctx().get_acl_filter().get_stats(); + let stats = insts[2].get_core_instance().acl_stats(); println!("stats: {}", stats); } // remove acl, 8080 should succ - insts[2] - .get_global_ctx() - .get_acl_filter() - .reload_rules(None); + reload_instance_acl(&insts[2], None).await; drop_insts(insts).await; } @@ -2764,11 +2676,7 @@ pub async fn acl_rule_test_subnet_proxy( #[values(true, false)] enable_kcp_proxy: bool, #[values(true, false)] enable_quic_proxy: bool, ) { - use crate::tunnel::{ - common::tests::_tunnel_pingpong_netns_with_timeout, - tcp::{TcpTunnelConnector, TcpTunnelListener}, - udp::{UdpTunnelConnector, UdpTunnelListener}, - }; + use crate::tunnel::common::tests::_tunnel_pingpong_netns_with_timeout; use rand::Rng; let insts = init_three_node_ex( @@ -2792,7 +2700,7 @@ pub async fn acl_rule_test_subnet_proxy( // 等待代理路由出现 wait_proxy_route_appear( - &insts[0].get_peer_manager(), + &insts[0].get_core_instance(), "10.144.144.3/24", insts[2].peer_id(), "10.1.2.0/24", @@ -2861,22 +2769,19 @@ pub async fn acl_rule_test_subnet_proxy( acl.acl_v1 = Some(acl_v1); // 在 inst3 上应用 ACL 规则 - insts[2] - .get_global_ctx() - .get_acl_filter() - .reload_rules(Some(&acl)); + reload_instance_acl(&insts[2], Some(&acl)).await; // TCP 测试部分 - 测试子网代理的 ACL 规则 { // 在 net_d (10.1.2.4) 上监听多个端口 - let listener_8080 = TcpTunnelListener::new("tcp://0.0.0.0:8080".parse().unwrap()); - let listener_8081 = TcpTunnelListener::new("tcp://0.0.0.0:8081".parse().unwrap()); - let listener_8082 = TcpTunnelListener::new("tcp://0.0.0.0:8082".parse().unwrap()); + let listener_8080 = core_tcp_listener("tcp://0.0.0.0:8080".parse().unwrap()); + let listener_8081 = core_tcp_listener("tcp://0.0.0.0:8081".parse().unwrap()); + let listener_8082 = core_tcp_listener("tcp://0.0.0.0:8082".parse().unwrap()); // 从 inst1 (net_a) 连接到子网代理 - let connector_8080 = TcpTunnelConnector::new("tcp://10.1.2.4:8080".parse().unwrap()); - let connector_8081 = TcpTunnelConnector::new("tcp://10.1.2.4:8081".parse().unwrap()); - let connector_8082 = TcpTunnelConnector::new("tcp://10.1.2.4:8082".parse().unwrap()); + let connector_8080 = core_tcp_dialer("tcp://10.1.2.4:8080".parse().unwrap()); + let connector_8081 = core_tcp_dialer("tcp://10.1.2.4:8081".parse().unwrap()); + let connector_8082 = core_tcp_dialer("tcp://10.1.2.4:8082".parse().unwrap()); let mut buf = vec![0; 32]; rand::thread_rng().fill(&mut buf[..]); @@ -2925,17 +2830,17 @@ pub async fn acl_rule_test_subnet_proxy( "TCP 连接子网代理 8081 应被 ACL 拦截,不能成功" ); - let stats = insts[2].get_global_ctx().get_acl_filter().get_stats(); + let stats = insts[2].get_core_instance().acl_stats(); println!("ACL stats after TCP tests: {:?}", stats); } // UDP 测试部分 - 测试子网代理的 ACL 规则 { - let listener_8080 = UdpTunnelListener::new("udp://0.0.0.0:8080".parse().unwrap()); - let listener_8082 = UdpTunnelListener::new("udp://0.0.0.0:8082".parse().unwrap()); + let listener_8080 = core_udp_listener("udp://0.0.0.0:8080".parse().unwrap()); + let listener_8082 = core_udp_listener("udp://0.0.0.0:8082".parse().unwrap()); - let connector_8080 = UdpTunnelConnector::new("udp://10.1.2.4:8080".parse().unwrap()); - let connector_8082 = UdpTunnelConnector::new("udp://10.1.2.4:8082".parse().unwrap()); + let connector_8080 = core_udp_dialer("udp://10.1.2.4:8080".parse().unwrap()); + let connector_8082 = core_udp_dialer("udp://10.1.2.4:8082".parse().unwrap()); let mut buf = vec![0; 32]; rand::thread_rng().fill(&mut buf[..]); @@ -2963,7 +2868,7 @@ pub async fn acl_rule_test_subnet_proxy( ) .await; - let stats = insts[2].get_global_ctx().get_acl_filter().get_stats(); + let stats = insts[2].get_core_instance().acl_stats(); println!("ACL stats after UDP tests: {}", stats); assert!( @@ -2981,10 +2886,7 @@ pub async fn acl_rule_test_subnet_proxy( .unwrap_err(); // 移除 ACL 规则 - insts[2] - .get_global_ctx() - .get_acl_filter() - .reload_rules(None); + reload_instance_acl(&insts[2], None).await; // 验证移除 ACL 后,ICMP 可以正常工作 wait_for_condition( @@ -3018,13 +2920,12 @@ where } async fn wait_route_cost(inst: &Instance, peer_id: u32, cost: i32, timeout: Duration) { - let peer_manager = inst.get_peer_manager(); + let core = inst.get_core_instance(); wait_for_condition( move || { - let peer_manager = peer_manager.clone(); + let core = core.clone(); async move { - peer_manager - .list_routes() + core.route_snapshots() .await .iter() .any(|route| route.peer_id == peer_id && route.cost == cost) @@ -3043,8 +2944,6 @@ pub async fn p2p_only_test( #[values(true, false)] enable_kcp_proxy: bool, #[values(true, false)] enable_quic_proxy: bool, ) { - use crate::peers::tests::wait_route_appear_with_cost; - let insts = init_three_node_ex( "udp", |cfg| { @@ -3067,18 +2966,8 @@ pub async fn p2p_only_test( .await; if has_p2p_conn { - insts[2] - .get_conn_manager() - .add_connector(RingTunnelConnector::new( - format!("ring://{}", insts[0].id()).parse().unwrap(), - )); - wait_route_appear_with_cost( - insts[2].get_peer_manager(), - insts[0].get_peer_manager().my_peer_id(), - Some(1), - ) - .await - .unwrap(); + insts[2].add_connector_url(insts[0].ring_listener_url()); + wait_route_cost(&insts[2], insts[0].peer_id(), 1, Duration::from_secs(5)).await; } let target_ip = "10.1.2.4"; @@ -3123,12 +3012,7 @@ pub async fn acl_group_base_test( #[values(true, false)] enable_kcp_proxy: bool, #[values(true, false)] enable_quic_proxy: bool, ) { - use crate::tunnel::{ - TunnelConnector, TunnelListener, - common::tests::_tunnel_pingpong_netns_with_timeout, - tcp::{TcpTunnelConnector, TcpTunnelListener}, - udp::{UdpTunnelConnector, UdpTunnelListener}, - }; + use crate::tunnel::common::tests::_tunnel_pingpong_netns_with_timeout; use rand::Rng; // 构造 ACL 配置,包含组信息 @@ -3240,24 +3124,24 @@ pub async fn acl_group_base_test( println!("Testing group-based ACL rules..."); - let make_listener = |port: u16| -> Box { + let make_listener = |port: u16| -> Box> + Sync> { match protocol { - "tcp" => Box::new(TcpTunnelListener::new( + "tcp" => Box::new(core_tcp_listener( format!("tcp://0.0.0.0:{}", port).parse().unwrap(), )), - "udp" => Box::new(UdpTunnelListener::new( + "udp" => Box::new(core_udp_listener( format!("udp://0.0.0.0:{}", port).parse().unwrap(), )), _ => panic!("unsupported protocol: {}", protocol), } }; - let make_connector = |port: u16| -> Box { + let make_connector = |port: u16| -> Box { match protocol { - "tcp" => Box::new(TcpTunnelConnector::new( + "tcp" => Box::new(core_tcp_dialer( format!("tcp://10.144.144.3:{}", port).parse().unwrap(), )), - "udp" => Box::new(UdpTunnelConnector::new( + "udp" => Box::new(core_udp_dialer( format!("udp://10.144.144.3:{}", port).parse().unwrap(), )), _ => panic!("unsupported protocol: {}", protocol), @@ -3348,7 +3232,7 @@ pub async fn acl_group_base_test( protocol ); - let stats = insts[2].get_global_ctx().get_acl_filter().get_stats(); + let stats = insts[2].get_core_instance().acl_stats(); println!("ACL stats after group {} tests: {:?}", protocol, stats); println!("✓ All group-based ACL tests completed successfully"); @@ -3372,9 +3256,10 @@ pub async fn lazy_p2p_builds_direct_connection_on_demand() { let inst3_peer_id = insts[2].peer_id(); assert!( !insts[0] - .get_peer_manager() - .get_peer_map() - .has_peer(inst3_peer_id), + .get_core_instance() + .connected_peers() + .await + .contains(&inst3_peer_id), "inst1 should not proactively connect to inst3 when lazy_p2p is enabled" ); wait_route_cost(&insts[0], inst3_peer_id, 2, Duration::from_secs(5)).await; @@ -3387,9 +3272,10 @@ pub async fn lazy_p2p_builds_direct_connection_on_demand() { wait_for_condition( || async { insts[0] - .get_peer_manager() - .get_peer_map() - .has_peer(inst3_peer_id) + .get_core_instance() + .connected_peers() + .await + .contains(&inst3_peer_id) }, Duration::from_secs(10), ) @@ -3420,9 +3306,10 @@ pub async fn need_p2p_overrides_lazy_p2p() { wait_for_condition( || async { insts[0] - .get_peer_manager() - .get_peer_map() - .has_peer(inst3_peer_id) + .get_core_instance() + .connected_peers() + .await + .contains(&inst3_peer_id) }, Duration::from_secs(10), ) @@ -3453,9 +3340,10 @@ pub async fn disable_p2p_still_connects_to_need_p2p_peers() { wait_for_condition( || async { insts[0] - .get_peer_manager() - .get_peer_map() - .has_peer(inst3_peer_id) + .get_core_instance() + .connected_peers() + .await + .contains(&inst3_peer_id) }, Duration::from_secs(10), ) @@ -3489,9 +3377,10 @@ pub async fn ordinary_nodes_do_not_proactively_connect_to_disable_p2p_peers() { assert!( !insts[0] - .get_peer_manager() - .get_peer_map() - .has_peer(inst3_peer_id), + .get_core_instance() + .connected_peers() + .await + .contains(&inst3_peer_id), "ordinary nodes should not proactively establish p2p with disable-p2p peers" ); wait_route_cost(&insts[0], inst3_peer_id, 2, Duration::from_secs(3)).await; @@ -3523,9 +3412,10 @@ pub async fn lazy_p2p_warms_up_before_p2p_only_send() { wait_for_condition( || async { insts[0] - .get_peer_manager() - .get_peer_map() - .has_peer(inst3_peer_id) + .get_core_instance() + .connected_peers() + .await + .contains(&inst3_peer_id) }, Duration::from_secs(10), ) @@ -3549,12 +3439,7 @@ pub async fn acl_group_self_test( #[values(true, false)] enable_kcp_proxy: bool, #[values(true, false)] enable_quic_proxy: bool, ) { - use crate::tunnel::{ - TunnelConnector, TunnelListener, - common::tests::_tunnel_pingpong_netns_with_timeout, - tcp::{TcpTunnelConnector, TcpTunnelListener}, - udp::{UdpTunnelConnector, UdpTunnelListener}, - }; + use crate::tunnel::common::tests::_tunnel_pingpong_netns_with_timeout; use rand::Rng; // 构造 ACL 配置,包含组信息 @@ -3637,24 +3522,24 @@ pub async fn acl_group_self_test( println!("Testing group-based ACL rules..."); - let make_listener = |port: u16| -> Box { + let make_listener = |port: u16| -> Box> + Sync> { match protocol { - "tcp" => Box::new(TcpTunnelListener::new( + "tcp" => Box::new(core_tcp_listener( format!("tcp://0.0.0.0:{}", port).parse().unwrap(), )), - "udp" => Box::new(UdpTunnelListener::new( + "udp" => Box::new(core_udp_listener( format!("udp://0.0.0.0:{}", port).parse().unwrap(), )), _ => panic!("unsupported protocol: {}", protocol), } }; - let make_connector = |port: u16| -> Box { + let make_connector = |port: u16| -> Box { match protocol { - "tcp" => Box::new(TcpTunnelConnector::new( + "tcp" => Box::new(core_tcp_dialer( format!("tcp://10.144.144.3:{}", port).parse().unwrap(), )), - "udp" => Box::new(UdpTunnelConnector::new( + "udp" => Box::new(core_udp_dialer( format!("udp://10.144.144.3:{}", port).parse().unwrap(), )), _ => panic!("unsupported protocol: {}", protocol), @@ -3705,7 +3590,7 @@ pub async fn acl_group_self_test( protocol ); - let stats = insts[2].get_global_ctx().get_acl_filter().get_stats(); + let stats = insts[2].get_core_instance().acl_stats(); println!("ACL stats after group {} tests: {:?}", protocol, stats); println!("✓ All group-based ACL tests completed successfully"); @@ -3743,39 +3628,33 @@ pub async fn whitelist_test( ) .await; - use crate::tunnel::{ - TunnelConnector, TunnelListener, - common::tests::_tunnel_pingpong_netns_with_timeout, - tcp::{TcpTunnelConnector, TcpTunnelListener}, - udp::{UdpTunnelConnector, UdpTunnelListener}, - }; + use crate::tunnel::common::tests::_tunnel_pingpong_netns_with_timeout; use rand::Rng; let make_listener = - |protocol: &str, port: u16| -> Box { + |protocol: &str, port: u16| -> Box> + Sync> { match protocol { - "tcp" => Box::new(TcpTunnelListener::new( + "tcp" => Box::new(core_tcp_listener( format!("tcp://0.0.0.0:{}", port).parse().unwrap(), )), - "udp" => Box::new(UdpTunnelListener::new( + "udp" => Box::new(core_udp_listener( format!("udp://0.0.0.0:{}", port).parse().unwrap(), )), _ => panic!("unsupported protocol: {}", protocol), } }; - let make_connector = - |protocol: &str, port: u16| -> Box { - match protocol { - "tcp" => Box::new(TcpTunnelConnector::new( - format!("tcp://10.144.144.3:{}", port).parse().unwrap(), - )), - "udp" => Box::new(UdpTunnelConnector::new( - format!("udp://10.144.144.3:{}", port).parse().unwrap(), - )), - _ => panic!("unsupported protocol: {}", protocol), - } - }; + let make_connector = |protocol: &str, port: u16| -> Box { + match protocol { + "tcp" => Box::new(core_tcp_dialer( + format!("tcp://10.144.144.3:{}", port).parse().unwrap(), + )), + "udp" => Box::new(core_udp_dialer( + format!("udp://10.144.144.3:{}", port).parse().unwrap(), + )), + _ => panic!("unsupported protocol: {}", protocol), + } + }; let mut buf = vec![0; 32]; rand::thread_rng().fill(&mut buf[..]); @@ -3845,13 +3724,13 @@ pub async fn config_patch_test() { check_route( "10.144.144.2/24", insts[1].peer_id(), - insts[0].get_peer_manager().list_routes().await, + insts[0].get_core_instance().route_snapshots().await, ); check_route( "10.144.144.3/24", insts[2].peer_id(), - insts[0].get_peer_manager().list_routes().await, + insts[0].get_core_instance().route_snapshots().await, ); // 测试1: 修改hostname、ip、子网代理 @@ -3877,7 +3756,7 @@ pub async fn config_patch_test() { ); tokio::time::sleep(Duration::from_secs(1)).await; check_route_ex( - insts[0].get_peer_manager().list_routes().await, + insts[0].get_core_instance().route_snapshots().await, insts[1].peer_id(), |r| { assert_eq!(r.hostname, "new_inst1"); @@ -3935,14 +3814,18 @@ pub async fn config_patch_test() { ); assert!( insts[1] - .get_global_ctx() - .get_feature_flags() + .get_core_instance() + .node_snapshot() + .await + .feature_flags .ipv6_public_addr_provider ); assert_eq!( insts[1] - .get_global_ctx() - .get_advertised_ipv6_public_addr_prefix(), + .get_core_instance() + .node_snapshot() + .await + .ipv6_public_addr_prefix, Some(public_prefix.parse().unwrap()) ); @@ -3966,8 +3849,8 @@ pub async fn config_patch_test() { let mut buf = vec![0; 32]; rand::thread_rng().fill(&mut buf[..]); - let tcp_listener = TcpTunnelListener::new("tcp://0.0.0.0:23457".parse().unwrap()); - let tcp_connector = TcpTunnelConnector::new("tcp://127.0.0.1:23458".parse().unwrap()); + let tcp_listener = core_tcp_listener("tcp://0.0.0.0:23457".parse().unwrap()); + let tcp_connector = core_tcp_dialer("tcp://127.0.0.1:23458".parse().unwrap()); let result = _tunnel_pingpong_netns_with_timeout( tcp_listener, tcp_connector, @@ -4003,13 +3886,15 @@ pub async fn config_patch_disable_relay_data_test() { assert!(!insts[1].get_global_ctx().get_flags().disable_relay_data); assert!( !insts[1] - .get_global_ctx() - .get_feature_flags() + .get_core_instance() + .node_snapshot() + .await + .feature_flags .avoid_relay_data ); check_route_ex( - insts[0].get_peer_manager().list_routes().await, + insts[0].get_core_instance().route_snapshots().await, dst_peer_id, |route| { assert_eq!(route.next_hop_peer_id, relay_peer_id); @@ -4042,16 +3927,18 @@ pub async fn config_patch_disable_relay_data_test() { ); assert!( insts[1] - .get_global_ctx() - .get_feature_flags() + .get_core_instance() + .node_snapshot() + .await + .feature_flags .avoid_relay_data ); wait_for_condition( || { - let peer_mgr = insts[0].get_peer_manager().clone(); + let core = insts[0].get_core_instance(); async move { - peer_mgr.list_routes().await.iter().any(|route| { + core.route_snapshots().await.iter().any(|route| { route.peer_id == relay_peer_id && route .feature_flag @@ -4066,7 +3953,7 @@ pub async fn config_patch_disable_relay_data_test() { .await; check_route_ex( - insts[0].get_peer_manager().list_routes().await, + insts[0].get_core_instance().route_snapshots().await, dst_peer_id, |route| { assert_eq!(route.next_hop_peer_id, relay_peer_id); @@ -4097,16 +3984,18 @@ pub async fn config_patch_disable_relay_data_test() { ); assert!( !insts[1] - .get_global_ctx() - .get_feature_flags() + .get_core_instance() + .node_snapshot() + .await + .feature_flags .avoid_relay_data ); wait_for_condition( || { - let peer_mgr = insts[0].get_peer_manager().clone(); + let core = insts[0].get_core_instance(); async move { - peer_mgr.list_routes().await.iter().any(|route| { + core.route_snapshots().await.iter().any(|route| { route.peer_id == relay_peer_id && route .feature_flag @@ -4156,8 +4045,6 @@ pub fn generate_secure_mode_config() -> SecureModeConfig { #[tokio::test] #[serial_test::serial] pub async fn relay_peer_e2e_encryption(#[values("tcp", "udp")] proto: &str) { - use crate::peers::route_trait::NextHopPolicy; - let insts = init_three_node_ex( proto, |cfg| { @@ -4191,7 +4078,7 @@ pub async fn relay_peer_e2e_encryption(#[values("tcp", "udp")] proto: &str) { // Wait for routes to be established wait_for_condition( || async { - let routes = insts[0].get_peer_manager().list_routes().await; + let routes = insts[0].get_core_instance().route_snapshots().await; routes.len() == 2 }, Duration::from_secs(10), @@ -4200,10 +4087,12 @@ pub async fn relay_peer_e2e_encryption(#[values("tcp", "udp")] proto: &str) { // Verify inst1 sees inst3 via inst2 (non-direct path) let next_hop_to_inst3 = insts[0] - .get_peer_manager() - .get_peer_map() - .get_gateway_peer_id(inst3_peer_id, NextHopPolicy::LeastHop) - .await; + .get_core_instance() + .route_snapshots() + .await + .into_iter() + .find(|route| route.peer_id == inst3_peer_id) + .map(|route| route.next_hop_peer_id); println!("Next hop from inst1 to inst3: {:?}", next_hop_to_inst3); assert_eq!( next_hop_to_inst3, @@ -4214,34 +4103,30 @@ pub async fn relay_peer_e2e_encryption(#[values("tcp", "udp")] proto: &str) { // Verify inst1 has no direct connection to inst3 assert!( !insts[0] - .get_peer_manager() - .get_peer_map() - .has_peer(inst3_peer_id), + .get_core_instance() + .connected_peers() + .await + .contains(&inst3_peer_id), "inst1 should NOT have direct connection to inst3" ); // Check if noise_static_pubkey is available for relay handshake - let route_info_inst3 = insts[0] - .get_peer_manager() - .get_peer_map() - .get_route_peer_info(inst3_peer_id) + let route_has_static_key = insts[0] + .get_core_instance() + .relay_route_has_static_key_for_test(inst3_peer_id) .await; println!( - "Route info for inst3 on inst1: noise_static_pubkey len = {:?}", - route_info_inst3 - .as_ref() - .map(|i| i.noise_static_pubkey.len()) + "Route info for inst3 on inst1 has a relay static key: {}", + route_has_static_key ); // Wait until relay route info includes inst3 static pubkey for IK handshake. wait_for_condition( || async { insts[0] - .get_peer_manager() - .get_peer_map() - .get_route_peer_info(inst3_peer_id) + .get_core_instance() + .relay_route_has_static_key_for_test(inst3_peer_id) .await - .is_some_and(|info| !info.noise_static_pubkey.is_empty()) }, Duration::from_secs(10), ) @@ -4256,13 +4141,16 @@ pub async fn relay_peer_e2e_encryption(#[values("tcp", "udp")] proto: &str) { ); // Verify relay sessions are established - let relay_map_1 = insts[0].get_peer_manager().get_relay_peer_map(); - let relay_map_3 = insts[2].get_peer_manager().get_relay_peer_map(); + let relay_1 = insts[0] + .get_core_instance() + .relay_session_snapshot_for_test(inst3_peer_id); + let relay_3 = insts[2] + .get_core_instance() + .relay_session_snapshot_for_test(inst1_peer_id); println!( "Relay states after ping: inst1->inst3: {}, inst3->inst1: {}", - relay_map_1.has_state(inst3_peer_id), - relay_map_3.has_state(inst1_peer_id) + relay_1.has_state, relay_3.has_state ); // Test bidirectional connectivity @@ -4300,7 +4188,7 @@ pub async fn relay_peer_e2e_encryption_udp() { wait_for_condition( || async { - let routes = insts[0].get_peer_manager().list_routes().await; + let routes = insts[0].get_core_instance().route_snapshots().await; routes.len() == 2 }, Duration::from_secs(10), @@ -4322,36 +4210,17 @@ pub async fn relay_peer_e2e_encryption_udp() { wait_for_condition( || async { - insts[0] - .get_global_ctx() - .stats_manager() - .get_metric(MetricName::TrafficBytesTx, &tx_labels) - .is_none() - && insts[0] - .get_global_ctx() - .stats_manager() - .get_metric(MetricName::TrafficPacketsTx, &tx_labels) - .is_none() - && insts[0] - .get_global_ctx() - .stats_manager() - .get_metric(MetricName::TrafficBytesTx, &total_labels) - .is_some_and(|metric| metric.value > 0) - && insts[0] - .get_global_ctx() - .stats_manager() - .get_metric(MetricName::TrafficPacketsTx, &total_labels) - .is_some_and(|metric| metric.value > 0) - && insts[0] - .get_global_ctx() - .stats_manager() - .get_metric(MetricName::TrafficBytesTxByInstance, &tx_labels) - .is_some_and(|metric| metric.value > 0) - && insts[0] - .get_global_ctx() - .stats_manager() - .get_metric(MetricName::TrafficPacketsTxByInstance, &tx_labels) - .is_some_and(|metric| metric.value > 0) + let metrics = insts[0].get_core_instance().metric_snapshots(); + metric_value(&metrics, MetricName::TrafficBytesTx, &tx_labels).is_none() + && metric_value(&metrics, MetricName::TrafficPacketsTx, &tx_labels).is_none() + && metric_value(&metrics, MetricName::TrafficBytesTx, &total_labels) + .is_some_and(|value| value > 0) + && metric_value(&metrics, MetricName::TrafficPacketsTx, &total_labels) + .is_some_and(|value| value > 0) + && metric_value(&metrics, MetricName::TrafficBytesTxByInstance, &tx_labels) + .is_some_and(|value| value > 0) + && metric_value(&metrics, MetricName::TrafficPacketsTxByInstance, &tx_labels) + .is_some_and(|value| value > 0) }, Duration::from_secs(10), ) @@ -4359,36 +4228,17 @@ pub async fn relay_peer_e2e_encryption_udp() { wait_for_condition( || async { - insts[2] - .get_global_ctx() - .stats_manager() - .get_metric(MetricName::TrafficBytesRx, &rx_labels) - .is_none() - && insts[2] - .get_global_ctx() - .stats_manager() - .get_metric(MetricName::TrafficPacketsRx, &rx_labels) - .is_none() - && insts[2] - .get_global_ctx() - .stats_manager() - .get_metric(MetricName::TrafficBytesRx, &total_labels) - .is_some_and(|metric| metric.value > 0) - && insts[2] - .get_global_ctx() - .stats_manager() - .get_metric(MetricName::TrafficPacketsRx, &total_labels) - .is_some_and(|metric| metric.value > 0) - && insts[2] - .get_global_ctx() - .stats_manager() - .get_metric(MetricName::TrafficBytesRxByInstance, &rx_labels) - .is_some_and(|metric| metric.value > 0) - && insts[2] - .get_global_ctx() - .stats_manager() - .get_metric(MetricName::TrafficPacketsRxByInstance, &rx_labels) - .is_some_and(|metric| metric.value > 0) + let metrics = insts[2].get_core_instance().metric_snapshots(); + metric_value(&metrics, MetricName::TrafficBytesRx, &rx_labels).is_none() + && metric_value(&metrics, MetricName::TrafficPacketsRx, &rx_labels).is_none() + && metric_value(&metrics, MetricName::TrafficBytesRx, &total_labels) + .is_some_and(|value| value > 0) + && metric_value(&metrics, MetricName::TrafficPacketsRx, &total_labels) + .is_some_and(|value| value > 0) + && metric_value(&metrics, MetricName::TrafficBytesRxByInstance, &rx_labels) + .is_some_and(|value| value > 0) + && metric_value(&metrics, MetricName::TrafficPacketsRxByInstance, &rx_labels) + .is_some_and(|value| value > 0) }, Duration::from_secs(10), ) @@ -4401,8 +4251,6 @@ pub async fn relay_peer_e2e_encryption_udp() { #[tokio::test] #[serial_test::serial] pub async fn relay_peer_session_cleanup() { - use crate::peers::route_trait::NextHopPolicy; - let mut insts = init_three_node_ex( "tcp", |cfg| { @@ -4415,7 +4263,7 @@ pub async fn relay_peer_session_cleanup() { let inst2_peer_id = insts[1].peer_id(); let inst3_peer_id = insts[2].peer_id(); - let relay_map_1 = insts[0].get_peer_manager().get_relay_peer_map(); + let core_1 = insts[0].get_core_instance(); wait_for_condition( || async { ping_test("net_a", "10.144.144.3", None).await }, @@ -4424,16 +4272,21 @@ pub async fn relay_peer_session_cleanup() { .await; wait_for_condition( - || async { relay_map_1.has_state(inst3_peer_id) && relay_map_1.has_session(inst3_peer_id) }, + || async { + let relay = core_1.relay_session_snapshot_for_test(inst3_peer_id); + relay.has_state && relay.has_session + }, Duration::from_secs(3), ) .await; let next_hop = insts[0] - .get_peer_manager() - .get_peer_map() - .get_gateway_peer_id(inst3_peer_id, NextHopPolicy::LeastHop) - .await; + .get_core_instance() + .route_snapshots() + .await + .into_iter() + .find(|route| route.peer_id == inst3_peer_id) + .map(|route| route.next_hop_peer_id); assert_eq!(next_hop, Some(inst2_peer_id)); let mut inst2 = insts.remove(1); @@ -4442,26 +4295,32 @@ pub async fn relay_peer_session_cleanup() { wait_for_condition( || async { - let routes = insts[0].get_peer_manager().list_routes().await; + let routes = insts[0].get_core_instance().route_snapshots().await; !routes.iter().any(|r| r.peer_id == inst3_peer_id) }, Duration::from_secs(6), ) .await; - relay_map_1.evict_idle_sessions(Duration::from_millis(0)); - assert!(!relay_map_1.has_state(inst3_peer_id)); + core_1.evict_idle_relay_sessions_for_test(Duration::from_millis(0)); + assert!( + !core_1 + .relay_session_snapshot_for_test(inst3_peer_id) + .has_state + ); - insts[0] - .get_peer_manager() - .get_peer_session_store() - .evict_unused_sessions_idle(Duration::from_millis(0)); + core_1.evict_unused_peer_sessions_for_test(Duration::from_millis(0)); wait_for_condition( - || async { !relay_map_1.has_session(inst3_peer_id) }, + || async { + !core_1 + .relay_session_snapshot_for_test(inst3_peer_id) + .has_session + }, Duration::from_secs(1), ) .await; + drop(core_1); drop_insts(insts).await; } diff --git a/easytier/src/tests/upnp_test.rs b/easytier/src/tests/upnp_test.rs index 8e12c72c..1440555d 100644 --- a/easytier/src/tests/upnp_test.rs +++ b/easytier/src/tests/upnp_test.rs @@ -8,42 +8,32 @@ use std::{ }; use anyhow::{Context, anyhow, bail}; +use easytier_core::{ + connectivity::stun::{StunInfoProvider, StunSocketMapper}, + process_runtime::CoreProcessRuntime, +}; use igd_next::{ GetGenericPortMappingEntryError, PortMappingEntry, PortMappingProtocol, SearchOptions, aio::tokio::search_gateway, }; use tempfile::TempDir; -use tokio::net::UdpSocket; -use super::{create_netns, del_netns, drop_insts, get_host_veth_name, ping_test}; +use super::{ + InstanceTestExt as _, create_netns, del_netns, drop_insts, get_host_veth_name, ping_test, +}; use crate::{ common::{ config::{ConfigLoader, TomlConfigLoader}, error::Error, - global_ctx::{GlobalCtx, GlobalCtxEvent}, + global_ctx::GlobalCtxEvent, netns::NetNS, - stun::{MockStunInfoCollector, StunInfoCollectorTrait}, - }, - connector::udp_hole_punch::{UdpHolePunchConnector, common::UdpHolePunchListener}, - instance::instance::Instance, - peers::{ - create_packet_recv_chan, - peer_manager::{PeerManager, RouteAlgoType}, - tests::{connect_peer_manager, wait_route_appear, wait_route_appear_with_cost}, + stun::MockStunInfoCollector, }, + instance::test_instance::TestInstance as Instance, proto::common::{NatType, StunInfo}, - tunnel::{common::tests::wait_for_condition, ring::RingTunnelConnector}, + tunnel::common::tests::wait_for_condition, }; -const TEST_NS_A: &str = "upnp_a"; -const TEST_NS_C: &str = "upnp_c"; -const TEST_BRIDGE: &str = "br_upnp"; -const TEST_WAN_IF: &str = "upnp_wan0"; -const TEST_GATEWAY_IP: Ipv4Addr = Ipv4Addr::new(172, 31, 255, 1); -const TEST_CLIENT_A_IP: Ipv4Addr = Ipv4Addr::new(172, 31, 255, 2); -const TEST_CLIENT_C_IP: Ipv4Addr = Ipv4Addr::new(172, 31, 255, 3); -const TEST_EXTERNAL_IP: Ipv4Addr = Ipv4Addr::new(11, 22, 33, 44); -const TEST_CONTROL_PORT: u16 = 5000; const TEST_IGD_DESCRIPTION: &str = "EasyTier udp hole punch"; const DUAL_NS_A: &str = "upnp2_a"; @@ -66,90 +56,6 @@ const DUAL_WAN_IF_C_PEER: &str = "upnp2_wan_c_p"; const DUAL_GW_NS_A: &str = "upnp2_gw_a"; const DUAL_GW_NS_C: &str = "upnp2_gw_c"; -struct UpnpIntegrationEnv { - _tempdir: TempDir, - child: Option, -} - -impl UpnpIntegrationEnv { - async fn new() -> anyhow::Result { - cleanup_miniupnpd_processes(); - cleanup_test_net(); - create_test_net()?; - - let tempdir = tempfile::tempdir().context("create miniupnpd tempdir")?; - let conf_path = tempdir.path().join("miniupnpd.conf"); - let leases_path = tempdir.path().join("miniupnpd.leases"); - std::fs::write(&leases_path, "").context("create miniupnpd lease file")?; - std::fs::write( - &conf_path, - format!( - "\ -ext_ifname={TEST_WAN_IF} -listening_ip={TEST_BRIDGE} -port={TEST_CONTROL_PORT} -enable_natpmp=no -enable_upnp=yes -secure_mode=no -system_uptime=yes -lease_file={} -ext_ip={} -friendly_name=EasyTier Test IGD -model_name=EasyTier Test -serial=12345678 -uuid=9f0c5a3a-c4f0-4f1e-b4df-8a8c7b1e2d00 -allow 1024-65535 172.31.255.0/24 1024-65535 -deny 0-65535 0.0.0.0/0 0-65535 -", - leases_path.display(), - TEST_EXTERNAL_IP - ), - ) - .context("write miniupnpd config")?; - - let miniupnpd_bin = find_miniupnpd_bin()?; - let child = Command::new(miniupnpd_bin) - .args(["-d", "-f"]) - .arg(&conf_path) - .stdin(Stdio::null()) - .stdout(Stdio::inherit()) - .stderr(Stdio::inherit()) - .spawn() - .context("spawn miniupnpd")?; - - let env = Self { - _tempdir: tempdir, - child: Some(child), - }; - env.wait_ready().await?; - Ok(env) - } - - async fn wait_ready(&self) -> anyhow::Result<()> { - wait_for_condition( - || async { - tokio::net::TcpStream::connect((TEST_GATEWAY_IP, TEST_CONTROL_PORT)) - .await - .is_ok() - }, - Duration::from_secs(10), - ) - .await; - Ok(()) - } -} - -impl Drop for UpnpIntegrationEnv { - fn drop(&mut self) { - if let Some(mut child) = self.child.take() { - let _ = child.kill(); - let _ = child.wait(); - } - cleanup_miniupnpd_processes(); - cleanup_test_net(); - } -} - struct DualGatewayUpnpIntegrationEnv { _tempdir: TempDir, children: Vec, @@ -294,7 +200,7 @@ struct GatewayBackedStunCollector { } #[async_trait::async_trait] -impl StunInfoCollectorTrait for GatewayBackedStunCollector { +impl StunInfoProvider for GatewayBackedStunCollector { fn get_stun_info(&self) -> StunInfo { StunInfo { udp_nat_type: NatType::PortRestricted as i32, @@ -306,104 +212,32 @@ impl StunInfoCollectorTrait for GatewayBackedStunCollector { } } - async fn get_udp_port_mapping(&self, local_port: u16) -> Result { - query_udp_mapping(self.netns, self.external_ip, self.client_ip, local_port).await + async fn get_udp_port_mapping(&self, local_port: u16) -> anyhow::Result { + Ok(query_udp_mapping(self.netns, self.external_ip, self.client_ip, local_port).await?) } + async fn get_tcp_port_mapping(&self, local_port: u16) -> anyhow::Result { + Ok(SocketAddr::new(IpAddr::V4(self.external_ip), local_port)) + } + + fn update_stun_info(&self) {} +} + +#[async_trait::async_trait] +impl StunSocketMapper for GatewayBackedStunCollector { async fn get_udp_port_mapping_with_socket( &self, - udp: Arc, - ) -> Result { - query_udp_mapping( + udp: Arc, + ) -> anyhow::Result { + use easytier_core::socket::udp::VirtualUdpSocket as _; + Ok(query_udp_mapping( self.netns, self.external_ip, self.client_ip, udp.local_addr()?.port(), ) - .await + .await?) } - - async fn get_tcp_port_mapping(&self, local_port: u16) -> Result { - Ok(SocketAddr::new(IpAddr::V4(self.external_ip), local_port)) - } -} - -fn create_test_net() -> anyhow::Result<()> { - create_netns(TEST_NS_A, &format!("{TEST_CLIENT_A_IP}/24"), "fd10::2/64"); - create_netns(TEST_NS_C, &format!("{TEST_CLIENT_C_IP}/24"), "fd10::3/64"); - run_cmd( - "ip", - &["link", "add", "name", TEST_BRIDGE, "type", "bridge"], - )?; - for netns in [TEST_NS_A, TEST_NS_C] { - run_cmd( - "ip", - &[ - "link", - "set", - get_host_veth_name(netns), - "master", - TEST_BRIDGE, - ], - )?; - } - - run_cmd( - "ip", - &[ - "addr", - "add", - &format!("{TEST_GATEWAY_IP}/24"), - "dev", - TEST_BRIDGE, - ], - )?; - run_cmd("ip", &["link", "add", TEST_WAN_IF, "type", "dummy"])?; - run_cmd( - "ip", - &[ - "addr", - "add", - &format!("{TEST_EXTERNAL_IP}/24"), - "dev", - TEST_WAN_IF, - ], - )?; - run_cmd("ip", &["link", "set", TEST_WAN_IF, "up"])?; - run_cmd("ip", &["link", "set", TEST_BRIDGE, "up"])?; - setup_iptables_rules()?; - for (netns, guest_veth) in [(TEST_NS_A, "veth_upnp_a_g"), (TEST_NS_C, "veth_upnp_c_g")] { - run_cmd( - "ip", - &[ - "netns", - "exec", - netns, - "ip", - "route", - "add", - "default", - "via", - &TEST_GATEWAY_IP.to_string(), - "dev", - guest_veth, - ], - )?; - } - run_cmd("sysctl", &["-w", "net.ipv4.ip_forward=1"])?; - Ok(()) -} - -fn cleanup_test_net() { - cleanup_iptables_rules(); - del_netns(TEST_NS_A); - del_netns(TEST_NS_C); - let _ = Command::new("ip") - .args(["link", "del", TEST_BRIDGE]) - .output(); - let _ = Command::new("ip") - .args(["link", "del", TEST_WAN_IF]) - .output(); } fn write_dual_gateway_config(dir: &Path, config: DualGatewayConfig<'_>) -> anyhow::Result { @@ -739,151 +573,6 @@ fn cleanup_miniupnpd_processes() { let _ = Command::new("pkill").args(["-x", "miniupnpd"]).output(); } -fn setup_gateway_iptables_rules( - ext_if: &str, - lan_bridge: &str, - chain_name: &str, - postrouting_chain_name: &str, -) -> anyhow::Result<()> { - cleanup_gateway_iptables_rules(ext_if, lan_bridge, chain_name, postrouting_chain_name); - let iptables = find_iptables_legacy_bin()?; - - run_cmd(&iptables, &["-t", "nat", "-N", chain_name])?; - run_cmd( - &iptables, - &[ - "-t", - "nat", - "-A", - "PREROUTING", - "-i", - ext_if, - "-j", - chain_name, - ], - )?; - run_cmd(&iptables, &["-t", "nat", "-N", postrouting_chain_name])?; - run_cmd( - &iptables, - &[ - "-t", - "nat", - "-A", - "POSTROUTING", - "-o", - ext_if, - "-j", - postrouting_chain_name, - ], - )?; - run_cmd(&iptables, &["-t", "mangle", "-N", chain_name])?; - run_cmd( - &iptables, - &[ - "-t", - "mangle", - "-A", - "PREROUTING", - "-i", - ext_if, - "-j", - chain_name, - ], - )?; - run_cmd(&iptables, &["-N", chain_name])?; - run_cmd( - &iptables, - &[ - "-A", "FORWARD", "-i", ext_if, "!", "-o", ext_if, "-j", chain_name, - ], - )?; - run_cmd( - &iptables, - &[ - "-A", "FORWARD", "-i", lan_bridge, "-o", ext_if, "-j", "ACCEPT", - ], - )?; - Ok(()) -} - -fn cleanup_gateway_iptables_rules( - ext_if: &str, - lan_bridge: &str, - chain_name: &str, - postrouting_chain_name: &str, -) { - let Ok(iptables) = find_iptables_legacy_bin() else { - return; - }; - - let _ = Command::new(&iptables) - .args([ - "-t", - "nat", - "-D", - "PREROUTING", - "-i", - ext_if, - "-j", - chain_name, - ]) - .output(); - let _ = Command::new(&iptables) - .args([ - "-t", - "nat", - "-D", - "POSTROUTING", - "-o", - ext_if, - "-j", - postrouting_chain_name, - ]) - .output(); - let _ = Command::new(&iptables) - .args([ - "-t", - "mangle", - "-D", - "PREROUTING", - "-i", - ext_if, - "-j", - chain_name, - ]) - .output(); - let _ = Command::new(&iptables) - .args([ - "-D", "FORWARD", "-i", ext_if, "!", "-o", ext_if, "-j", chain_name, - ]) - .output(); - let _ = Command::new(&iptables) - .args([ - "-D", "FORWARD", "-i", lan_bridge, "-o", ext_if, "-j", "ACCEPT", - ]) - .output(); - let _ = Command::new(&iptables).args(["-F", chain_name]).output(); - let _ = Command::new(&iptables).args(["-X", chain_name]).output(); - let _ = Command::new(&iptables) - .args(["-t", "mangle", "-F", chain_name]) - .output(); - let _ = Command::new(&iptables) - .args(["-t", "mangle", "-X", chain_name]) - .output(); - let _ = Command::new(&iptables) - .args(["-t", "nat", "-F", chain_name]) - .output(); - let _ = Command::new(&iptables) - .args(["-t", "nat", "-X", chain_name]) - .output(); - let _ = Command::new(&iptables) - .args(["-t", "nat", "-F", postrouting_chain_name]) - .output(); - let _ = Command::new(&iptables) - .args(["-t", "nat", "-X", postrouting_chain_name]) - .output(); -} - fn setup_gateway_iptables_rules_in_netns( netns: &str, ext_if: &str, @@ -1052,174 +741,6 @@ fn cleanup_gateway_iptables_rules_in_netns( ); } -fn setup_iptables_rules() -> anyhow::Result<()> { - cleanup_iptables_rules(); - let iptables = find_iptables_legacy_bin()?; - - run_cmd(&iptables, &["-t", "nat", "-N", "MINIUPNPD"])?; - run_cmd( - &iptables, - &[ - "-t", - "nat", - "-A", - "PREROUTING", - "-i", - TEST_WAN_IF, - "-j", - "MINIUPNPD", - ], - )?; - run_cmd(&iptables, &["-t", "nat", "-N", "MINIUPNPD-POSTROUTING"])?; - run_cmd( - &iptables, - &[ - "-t", - "nat", - "-A", - "POSTROUTING", - "-o", - TEST_WAN_IF, - "-j", - "MINIUPNPD-POSTROUTING", - ], - )?; - - run_cmd(&iptables, &["-t", "mangle", "-N", "MINIUPNPD"])?; - run_cmd( - &iptables, - &[ - "-t", - "mangle", - "-A", - "PREROUTING", - "-i", - TEST_WAN_IF, - "-j", - "MINIUPNPD", - ], - )?; - - run_cmd(&iptables, &["-N", "MINIUPNPD"])?; - run_cmd( - &iptables, - &[ - "-A", - "FORWARD", - "-i", - TEST_WAN_IF, - "!", - "-o", - TEST_WAN_IF, - "-j", - "MINIUPNPD", - ], - )?; - run_cmd( - &iptables, - &[ - "-A", - "FORWARD", - "-i", - TEST_BRIDGE, - "-o", - TEST_WAN_IF, - "-j", - "ACCEPT", - ], - )?; - - Ok(()) -} - -fn cleanup_iptables_rules() { - let Ok(iptables) = find_iptables_legacy_bin() else { - return; - }; - - let _ = Command::new(&iptables) - .args([ - "-t", - "nat", - "-D", - "PREROUTING", - "-i", - TEST_WAN_IF, - "-j", - "MINIUPNPD", - ]) - .output(); - let _ = Command::new(&iptables) - .args([ - "-t", - "nat", - "-D", - "POSTROUTING", - "-o", - TEST_WAN_IF, - "-j", - "MINIUPNPD-POSTROUTING", - ]) - .output(); - let _ = Command::new(&iptables) - .args([ - "-t", - "mangle", - "-D", - "PREROUTING", - "-i", - TEST_WAN_IF, - "-j", - "MINIUPNPD", - ]) - .output(); - let _ = Command::new(&iptables) - .args([ - "-D", - "FORWARD", - "-i", - TEST_WAN_IF, - "!", - "-o", - TEST_WAN_IF, - "-j", - "MINIUPNPD", - ]) - .output(); - let _ = Command::new(&iptables) - .args([ - "-D", - "FORWARD", - "-i", - TEST_BRIDGE, - "-o", - TEST_WAN_IF, - "-j", - "ACCEPT", - ]) - .output(); - let _ = Command::new(&iptables).args(["-F", "MINIUPNPD"]).output(); - let _ = Command::new(&iptables).args(["-X", "MINIUPNPD"]).output(); - let _ = Command::new(&iptables) - .args(["-t", "mangle", "-F", "MINIUPNPD"]) - .output(); - let _ = Command::new(&iptables) - .args(["-t", "mangle", "-X", "MINIUPNPD"]) - .output(); - let _ = Command::new(&iptables) - .args(["-t", "nat", "-F", "MINIUPNPD"]) - .output(); - let _ = Command::new(&iptables) - .args(["-t", "nat", "-X", "MINIUPNPD"]) - .output(); - let _ = Command::new(&iptables) - .args(["-t", "nat", "-F", "MINIUPNPD-POSTROUTING"]) - .output(); - let _ = Command::new(&iptables) - .args(["-t", "nat", "-X", "MINIUPNPD-POSTROUTING"]) - .output(); -} - fn run_cmd>(cmd: S, args: &[&str]) -> anyhow::Result<()> { let cmd = cmd.as_ref(); let output = Command::new(cmd) @@ -1366,36 +887,6 @@ async fn query_udp_mapping( )) } -async fn mapping_exists(local_port: u16) -> bool { - query_udp_mapping(TEST_NS_A, TEST_EXTERNAL_IP, TEST_CLIENT_A_IP, local_port) - .await - .is_ok() -} - -async fn create_test_peer_manager( - inst_name: &str, - netns: Option<&str>, - disable_upnp: bool, - stun_collector: Box, -) -> Arc { - let config = TomlConfigLoader::default(); - config.set_inst_name(inst_name.to_owned()); - config.set_netns(netns.map(ToOwned::to_owned)); - - let global_ctx = Arc::new(GlobalCtx::new(config)); - if disable_upnp { - let mut flags = global_ctx.get_flags(); - flags.disable_upnp = true; - global_ctx.set_flags(flags); - } - global_ctx.replace_stun_info_collector(stun_collector); - - let (packet_tx, _packet_rx) = create_packet_recv_chan(); - let peer_mgr = Arc::new(PeerManager::new(RouteAlgoType::Ospf, global_ctx, packet_tx)); - peer_mgr.run().await.unwrap(); - peer_mgr -} - fn create_test_instance_config( inst_name: &str, netns: Option<&str>, @@ -1411,13 +902,14 @@ fn create_test_instance_config( config } -fn create_test_instance( +fn create_test_instance_with_process_runtime( inst_name: &str, netns: Option<&str>, ipv4: &str, ipv6: &str, - stun_collector: Box, + stun_collector: Box>, configure_flags: impl FnOnce(&mut crate::common::config::Flags), + process_runtime: Arc, ) -> Instance { let config = create_test_instance_config(inst_name, netns, ipv4, ipv6); let mut flags = config.get_flags(); @@ -1425,11 +917,7 @@ fn create_test_instance( configure_flags(&mut flags); config.set_flags(flags); - let instance = Instance::new(config); - instance - .get_global_ctx() - .replace_stun_info_collector(stun_collector); - instance + Instance::new_with_process_runtime_and_stun_provider(config, process_runtime, stun_collector) } async fn wait_for_port_mapping_event( @@ -1457,15 +945,20 @@ where } async fn peer_has_udp_conn_to_remote_addr( - peer_mgr: Arc, + core: Arc, peer_id: u32, expected_remote_addr: SocketAddr, ) -> bool { - let Some(conns) = peer_mgr.get_peer_map().list_peer_conns(peer_id).await else { + let Some(peer) = core + .peer_snapshots() + .await + .into_iter() + .find(|peer| peer.peer_id == peer_id) + else { return false; }; - conns.iter().any(|conn| { + peer.conns.iter().any(|conn| { let Some(tunnel) = conn.tunnel.as_ref() else { return false; }; @@ -1489,263 +982,34 @@ async fn peer_has_udp_conn_to_remote_addr( }) } -#[tokio::test] -#[serial_test::serial(upnp)] -async fn udp_hole_punch_listener_establishes_upnp_mapping() { - let _env = UpnpIntegrationEnv::new().await.unwrap(); - let peer_mgr = create_test_peer_manager( - "upnp-test-listener", - Some(TEST_NS_A), - false, - Box::new(GatewayBackedStunCollector { - netns: TEST_NS_A, - client_ip: TEST_CLIENT_A_IP, - external_ip: TEST_EXTERNAL_IP, - }), - ) - .await; - let mut event_rx = peer_mgr.get_global_ctx().subscribe(); - - let listener = UdpHolePunchListener::new(peer_mgr.clone()).await.unwrap(); - let local_port = listener.get_socket().await.local_addr().unwrap().port(); - - let event = wait_for_port_mapping_event(&mut event_rx).await; - let mapped_addr = query_udp_mapping(TEST_NS_A, TEST_EXTERNAL_IP, TEST_CLIENT_A_IP, local_port) - .await - .unwrap(); - - match event { - GlobalCtxEvent::ListenerPortMappingEstablished { - local_listener, - mapped_listener, - backend, - } => { - let expected_external_ip = TEST_EXTERNAL_IP.to_string(); - assert_eq!(backend, "igd"); - assert_eq!(local_listener.scheme(), "udp"); - assert_eq!(local_listener.port(), Some(local_port)); - assert_eq!(mapped_listener.scheme(), "udp"); - assert_eq!( - mapped_listener.host_str(), - Some(expected_external_ip.as_str()) - ); - assert_eq!(mapped_listener.port(), Some(mapped_addr.port())); +async fn wait_instance_route( + inst: &Instance, + peer_id: u32, + cost: Option, +) -> Result<(), Error> { + let core = inst.get_core_instance(); + let now = std::time::Instant::now(); + while now.elapsed() < Duration::from_secs(5) { + if core + .route_snapshots() + .await + .iter() + .any(|route| route.peer_id == peer_id && cost.is_none_or(|cost| route.cost == cost)) + { + return Ok(()); } - other => panic!("unexpected event: {other:?}"), + tokio::time::sleep(Duration::from_millis(50)).await; } - - assert!(mapping_exists(local_port).await); - - drop(listener); - - wait_for_condition( - || async { !mapping_exists(local_port).await }, - Duration::from_secs(10), - ) - .await; -} - -#[tokio::test] -#[serial_test::serial(upnp)] -async fn udp_hole_punch_listener_skips_upnp_when_disabled() { - let _env = UpnpIntegrationEnv::new().await.unwrap(); - let peer_mgr = create_test_peer_manager( - "upnp-test-disabled", - Some(TEST_NS_A), - true, - Box::new(MockStunInfoCollector { - udp_nat_type: NatType::PortRestricted, - }), - ) - .await; - let mut event_rx = peer_mgr.get_global_ctx().subscribe(); - - let listener = UdpHolePunchListener::new(peer_mgr.clone()).await.unwrap(); - let local_port = listener.get_socket().await.local_addr().unwrap().port(); - - let event = tokio::time::timeout(Duration::from_secs(2), event_rx.recv()).await; - assert!(event.is_err(), "unexpected port mapping event: {event:?}"); - assert!(!mapping_exists(local_port).await); - - drop(listener); -} - -#[tokio::test] -#[serial_test::serial(upnp)] -async fn udp_hole_punch_succeeds_via_upnp_mappings_with_different_external_ports() { - let _env = DualGatewayUpnpIntegrationEnv::new().await.unwrap(); - - let p_a = create_test_peer_manager( - "upnp-test-a", - Some(DUAL_NS_A), - false, - Box::new(GatewayBackedStunCollector { - netns: DUAL_NS_A, - client_ip: DUAL_CLIENT_A_IP, - external_ip: DUAL_EXTERNAL_A_IP, - }), - ) - .await; - let mut event_rx_a = p_a.get_global_ctx().subscribe(); - let p_b = create_test_peer_manager( - "upnp-test-b", - None, - false, - Box::new(MockStunInfoCollector { - udp_nat_type: NatType::Unknown, - }), - ) - .await; - let p_c = create_test_peer_manager( - "upnp-test-c", - Some(DUAL_NS_C), - false, - Box::new(GatewayBackedStunCollector { - netns: DUAL_NS_C, - client_ip: DUAL_CLIENT_C_IP, - external_ip: DUAL_EXTERNAL_C_IP, - }), - ) - .await; - let mut event_rx_c = p_c.get_global_ctx().subscribe(); - - connect_peer_manager(p_a.clone(), p_b.clone()).await; - connect_peer_manager(p_b.clone(), p_c.clone()).await; - timeout_stage( - "wait_route_appear(a,c)", - Duration::from_secs(10), - wait_route_appear(p_a.clone(), p_c.clone()), - ) - .await - .unwrap(); - timeout_stage( - "wait_route_appear_with_cost(a,c,2)", - Duration::from_secs(10), - wait_route_appear_with_cost(p_a.clone(), p_c.my_peer_id(), Some(2)), - ) - .await - .unwrap(); - let mut hole_punching_a = UdpHolePunchConnector::new(p_a.clone()); - let mut hole_punching_c = UdpHolePunchConnector::new(p_c.clone()); - hole_punching_a.run_as_client().await.unwrap(); - hole_punching_c.run_as_server().await.unwrap(); - - timeout_stage( - "udp_hole_punch_run_immediately(a)", - Duration::from_secs(10), - hole_punching_a.run_immediately_for_test(), - ) - .await; - - let event_a = timeout_stage( - "wait_port_mapping_event(a)", - Duration::from_secs(15), - wait_for_port_mapping_event(&mut event_rx_a), - ) - .await; - let event_c = timeout_stage( - "wait_port_mapping_event(c)", - Duration::from_secs(15), - wait_for_port_mapping_event(&mut event_rx_c), - ) - .await; - - let (local_port_a, mapped_port_a) = match event_a { - GlobalCtxEvent::ListenerPortMappingEstablished { - local_listener, - mapped_listener, - backend, - } => { - assert_eq!(backend, "igd"); - ( - local_listener.port().unwrap(), - mapped_listener.port().unwrap(), - ) - } - other => panic!("unexpected event for a: {other:?}"), - }; - let (local_port_c, mapped_port_c) = match event_c { - GlobalCtxEvent::ListenerPortMappingEstablished { - local_listener, - mapped_listener, - backend, - } => { - assert_eq!(backend, "igd"); - ( - local_listener.port().unwrap(), - mapped_listener.port().unwrap(), - ) - } - other => panic!("unexpected event for c: {other:?}"), - }; - - assert_ne!(mapped_port_a, local_port_a); - assert_ne!(mapped_port_c, local_port_c); - - let mapped_addr_a = timeout_stage( - "query_udp_mapping(a)", - Duration::from_secs(10), - query_udp_mapping( - DUAL_NS_A, - DUAL_EXTERNAL_A_IP, - DUAL_CLIENT_A_IP, - local_port_a, - ), - ) - .await - .unwrap(); - let mapped_addr_c = timeout_stage( - "query_udp_mapping(c)", - Duration::from_secs(10), - query_udp_mapping( - DUAL_NS_C, - DUAL_EXTERNAL_C_IP, - DUAL_CLIENT_C_IP, - local_port_c, - ), - ) - .await - .unwrap(); - - assert_eq!(mapped_addr_a.port(), mapped_port_a); - assert_eq!(mapped_addr_c.port(), mapped_port_c); - - timeout_stage( - "wait_route_cost_1_after_udp_hole_punch", - Duration::from_secs(15), - wait_for_condition( - || { - let p_a = p_a.clone(); - let p_c = p_c.clone(); - async move { - let a_ok = p_a - .list_routes() - .await - .iter() - .any(|route| route.peer_id == p_c.my_peer_id() && route.cost == 1); - let c_ok = p_c - .list_routes() - .await - .iter() - .any(|route| route.peer_id == p_a.my_peer_id() && route.cost == 1); - a_ok && c_ok - } - }, - Duration::from_secs(15), - ), - ) - .await; - - assert_ne!(mapped_addr_a.port(), local_port_a); - assert_ne!(mapped_addr_c.port(), local_port_c); + Err(Error::NotFound) } #[tokio::test] #[serial_test::serial(upnp)] async fn instances_build_direct_connection_via_upnp_udp_hole_punch() { let _env = DualGatewayUpnpIntegrationEnv::new().await.unwrap(); + let process_runtime = CoreProcessRuntime::new(); - let mut inst_a = create_test_instance( + let mut inst_a = create_test_instance_with_process_runtime( "upnp-inst-a", Some(DUAL_NS_A), "10.144.200.1/24", @@ -1756,10 +1020,11 @@ async fn instances_build_direct_connection_via_upnp_udp_hole_punch() { external_ip: DUAL_EXTERNAL_A_IP, }), |flags| flags.need_p2p = true, + process_runtime.clone(), ); let mut event_rx_a = inst_a.get_global_ctx().subscribe(); - let mut inst_b = create_test_instance( + let mut inst_b = create_test_instance_with_process_runtime( "upnp-inst-b", None, "10.144.200.2/24", @@ -1768,9 +1033,10 @@ async fn instances_build_direct_connection_via_upnp_udp_hole_punch() { udp_nat_type: NatType::Unknown, }), |_| {}, + process_runtime.clone(), ); - let mut inst_c = create_test_instance( + let mut inst_c = create_test_instance_with_process_runtime( "upnp-inst-c", Some(DUAL_NS_C), "10.144.200.3/24", @@ -1781,6 +1047,7 @@ async fn instances_build_direct_connection_via_upnp_udp_hole_punch() { external_ip: DUAL_EXTERNAL_C_IP, }), |flags| flags.need_p2p = true, + process_runtime, ); let mut event_rx_c = inst_c.get_global_ctx().subscribe(); @@ -1788,28 +1055,23 @@ async fn instances_build_direct_connection_via_upnp_udp_hole_punch() { inst_b.run().await.unwrap(); inst_c.run().await.unwrap(); - inst_a - .get_conn_manager() - .add_connector(RingTunnelConnector::new( - format!("ring://{}", inst_b.id()).parse().unwrap(), - )); - inst_c - .get_conn_manager() - .add_connector(RingTunnelConnector::new( - format!("ring://{}", inst_b.id()).parse().unwrap(), - )); + inst_a.add_connector_url(inst_b.ring_listener_url()); + inst_c.add_connector_url(inst_b.ring_listener_url()); timeout_stage( "wait_route_appear(inst_a, inst_c)", Duration::from_secs(10), - wait_route_appear(inst_a.get_peer_manager(), inst_c.get_peer_manager()), + async { + wait_instance_route(&inst_a, inst_c.peer_id(), None).await?; + wait_instance_route(&inst_c, inst_a.peer_id(), None).await + }, ) .await .unwrap(); timeout_stage( "wait_route_cost_2(inst_a -> inst_c)", Duration::from_secs(10), - wait_route_appear_with_cost(inst_a.get_peer_manager(), inst_c.peer_id(), Some(2)), + wait_instance_route(&inst_a, inst_c.peer_id(), Some(2)), ) .await .unwrap(); @@ -1902,31 +1164,31 @@ async fn instances_build_direct_connection_via_upnp_udp_hole_punch() { Duration::from_secs(20), wait_for_condition( || { - let peer_mgr_a = inst_a.get_peer_manager(); - let peer_mgr_c = inst_c.get_peer_manager(); + let core_a = inst_a.get_core_instance(); + let core_c = inst_c.get_core_instance(); let peer_id_a = inst_a.peer_id(); let peer_id_c = inst_c.peer_id(); async move { - peer_mgr_a.get_peer_map().has_peer(peer_id_c) - && peer_mgr_c.get_peer_map().has_peer(peer_id_a) - && peer_mgr_a - .list_routes() + core_a.connected_peers().await.contains(&peer_id_c) + && core_c.connected_peers().await.contains(&peer_id_a) + && core_a + .route_snapshots() .await .iter() .any(|route| route.peer_id == peer_id_c && route.cost == 1) - && peer_mgr_c - .list_routes() + && core_c + .route_snapshots() .await .iter() .any(|route| route.peer_id == peer_id_a && route.cost == 1) && peer_has_udp_conn_to_remote_addr( - peer_mgr_a.clone(), + core_a.clone(), peer_id_c, mapped_addr_c, ) .await && peer_has_udp_conn_to_remote_addr( - peer_mgr_c.clone(), + core_c.clone(), peer_id_a, mapped_addr_a, ) diff --git a/easytier/src/tunnel/buf.rs b/easytier/src/tunnel/buf.rs deleted file mode 100644 index 01363096..00000000 --- a/easytier/src/tunnel/buf.rs +++ /dev/null @@ -1,92 +0,0 @@ -use std::collections::VecDeque; -use std::io::IoSlice; - -use bytes::{Buf, BufMut, Bytes, BytesMut}; - -pub(crate) struct BufList { - bufs: VecDeque, -} - -impl BufList { - pub(crate) fn new() -> BufList { - BufList { - bufs: VecDeque::new(), - } - } - - #[inline] - pub(crate) fn push(&mut self, buf: T) { - debug_assert!(buf.has_remaining()); - self.bufs.push_back(buf); - } - - #[inline] - pub(crate) fn bufs_cnt(&self) -> usize { - self.bufs.len() - } -} - -impl Buf for BufList { - #[inline] - fn remaining(&self) -> usize { - self.bufs.iter().map(|buf| buf.remaining()).sum() - } - - #[inline] - fn chunk(&self) -> &[u8] { - self.bufs.front().map(Buf::chunk).unwrap_or_default() - } - - #[inline] - fn advance(&mut self, mut cnt: usize) { - while cnt > 0 { - { - let front = &mut self.bufs[0]; - let rem = front.remaining(); - if rem > cnt { - front.advance(cnt); - return; - } else { - front.advance(rem); - cnt -= rem; - } - } - self.bufs.pop_front(); - } - } - - #[inline] - fn chunks_vectored<'t>(&'t self, dst: &mut [IoSlice<'t>]) -> usize { - if dst.is_empty() { - return 0; - } - let mut vecs = 0; - for buf in &self.bufs { - vecs += buf.chunks_vectored(&mut dst[vecs..]); - if vecs == dst.len() { - break; - } - } - vecs - } - - #[inline] - fn copy_to_bytes(&mut self, len: usize) -> Bytes { - // Our inner buffer may have an optimized version of copy_to_bytes, and if the whole - // request can be fulfilled by the front buffer, we can take advantage. - match self.bufs.front_mut() { - Some(front) if front.remaining() == len => { - let b = front.copy_to_bytes(len); - self.bufs.pop_front(); - b - } - Some(front) if front.remaining() > len => front.copy_to_bytes(len), - _ => { - assert!(len <= self.remaining(), "`len` greater than remaining"); - let mut bm = BytesMut::with_capacity(len); - bm.put(self.take(len)); - bm.freeze() - } - } - } -} diff --git a/easytier/src/tunnel/common.rs b/easytier/src/tunnel/common.rs index 03a75acf..bffebfda 100644 --- a/easytier/src/tunnel/common.rs +++ b/easytier/src/tunnel/common.rs @@ -1,336 +1,10 @@ use bon::builder; -use futures::{Future, Sink, Stream, stream::FuturesUnordered}; use network_interface::NetworkInterfaceConfig as _; -use pin_project_lite::pin_project; -use std::{ - any::Any, - net::{IpAddr, SocketAddr}, - pin::Pin, - sync::{Arc, Mutex}, - task::{Poll, ready}, -}; -use tokio::io::{AsyncRead, AsyncWrite, ReadBuf}; +use std::net::{IpAddr, SocketAddr}; -use super::TunnelInfo; -use super::{ - SinkItem, StreamItem, Tunnel, TunnelError, ZCPacketSink, ZCPacketStream, - buf::BufList, - packet_def::{TCP_TUNNEL_HEADER_SIZE, TCPTunnelHeader, ZCPacketType}, -}; use crate::common::netns::NetNS; -use crate::tunnel::packet_def::{PEER_MANAGER_HEADER_SIZE, ZCPacket}; -use bytes::{Buf, BufMut, Bytes, BytesMut}; +use easytier_core::tunnel::TunnelError; use tokio::net::{TcpListener, TcpSocket, UdpSocket}; -use tokio_stream::StreamExt; -use tokio_util::io::poll_write_buf; -use zerocopy::FromBytes as _; - -pub struct TunnelWrapper { - reader: Arc>>, - writer: Arc>>, - info: Option, - associate_data: Option>, -} - -impl TunnelWrapper { - pub fn new(reader: R, writer: W, info: Option) -> Self { - Self::new_with_associate_data(reader, writer, info, None) - } - - pub fn new_with_associate_data( - reader: R, - writer: W, - info: Option, - associate_data: Option>, - ) -> Self { - TunnelWrapper { - reader: Arc::new(Mutex::new(Some(reader))), - writer: Arc::new(Mutex::new(Some(writer))), - info, - associate_data, - } - } -} - -impl Tunnel for TunnelWrapper -where - R: ZCPacketStream + Send + 'static, - W: ZCPacketSink + Send + 'static, -{ - fn split(&self) -> (Pin>, Pin>) { - let reader = self.reader.lock().unwrap().take().unwrap(); - let writer = self.writer.lock().unwrap().take().unwrap(); - (Box::pin(reader), Box::pin(writer)) - } - - fn info(&self) -> Option { - self.info.clone() - } -} - -// a length delimited codec for async reader -pin_project! { - pub struct FramedReader { - #[pin] - reader: R, - buf: BytesMut, - max_packet_size: usize, - associate_data: Option>, - error: Option, - } -} - -impl FramedReader { - pub fn new(reader: R, max_packet_size: usize) -> Self { - Self::new_with_associate_data(reader, max_packet_size, None) - } - - pub fn new_with_associate_data( - reader: R, - max_packet_size: usize, - associate_data: Option>, - ) -> Self { - FramedReader { - reader, - buf: BytesMut::with_capacity(max_packet_size), - max_packet_size, - associate_data, - error: None, - } - } - - fn extract_one_packet( - buf: &mut BytesMut, - max_packet_size: usize, - ) -> Option> { - if buf.len() < TCP_TUNNEL_HEADER_SIZE { - // header is not complete - return None; - } - - let header = TCPTunnelHeader::ref_from_prefix(&buf[..]).unwrap(); - let body_len = header.len.get() as usize; - if body_len > max_packet_size { - // body is too long - return Some(Err(TunnelError::InvalidPacket("body too long".to_string()))); - } - - if body_len < PEER_MANAGER_HEADER_SIZE { - return Some(Err(TunnelError::InvalidPacket( - "body too short".to_string(), - ))); - } - - if buf.len() < TCP_TUNNEL_HEADER_SIZE + body_len { - // body is not complete - return None; - } - - // extract one packet - let packet_buf = buf.split_to(TCP_TUNNEL_HEADER_SIZE + body_len); - Some(Ok(ZCPacket::new_from_buf(packet_buf, ZCPacketType::TCP))) - } -} - -impl Stream for FramedReader -where - R: AsyncRead + Send + 'static + Unpin, -{ - type Item = StreamItem; - - fn poll_next( - self: Pin<&mut Self>, - cx: &mut std::task::Context<'_>, - ) -> std::task::Poll> { - let mut self_mut = self.project(); - - loop { - if let Some(e) = self_mut.error.as_ref() { - tracing::warn!("poll_next on a failed FramedReader, {:?}", e); - return Poll::Ready(None); - } - - if let Some(packet) = Self::extract_one_packet(self_mut.buf, *self_mut.max_packet_size) - { - if let Err(TunnelError::InvalidPacket(msg)) = packet.as_ref() { - self_mut - .error - .replace(TunnelError::InvalidPacket(msg.clone())); - } - return Poll::Ready(Some(packet)); - } - - reserve_buf( - self_mut.buf, - *self_mut.max_packet_size, - *self_mut.max_packet_size * 2, - ); - - let cap = self_mut.buf.capacity() - self_mut.buf.len(); - let buf = self_mut.buf.chunk_mut().as_mut_ptr(); - let buf = unsafe { std::slice::from_raw_parts_mut(buf, cap) }; - let mut buf = ReadBuf::new(buf); - - let ret = ready!(self_mut.reader.as_mut().poll_read(cx, &mut buf)); - let len = buf.filled().len(); - unsafe { self_mut.buf.advance_mut(len) }; - - match ret { - Ok(_) => { - if len == 0 { - return Poll::Ready(None); - } - } - Err(e) => { - return Poll::Ready(Some(Err(TunnelError::IOError(e)))); - } - } - } - } -} - -pub trait ZCPacketToBytes { - fn zcpacket_into_bytes(&self, zc_packet: ZCPacket) -> Result; -} - -pub struct TcpZCPacketToBytes; -impl ZCPacketToBytes for TcpZCPacketToBytes { - fn zcpacket_into_bytes(&self, item: ZCPacket) -> Result { - let mut item = item.convert_type(ZCPacketType::TCP); - - let tcp_len = PEER_MANAGER_HEADER_SIZE + item.payload_len(); - let Some(header) = item.mut_tcp_tunnel_header() else { - return Err(TunnelError::InvalidPacket("packet too short".to_string())); - }; - header.len.set(tcp_len.try_into().unwrap()); - - Ok(item.into_bytes()) - } -} - -pin_project! { - pub struct FramedWriter { - #[pin] - writer: W, - sending_bufs: BufList, - associate_data: Option>, - - converter: C, - } -} - -impl FramedWriter { - fn max_buffer_count(&self) -> usize { - 64 - } -} - -impl FramedWriter { - pub fn new(writer: W) -> Self { - Self::new_with_associate_data(writer, None) - } - - pub fn new_with_associate_data( - writer: W, - associate_data: Option>, - ) -> Self { - FramedWriter { - writer, - sending_bufs: BufList::new(), - associate_data, - converter: TcpZCPacketToBytes {}, - } - } -} - -impl FramedWriter { - pub fn new_with_converter(writer: W, converter: C) -> Self { - Self::new_with_converter_and_associate_data(writer, converter, None) - } - - pub fn new_with_converter_and_associate_data( - writer: W, - converter: C, - associate_data: Option>, - ) -> Self { - FramedWriter { - writer, - sending_bufs: BufList::new(), - associate_data, - converter, - } - } -} - -impl Sink for FramedWriter -where - W: AsyncWrite + Send + 'static, - C: ZCPacketToBytes + Send + 'static, -{ - type Error = TunnelError; - - fn poll_ready( - mut self: Pin<&mut Self>, - cx: &mut std::task::Context<'_>, - ) -> std::task::Poll> { - let max_buffer_count = self.max_buffer_count(); - if self.sending_bufs.bufs_cnt() >= max_buffer_count { - self.as_mut().poll_flush(cx) - } else { - tracing::trace!(bufs_cnt = self.sending_bufs.bufs_cnt(), "ready to send"); - Poll::Ready(Ok(())) - } - } - - fn start_send(self: Pin<&mut Self>, item: ZCPacket) -> Result<(), Self::Error> { - let pinned = self.project(); - pinned - .sending_bufs - .push(pinned.converter.zcpacket_into_bytes(item)?); - - Ok(()) - } - - fn poll_flush( - self: Pin<&mut Self>, - cx: &mut std::task::Context<'_>, - ) -> Poll> { - let mut pinned = self.project(); - let mut remaining = pinned.sending_bufs.remaining(); - while remaining != 0 { - let n = ready!(poll_write_buf( - pinned.writer.as_mut(), - cx, - pinned.sending_bufs - ))?; - if n == 0 { - return Poll::Ready(Err(TunnelError::IOError(std::io::Error::new( - std::io::ErrorKind::WriteZero, - "failed to \ - write frame to transport", - )))); - } - remaining -= n; - } - - tracing::trace!(?remaining, "flushed"); - - // Try flushing the underlying IO - ready!(pinned.writer.poll_flush(cx))?; - - Poll::Ready(Ok(())) - } - - fn poll_close( - mut self: Pin<&mut Self>, - cx: &mut std::task::Context<'_>, - ) -> Poll> { - ready!(self.as_mut().poll_flush(cx))?; - ready!(self.project().writer.poll_shutdown(cx))?; - - Poll::Ready(Ok(())) - } -} pub(crate) fn get_interface_name_by_ip(local_ip: &IpAddr) -> Option { if local_ip.is_unspecified() || local_ip.is_multicast() { @@ -349,27 +23,6 @@ pub(crate) fn get_interface_name_by_ip(local_ip: &IpAddr) -> Option { None } -pub(crate) async fn wait_for_connect_futures( - mut futures: FuturesUnordered, -) -> Result -where - Fut: Future> + Send, - E: std::error::Error + Into + Send + 'static, -{ - // return last error - let mut last_err = None; - - while let Some(ret) = futures.next().await { - if let Err(e) = ret { - last_err = Some(e.into()); - } else { - return ret.map_err(|e| e.into()); - } - } - - Err(last_err.unwrap_or(TunnelError::Shutdown)) -} - // region bind pub trait Bindable: Sized { @@ -417,6 +70,8 @@ fn setup_socket2_ext( bind_addr: &SocketAddr, #[allow(unused_variables)] bind_dev: Option, only_v6: bool, + reuse_addr: bool, + reuse_port: bool, socket_mark: Option, ) -> Result<(), TunnelError> { #[cfg(target_os = "windows")] @@ -430,7 +85,15 @@ fn setup_socket2_ext( } socket2_socket.set_nonblocking(true)?; - socket2_socket.set_reuse_address(!cfg!(target_os = "windows"))?; + socket2_socket.set_reuse_address(reuse_addr)?; + #[cfg(all(unix, not(target_os = "solaris"), not(target_os = "illumos")))] + if reuse_port { + socket2_socket.set_reuse_port(true)?; + } + #[cfg(not(all(unix, not(target_os = "solaris"), not(target_os = "illumos"))))] + { + let _ = reuse_port; + } // SO_MARK must be set before bind() so the kernel applies the mark to // any source-address selection bind() triggers on unspecified binds. @@ -445,28 +108,27 @@ fn setup_socket2_ext( } } - // #[cfg(all(unix, not(target_os = "solaris"), not(target_os = "illumos")))] - // socket2_socket.set_reuse_port(true)?; - - if bind_addr.ip().is_unspecified() { - return Ok(()); - } - // linux/mac does not use interface of bind_addr to send packet, so we need to bind device // win can handle this with bind correctly #[cfg(any(target_os = "ios", target_os = "macos"))] if let Some(dev_name) = bind_dev { // use IP_BOUND_IF to bind device - unsafe { - let dev_idx = nix::libc::if_nametoindex(dev_name.as_str().as_ptr() as *const i8); - tracing::warn!(?dev_idx, ?dev_name, "bind device"); - if bind_addr.is_ipv4() { - socket2_socket.bind_device_by_index_v4(std::num::NonZeroU32::new(dev_idx))?; - } else { - socket2_socket.bind_device_by_index_v6(std::num::NonZeroU32::new(dev_idx))?; - } - tracing::warn!(?dev_idx, ?dev_name, "bind device doen"); + let c_dev_name = std::ffi::CString::new(dev_name.clone()).map_err(|err| { + TunnelError::InvalidAddr(format!("invalid interface name {dev_name}: {err}")) + })?; + let dev_idx = unsafe { nix::libc::if_nametoindex(c_dev_name.as_ptr()) }; + let Some(dev_idx) = std::num::NonZeroU32::new(dev_idx) else { + return Err(TunnelError::InvalidAddr(format!( + "network interface not found: {dev_name}" + ))); + }; + tracing::warn!(?dev_idx, ?dev_name, "bind device"); + if bind_addr.is_ipv4() { + socket2_socket.bind_device_by_index_v4(Some(dev_idx))?; + } else { + socket2_socket.bind_device_by_index_v6(Some(dev_idx))?; } + tracing::warn!(?dev_idx, ?dev_name, "bind device done"); } #[cfg(any( @@ -559,6 +221,8 @@ pub fn bind( #[builder(default, into)] dev: BindDev, net_ns: Option, #[builder(default)] only_v6: bool, + #[builder(default = !cfg!(target_os = "windows"))] reuse_addr: bool, + #[builder(default)] reuse_port: bool, /// Linux SO_MARK (fwmark) to apply to the socket. `None` leaves SO_MARK /// untouched; `Some(mark)` applies that exact value, including `Some(0)`. socket_mark: Option, @@ -570,35 +234,34 @@ pub fn bind( BindDev::Custom(s) => Some(s), }; let socket = socket2::Socket::new(socket2::Domain::for_address(addr), B::TYPE, B::PROTOCOL)?; - setup_socket2_ext(&socket, &addr, dev, only_v6, socket_mark)?; + setup_socket2_ext( + &socket, + &addr, + dev, + only_v6, + reuse_addr, + reuse_port, + socket_mark, + )?; B::finalize(socket) } -pub fn reserve_buf(buf: &mut BytesMut, min_size: usize, max_size: usize) { - if buf.capacity() < min_size { - buf.reserve(max_size); - } -} - // endregion -pub mod tests { +#[cfg(test)] +pub(crate) mod tests { use atomic_shim::AtomicU64; use std::{sync::Arc, time::Instant}; use futures::{Future, SinkExt, StreamExt}; use tokio_util::bytes::{BufMut, Bytes, BytesMut}; - use crate::{ - common::netns::NetNS, - tunnel::{TunnelConnector, TunnelListener, packet_def::ZCPacket}, + use easytier_core::{ + connectivity::protocol::raw::TunnelDialer, packet::ZCPacket, socket::SocketListener, + tunnel::Tunnel, }; - #[cfg(test)] - use crate::tunnel::{ - TunnelError, - packet_def::{PEER_MANAGER_HEADER_SIZE, TCP_TUNNEL_HEADER_SIZE}, - }; + use crate::common::netns::NetNS; #[cfg(any(target_os = "android", target_os = "fuchsia", target_os = "linux"))] #[test] @@ -653,21 +316,26 @@ pub mod tests { assert_eq!(read_so_mark(&socket), 0); } + #[cfg(any( + target_os = "android", + target_os = "fuchsia", + target_os = "linux", + target_env = "ohos" + ))] #[test] - fn framed_reader_rejects_short_peer_manager_body() { - let mut buf = BytesMut::new(); - buf.put_u32_le((PEER_MANAGER_HEADER_SIZE - 1) as u32); - buf.resize(TCP_TUNNEL_HEADER_SIZE + PEER_MANAGER_HEADER_SIZE - 1, 0); + fn bind_custom_device_is_applied_for_unspecified_addr() { + use std::net::SocketAddr; + use tokio::net::UdpSocket; - let ret = super::FramedReader::::extract_one_packet(&mut buf, 2000); - - assert!(matches!( - ret, - Some(Err(TunnelError::InvalidPacket(msg))) if msg == "body too short" - )); + let addr: SocketAddr = "0.0.0.0:0".parse().unwrap(); + let _err = super::bind::() + .addr(addr) + .dev("et/invalid-device-name") + .call() + .expect_err("custom device must not be skipped for unspecified bind addr"); } - pub async fn _tunnel_echo_server(tunnel: Box, once: bool) { + pub async fn _tunnel_echo_server(tunnel: Box, once: bool) { let (mut recv, mut send) = tunnel.split(); if !once { @@ -700,33 +368,15 @@ pub mod tests { tracing::warn!("echo server exit..."); } - pub(crate) async fn _tunnel_pingpong(listener: L, connector: C) - where - L: TunnelListener + Send + Sync + 'static, - C: TunnelConnector + Send + Sync + 'static, - { - _tunnel_pingpong_netns_with_timeout( - listener, - connector, - NetNS::new(None), - NetNS::new(None), - "12345678abcdefg".as_bytes().to_vec(), - // only used by tunnel test, so set a long timeout - tokio::time::Duration::from_secs(5), - ) - .await - .unwrap(); - } - async fn _tunnel_pingpong_netns( mut listener: L, - mut connector: C, + connector: C, l_netns: NetNS, c_netns: NetNS, buf: Vec, ) where - L: TunnelListener + Send + Sync + 'static, - C: TunnelConnector + Send + Sync + 'static, + L: SocketListener> + Sync + 'static, + C: TunnelDialer, { l_netns .run_async(|| async { @@ -791,8 +441,8 @@ pub mod tests { timeout: std::time::Duration, ) -> Result<(), anyhow::Error> where - L: TunnelListener + Send + Sync + 'static, - C: TunnelConnector + Send + Sync + 'static, + L: SocketListener> + Sync + 'static, + C: TunnelDialer, { let handle = tokio::spawn(async move { _tunnel_pingpong_netns(listener, connector, l_netns, c_netns, buf).await; @@ -821,23 +471,15 @@ pub mod tests { } } - pub(crate) async fn _tunnel_bench(listener: L, connector: C) - where - L: TunnelListener + Send + Sync + 'static, - C: TunnelConnector + Send + Sync + 'static, - { - _tunnel_bench_netns(listener, connector, NetNS::new(None), NetNS::new(None)).await; - } - pub(crate) async fn _tunnel_bench_netns( mut listener: L, - mut connector: C, + connector: C, netns_l: NetNS, netns_c: NetNS, ) -> usize where - L: TunnelListener + Send + Sync + 'static, - C: TunnelConnector + Send + Sync + 'static, + L: SocketListener> + Sync + 'static, + C: TunnelDialer, { { let _g = netns_l.guard(); diff --git a/easytier/src/tunnel/fake_tcp/mod.rs b/easytier/src/tunnel/fake_tcp/mod.rs deleted file mode 100644 index 9cd052f8..00000000 --- a/easytier/src/tunnel/fake_tcp/mod.rs +++ /dev/null @@ -1,594 +0,0 @@ -mod netfilter; -mod packet; -mod stack; - -use bytes::BytesMut; -use futures::{Sink, Stream}; -use network_interface::NetworkInterfaceConfig; -use pnet::util::MacAddr; -use std::{ - net::{IpAddr, Ipv4Addr, SocketAddr, UdpSocket}, - pin::Pin, - sync::Arc, - task::{Context as TaskContext, Poll}, -}; -use tokio::{io::AsyncReadExt, net::TcpStream}; - -use crate::tunnel::{ - FromUrl, IpVersion, SinkError, SinkItem, StreamItem, Tunnel, TunnelConnector, TunnelError, - TunnelInfo, TunnelListener, - common::TunnelWrapper, - fake_tcp::netfilter::create_tun, - packet_def::{PEER_MANAGER_HEADER_SIZE, TCP_TUNNEL_HEADER_SIZE, ZCPacket, ZCPacketType}, -}; - -use futures::Future; -use tokio_util::task::AbortOnDropHandle; - -use dashmap::DashMap; - -struct IpToIfNameCache { - ip_to_ifname: DashMap)>, -} - -impl IpToIfNameCache { - fn new() -> Self { - Self { - ip_to_ifname: DashMap::new(), - } - } - - fn reload_ip_to_ifname(&self) { - self.ip_to_ifname.clear(); - let Ok(interfaces) = network_interface::NetworkInterface::show() else { - tracing::warn!("failed to enumerate interfaces when reloading faketcp ip cache"); - return; - }; - for iface in interfaces { - let mac = iface.mac_addr.as_deref().and_then(|mac| { - mac.parse::().map_err(|e| { - tracing::debug!(iface = %iface.name, mac, ?e, "failed to parse interface mac") - }).ok() - }); - for ip in iface.addr.iter() { - self.ip_to_ifname.insert(ip.ip(), (iface.name.clone(), mac)); - } - } - } - - fn get_ifname(&self, ip: &IpAddr) -> Option<(String, Option)> { - if let Some(ifname) = self.ip_to_ifname.get(ip) { - Some(ifname.clone()) - } else { - self.reload_ip_to_ifname(); - self.ip_to_ifname.get(ip).map(|s| s.clone()) - } - } -} - -fn get_faketcp_tunnel_type_str(driver_type: &str) -> String { - format!("faketcp_{}", driver_type) -} - -async fn create_tun_off_runtime( - interface_name: String, - src_addr: Option, - dst_addr: SocketAddr, -) -> Result, TunnelError> { - tokio::task::spawn_blocking(move || create_tun(&interface_name, src_addr, dst_addr)) - .await - .map_err(|e| TunnelError::InternalError(format!("faketcp create_tun task failed: {e}")))? - .map_err(Into::into) -} - -pub struct FakeTcpTunnelListener { - addr: url::Url, - os_listener: Option, - // interface_name -> fake tcp stack - stack_map: DashMap>, - // a cache from ip addr to interface name - ip_to_ifname: IpToIfNameCache, -} - -impl FakeTcpTunnelListener { - pub fn new(addr: url::Url) -> Self { - // Define filter: Capture all packets (or refine this if needed) - // For FakeTCP, we probably want to capture packets destined to us? - // But `stack::Stack` handles IP/TCP logic. - // Maybe we just capture everything for now as a raw tunnel? - // Or better, filter based on some criteria? - // The user said "satisfy filter function". - // Let's create a filter that accepts everything for now, or maybe only IP packets? - FakeTcpTunnelListener { - addr, - os_listener: None, - stack_map: DashMap::new(), - ip_to_ifname: IpToIfNameCache::new(), - } - } - - async fn do_accept(&mut self) -> Result { - loop { - match self.os_listener.as_mut().unwrap().accept().await { - Ok((s, remote_addr)) => { - let Ok(local_addr) = s.local_addr() else { - tracing::warn!("accept fail with local_addr error"); - continue; - }; - let Some((interface_name, mac)) = - self.ip_to_ifname.get_ifname(&local_addr.ip()) - else { - tracing::warn!("accept fail with interface_name error"); - continue; - }; - return Ok(AcceptResult { - socket: s, - local_addr, - remote_addr, - interface_name, - mac, - }); - } - Err(e) => { - use std::io::ErrorKind::*; - if matches!( - e.kind(), - NotConnected | ConnectionAborted | ConnectionRefused | ConnectionReset - ) { - tracing::warn!(?e, "accept fail with retryable error: {:?}", e); - continue; - } - tracing::warn!(?e, "accept fail"); - return Err(e.into()); - } - } - } - } - - async fn get_stack( - &self, - accept_result: &AcceptResult, - ) -> Result, TunnelError> { - let local_socket_addr = accept_result.local_addr; - - let interface_name = &accept_result.interface_name; - - let (local_ip, local_ip6) = match local_socket_addr.ip() { - IpAddr::V4(ip) => (Some(ip), None), - IpAddr::V6(ip) => (None, Some(ip)), - }; - - if let Some(entry) = self.stack_map.get(interface_name) { - let stack = entry.clone(); - drop(entry); - - if !stack.is_closed() { - return Ok(stack); - } - - tracing::warn!( - interface_name, - "fake_tcp stack reader_task finished, recreating stack" - ); - self.stack_map.remove(interface_name); - } - - let tun = - create_tun_off_runtime(interface_name.to_string(), None, local_socket_addr).await?; - tracing::info!( - ?local_socket_addr, - "create new stack with interface_name: {:?}", - interface_name - ); - let stack = Arc::new(stack::Stack::new( - tun, - local_ip.unwrap_or(Ipv4Addr::UNSPECIFIED), - local_ip6, - accept_result.mac, - )); - self.stack_map - .insert(interface_name.to_string(), stack.clone()); - - Ok(stack) - } -} - -fn build_os_socket_reader_task(mut socket: TcpStream) -> AbortOnDropHandle<()> { - AbortOnDropHandle::new(tokio::spawn(async move { - // read the os socket until it's closed - let mut buf = [0u8; 1024]; - while let Ok(size) = socket.read(&mut buf).await { - tracing::trace!("read {} bytes from os socket", size); - if size == 0 { - break; - } - } - tracing::info!("FakeTcpTunnelListener os socket closed"); - })) -} - -#[derive(Debug)] -struct AcceptResult { - socket: TcpStream, - local_addr: SocketAddr, - remote_addr: SocketAddr, - interface_name: String, - mac: Option, -} - -#[async_trait::async_trait] -impl TunnelListener for FakeTcpTunnelListener { - async fn listen(&mut self) -> Result<(), TunnelError> { - let port = self.addr.port().unwrap_or(0); - let bind_addr = SocketAddr::from_url(self.addr.clone(), IpVersion::Both).await?; - let os_listener = tokio::net::TcpListener::bind(bind_addr).await?; - tracing::info!(port, "FakeTcpTunnelListener listening"); - self.os_listener = Some(os_listener); - Ok(()) - } - - async fn accept(&mut self) -> Result, TunnelError> { - tracing::debug!("FakeTcpTunnelListener waiting for accept"); - let (res, stack, socket) = loop { - let res = self.do_accept().await?; - let stack = self.get_stack(&res).await?; - let socket = stack.try_alloc_established_socket( - res.local_addr, - res.remote_addr, - stack::State::Established, - ); - let Some(socket) = socket else { - tracing::warn!( - interface_name = res.interface_name, - "fake_tcp stack closed while accepting connection, dropping accepted socket" - ); - self.stack_map.remove(&res.interface_name); - continue; - }; - break (res, stack, socket); - }; - - tracing::info!( - ?res, - remote = socket.remote_addr().to_string(), - "FakeTcpTunnelListener accepted connection" - ); - - let info = TunnelInfo { - tunnel_type: get_faketcp_tunnel_type_str(stack.driver_type()), - local_addr: Some(self.local_url().into()), - remote_addr: Some( - crate::tunnel::build_url_from_socket_addr( - &socket.remote_addr().to_string(), - "faketcp", - ) - .into(), - ), - resolved_remote_addr: Some( - crate::tunnel::build_url_from_socket_addr( - &socket.remote_addr().to_string(), - "faketcp", - ) - .into(), - ), - }; - - // We treat the fake tcp socket as a datagram tunnel directly - // The reader/writer will interface with the socket using recv_bytes/send - // We need to adapt the socket to ZCPacketStream and ZCPacketSink - - // Since FakeTcpTunnel is a datagram tunnel, we don't need FramedReader/Writer (which are for stream based tunnels like TCP) - // We should wrap the socket into something that produces/consumes ZCPacket directly. - - let socket = Arc::new(socket); - let reader = FakeTcpStream::new(socket.clone()); - let writer = FakeTcpSink::new(socket); - - Ok(Box::new(TunnelWrapper::new_with_associate_data( - reader, - writer, - Some(info), - Some(Box::new(build_os_socket_reader_task(res.socket))), - ))) - } - - fn local_url(&self) -> url::Url { - self.addr.clone() - } -} - -pub struct FakeTcpTunnelConnector { - addr: url::Url, - ip_to_if_name: IpToIfNameCache, - resolved_addr: Option, - socket_mark: Option, -} - -impl FakeTcpTunnelConnector { - pub fn new(addr: url::Url) -> Self { - FakeTcpTunnelConnector { - addr, - ip_to_if_name: IpToIfNameCache::new(), - resolved_addr: None, - socket_mark: None, - } - } -} - -fn get_local_ip_for_destination(destination: IpAddr) -> Option { - // 使用一个不可路由的、私有的、或回环地址创建一个临时的 socket,让内核自动选择源接口。 - // 对于 IPv4,使用 0.0.0.0; 对于 IPv6,使用 :: - let bind_addr = if destination.is_ipv4() { - IpAddr::V4(Ipv4Addr::new(0, 0, 0, 0)) - } else { - IpAddr::V6(std::net::Ipv6Addr::new(0, 0, 0, 0, 0, 0, 0, 0)) - }; - - // 绑定到一个临时端口 (0) - let socket = UdpSocket::bind((bind_addr, 0)).ok()?; - - // 尝试连接到目标地址。这不会真正发送数据包,只是让内核确定路由。 - socket.connect((destination, 80)).ok()?; // 使用一个常见的端口,例如 80 - - // 获取 socket 的本地地址信息 - socket.local_addr().map(|addr| addr.ip()).ok() -} - -#[async_trait::async_trait] -impl TunnelConnector for FakeTcpTunnelConnector { - async fn connect(&mut self) -> Result, TunnelError> { - let remote_addr = match self.resolved_addr { - Some(addr) => addr, - None => SocketAddr::from_url(self.addr.clone(), IpVersion::Both).await?, - }; - let local_ip = get_local_ip_for_destination(remote_addr.ip()) - .ok_or(TunnelError::InternalError("Failed to get local ip".into()))?; - - let os_socket = tokio::net::TcpSocket::new_v4()?; - // SO_MARK applies only to the kernel-visible "decoy" socket below. - // The actual FakeTCP payload travels via crafted segments written - // straight to the TUN device, which the kernel doesn't tag with - // SO_MARK. Operators relying on fwmark for FakeTCP must mark the - // TUN device's traffic with a separate nftables/iptables rule. - crate::tunnel::common::apply_socket_mark( - &socket2::SockRef::from(&os_socket), - self.socket_mark, - )?; - os_socket.bind("0.0.0.0:0".parse().unwrap())?; - let local_port = os_socket.local_addr()?.port(); - let local_addr = SocketAddr::new(local_ip, local_port); - - let (interface_name, mac) = - self.ip_to_if_name - .get_ifname(&local_ip) - .ok_or(TunnelError::InternalError( - "Failed to get interface name".into(), - ))?; - - let (local_ip, local_ip6) = match local_ip { - IpAddr::V4(ip) => (Some(ip), None), - IpAddr::V6(ip) => (None, Some(ip)), - }; - - let tun = - create_tun_off_runtime(interface_name.clone(), Some(remote_addr), local_addr).await?; - let local_ip = local_ip.unwrap_or("0.0.0.0".parse().unwrap()); - let stack = stack::Stack::new(tun, local_ip, local_ip6, mac); - let driver_type = stack.driver_type(); - - let socket = stack - .try_alloc_established_socket(local_addr, remote_addr, stack::State::SynSent) - .ok_or(TunnelError::InternalError( - "FakeTCP stack closed while allocating socket".into(), - ))?; - - let os_stream = os_socket.connect(remote_addr).await?; - - tracing::info!(?remote_addr, "FakeTcpTunnelConnector connecting"); - - let mut buf = BytesMut::new(); - socket - .recv(&mut buf) - .await - .ok_or(TunnelError::InternalError( - "Failed to recv bytes to establish connection".into(), - ))?; - - tracing::info!(local_addr = ?socket.local_addr(), "FakeTcpTunnelConnector connected"); - - let info = TunnelInfo { - tunnel_type: get_faketcp_tunnel_type_str(driver_type), - local_addr: Some( - crate::tunnel::build_url_from_socket_addr( - &socket.local_addr().to_string(), - "faketcp", - ) - .into(), - ), - remote_addr: Some(self.addr.clone().into()), - resolved_remote_addr: Some( - crate::tunnel::build_url_from_socket_addr(&remote_addr.to_string(), "faketcp") - .into(), - ), - }; - - let socket = Arc::new(socket); - let reader = FakeTcpStream::new(socket.clone()); - let writer = FakeTcpSink::new(socket); - - Ok(Box::new(TunnelWrapper::new_with_associate_data( - reader, - writer, - Some(info), - Some(Box::new((build_os_socket_reader_task(os_stream), stack))), - ))) - } - - fn remote_url(&self) -> url::Url { - self.addr.clone() - } - - fn set_resolved_addr(&mut self, addr: SocketAddr) { - self.resolved_addr = Some(addr); - } - - fn set_socket_mark(&mut self, socket_mark: Option) { - self.socket_mark = socket_mark; - } -} - -type RecvFut = Pin> + Send + Sync>>; - -enum FakeTcpStreamState { - ConsumingBuf(BytesMut), - PollFuture(RecvFut), - Closed, -} - -struct FakeTcpStream { - socket: Arc, - state: FakeTcpStreamState, -} - -impl FakeTcpStream { - fn new(socket: Arc) -> Self { - Self { - socket, - state: FakeTcpStreamState::ConsumingBuf(BytesMut::new()), - } - } -} - -impl Stream for FakeTcpStream { - type Item = StreamItem; - - fn poll_next(self: Pin<&mut Self>, cx: &mut TaskContext<'_>) -> Poll> { - let s = self.get_mut(); - loop { - let state = std::mem::replace(&mut s.state, FakeTcpStreamState::Closed); - match state { - FakeTcpStreamState::ConsumingBuf(buf) => { - let buf_len = buf.len(); - // check peer manager header and split buf out - let packet = ZCPacket::new_from_buf(buf, ZCPacketType::TCP); - if let Some(tcp_hdr) = packet.tcp_tunnel_header() { - let expected_payload_len = tcp_hdr.len.get() as usize; - let min_packet_len = TCP_TUNNEL_HEADER_SIZE + PEER_MANAGER_HEADER_SIZE; - if expected_payload_len < min_packet_len { - tracing::warn!( - "drop fake tcp packet with invalid length: expected_payload_len={}, min_required={}", - expected_payload_len, - min_packet_len - ); - s.state = FakeTcpStreamState::Closed; - return Poll::Ready(None); - } - - if expected_payload_len <= buf_len { - let mut buf = packet.inner(); - let new_inner = buf.split_to(expected_payload_len); - s.state = FakeTcpStreamState::ConsumingBuf(buf); - return Poll::Ready(Some(Ok(ZCPacket::new_from_buf( - new_inner, - ZCPacketType::TCP, - )))); - } - } - - let mut buf = packet.inner(); - buf.truncate(0); - - let socket = s.socket.clone(); - s.state = FakeTcpStreamState::PollFuture(Box::pin(async move { - let ret = socket.recv(&mut buf).await; - ret.map(|s| (buf, s)) - })); - } - FakeTcpStreamState::PollFuture(mut fut) => match fut.as_mut().poll(cx) { - Poll::Ready(Some((buf, _sz))) => { - s.state = FakeTcpStreamState::ConsumingBuf(buf); - } - Poll::Ready(None) => { - s.state = FakeTcpStreamState::Closed; - } - Poll::Pending => { - s.state = FakeTcpStreamState::PollFuture(fut); - return Poll::Pending; - } - }, - FakeTcpStreamState::Closed => { - return Poll::Ready(None); - } - } - } - } -} - -struct FakeTcpSink { - socket: Arc, -} - -impl FakeTcpSink { - fn new(socket: Arc) -> Self { - Self { socket } - } -} - -impl Sink for FakeTcpSink { - type Error = SinkError; - - fn poll_ready( - self: Pin<&mut Self>, - _cx: &mut TaskContext<'_>, - ) -> Poll> { - Poll::Ready(Ok(())) - } - - fn start_send(self: Pin<&mut Self>, item: SinkItem) -> Result<(), Self::Error> { - // We need to send the packet as bytes - // The item is ZCPacket, which has into_bytes() method - let mut packet = item.convert_type(ZCPacketType::TCP); - let len = packet.buf_len(); - packet.mut_tcp_tunnel_header().unwrap().len.set(len as u32); - self.socket.try_send(&packet.into_bytes()); - - Ok(()) - } - - fn poll_flush( - self: Pin<&mut Self>, - _cx: &mut TaskContext<'_>, - ) -> Poll> { - Poll::Ready(Ok(())) - } - - fn poll_close( - self: Pin<&mut Self>, - _cx: &mut TaskContext<'_>, - ) -> Poll> { - self.socket.close(); - Poll::Ready(Ok(())) - } -} - -#[cfg(test)] -mod tests { - use crate::tunnel::common::tests::_tunnel_pingpong; - - use super::*; - - #[tokio::test] - async fn faketcp_pingpong() { - #[cfg(target_family = "unix")] - { - if unsafe { nix::libc::geteuid() } != 0 { - return; - } - } - - let listener = FakeTcpTunnelListener::new("faketcp://0.0.0.0:31011".parse().unwrap()); - let connector = FakeTcpTunnelConnector::new("faketcp://127.0.0.1:31011".parse().unwrap()); - - _tunnel_pingpong(listener, connector).await - } -} diff --git a/easytier/src/tunnel/insecure_tls.rs b/easytier/src/tunnel/insecure_tls.rs deleted file mode 100644 index a4933ecf..00000000 --- a/easytier/src/tunnel/insecure_tls.rs +++ /dev/null @@ -1,86 +0,0 @@ -use std::sync::Arc; - -use rustls::pki_types::{CertificateDer, PrivateKeyDer, ServerName, UnixTime}; - -/// Dummy certificate verifier that treats any certificate as valid. -/// NOTE, such verification is vulnerable to MITM attacks, but convenient for testing. -#[derive(Debug)] -struct SkipServerVerification(Arc); - -impl SkipServerVerification { - fn new(provider: Arc) -> Arc { - Arc::new(Self(provider)) - } -} - -impl rustls::client::danger::ServerCertVerifier for SkipServerVerification { - fn verify_server_cert( - &self, - _end_entity: &CertificateDer<'_>, - _intermediates: &[CertificateDer<'_>], - _server_name: &ServerName<'_>, - _ocsp: &[u8], - _now: UnixTime, - ) -> Result { - Ok(rustls::client::danger::ServerCertVerified::assertion()) - } - - fn verify_tls12_signature( - &self, - message: &[u8], - cert: &CertificateDer<'_>, - dss: &rustls::DigitallySignedStruct, - ) -> Result { - rustls::crypto::verify_tls12_signature( - message, - cert, - dss, - &self.0.signature_verification_algorithms, - ) - } - - fn verify_tls13_signature( - &self, - message: &[u8], - cert: &CertificateDer<'_>, - dss: &rustls::DigitallySignedStruct, - ) -> Result { - rustls::crypto::verify_tls13_signature( - message, - cert, - dss, - &self.0.signature_verification_algorithms, - ) - } - - fn supported_verify_schemes(&self) -> Vec { - self.0.signature_verification_algorithms.supported_schemes() - } -} - -pub fn init_crypto_provider() { - let _ = - rustls::crypto::CryptoProvider::install_default(rustls::crypto::ring::default_provider()); -} - -pub fn get_insecure_tls_client_config() -> rustls::ClientConfig { - init_crypto_provider(); - let provider = rustls::crypto::CryptoProvider::get_default().unwrap(); - let mut config = rustls::ClientConfig::builder() - .dangerous() - .with_custom_certificate_verifier(SkipServerVerification::new(provider.clone())) - .with_no_client_auth(); - config.enable_sni = true; - config.enable_early_data = false; - config -} - -pub fn get_insecure_tls_cert<'a>() -> (Vec>, PrivateKeyDer<'a>) { - let cert = rcgen::generate_simple_self_signed(vec!["localhost".into()]).unwrap(); - let cert_der = cert.serialize_der().unwrap(); - let priv_key = cert.serialize_private_key_der(); - let priv_key = rustls::pki_types::PrivatePkcs8KeyDer::from(priv_key); - let cert_chain = vec![cert_der.into()]; - - (cert_chain, priv_key.into()) -} diff --git a/easytier/src/tunnel/mod.rs b/easytier/src/tunnel/mod.rs index e3e9de1b..5106c231 100644 --- a/easytier/src/tunnel/mod.rs +++ b/easytier/src/tunnel/mod.rs @@ -1,34 +1,15 @@ -use std::{ - collections::hash_map::DefaultHasher, hash::Hasher, net::SocketAddr, pin::Pin, sync::Arc, -}; +use std::{collections::hash_map::DefaultHasher, hash::Hasher, net::SocketAddr}; -use crate::{ - common::{dns::socket_addrs, error::Error}, - proto::common::TunnelInfo, -}; -use async_trait::async_trait; +#[cfg(any(feature = "faketcp", feature = "websocket", feature = "wireguard"))] +use crate::common::dns::socket_addrs; +use crate::common::error::Error; use derive_more::{From, TryInto}; -use futures::{Sink, Stream}; -use socket2::Protocol; -use std::fmt::Debug; +#[cfg(any(feature = "faketcp", feature = "websocket", feature = "wireguard"))] +use easytier_core::tunnel::{IpVersion, TunnelError}; use strum::{Display, EnumString, IntoStaticStr, VariantArray}; -use tokio::time::error::Elapsed; -use self::packet_def::ZCPacket; - -pub mod buf; pub mod common; -pub mod filter; -pub mod mpsc; -pub mod packet_def; -pub mod ring; -pub mod stats; -pub mod tcp; -pub mod udp; -pub(crate) mod udp_src; - -#[cfg(feature = "faketcp")] -pub mod fake_tcp; +pub(crate) mod protocol; #[cfg(feature = "wireguard")] pub mod wireguard; @@ -39,116 +20,6 @@ pub mod quic; #[cfg(feature = "websocket")] pub mod websocket; -#[cfg(any(feature = "quic", feature = "websocket"))] -pub mod insecure_tls; - -#[cfg(unix)] -pub mod unix; - -#[derive(thiserror::Error, Debug)] -pub enum TunnelError { - #[error("io error: {0}")] - IOError(#[from] std::io::Error), - #[error("invalid packet. msg: {0}")] - InvalidPacket(String), - #[error("exceed max packet size. max: {0}, input: {1}")] - ExceedMaxPacketSize(usize, usize), - - #[error("invalid protocol: {0}")] - InvalidProtocol(String), - #[error("invalid addr: {0}")] - InvalidAddr(String), - - #[error("internal error {0}")] - InternalError(String), - - #[error("conn id not match, expect: {0}, actual: {1}")] - ConnIdNotMatch(u32, u32), - #[error("buffer full")] - BufferFull, - - #[error("timeout")] - Timeout(#[from] Elapsed), - - #[error("anyhow error: {0}")] - Anyhow(#[from] anyhow::Error), - - #[error("shutdown")] - Shutdown, - - #[error("no dns record found")] - NoDnsRecordFound(IpVersion), - - #[cfg(feature = "websocket")] - #[error("websocket error: {0}")] - WebSocketError(#[from] tokio_websockets::Error), - - #[error("tunnel error: {0}")] - TunError(String), -} - -pub type StreamT = packet_def::ZCPacket; -pub type StreamItem = Result; -pub type SinkItem = packet_def::ZCPacket; -pub type SinkError = TunnelError; - -pub trait ZCPacketStream: Stream + Send {} -impl ZCPacketStream for T where T: Stream + Send {} -pub trait ZCPacketSink: Sink + Send {} -impl ZCPacketSink for T where T: Sink + Send {} - -pub type SplitTunnel = (Pin>, Pin>); - -#[auto_impl::auto_impl(Box, Arc)] -pub trait Tunnel: Send { - fn split(&self) -> SplitTunnel; - fn info(&self) -> Option; -} - -#[auto_impl::auto_impl(Arc)] -pub trait TunnelConnCounter: 'static + Send + Sync + Debug { - fn get(&self) -> Option; -} - -#[derive(Debug, Clone, Copy, PartialEq)] -pub enum IpVersion { - V4, - V6, - Both, -} - -#[async_trait] -#[auto_impl::auto_impl(Box)] -pub trait TunnelListener: Send { - async fn listen(&mut self) -> Result<(), TunnelError>; - async fn accept(&mut self) -> Result, TunnelError>; - fn local_url(&self) -> url::Url; - fn get_conn_counter(&self) -> Arc> { - #[derive(Debug)] - struct FakeTunnelConnCounter {} - impl TunnelConnCounter for FakeTunnelConnCounter { - fn get(&self) -> Option { - None - } - } - Arc::new(Box::new(FakeTunnelConnCounter {})) - } -} - -#[async_trait] -#[auto_impl::auto_impl(Box, &mut)] -pub trait TunnelConnector: Send { - async fn connect(&mut self) -> Result, TunnelError>; - fn remote_url(&self) -> url::Url; - fn set_bind_addrs(&mut self, _addrs: Vec) {} - fn set_ip_version(&mut self, _ip_version: IpVersion) {} - fn set_resolved_addr(&mut self, _addr: SocketAddr) {} - /// Linux SO_MARK to apply to outbound sockets. `None` leaves SO_MARK - /// untouched; `Some(mark)` applies that exact value (including `Some(0)`). - /// Default impl is a no-op; IP-based connectors override. - fn set_socket_mark(&mut self, _socket_mark: Option) {} -} - pub fn build_url_from_socket_addr(addr: &String, scheme: &str) -> url::Url { if let Ok(sock_addr) = addr.parse::() { let url_str = format!("{}://0.0.0.0", scheme); @@ -162,31 +33,8 @@ pub fn build_url_from_socket_addr(addr: &String, scheme: &str) -> url::Url { } } -impl std::fmt::Debug for dyn Tunnel { - fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { - f.debug_struct("Tunnel") - .field("info", &self.info()) - .finish() - } -} - -impl std::fmt::Debug for dyn TunnelConnector { - fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { - f.debug_struct("TunnelConnector") - .field("remote_url", &self.remote_url()) - .finish() - } -} - -impl std::fmt::Debug for dyn TunnelListener { - fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { - f.debug_struct("TunnelListener") - .field("local_url", &self.local_url()) - .finish() - } -} - #[async_trait::async_trait] +#[cfg(any(feature = "faketcp", feature = "websocket", feature = "wireguard"))] pub(crate) trait FromUrl { async fn from_url(url: url::Url, ip_version: IpVersion) -> Result where @@ -194,6 +42,7 @@ pub(crate) trait FromUrl { } #[async_trait::async_trait] +#[cfg(any(feature = "faketcp", feature = "websocket", feature = "wireguard"))] impl FromUrl for SocketAddr { async fn from_url(url: url::Url, ip_version: IpVersion) -> Result { let addrs = socket_addrs(&url, || { @@ -229,15 +78,6 @@ impl FromUrl for SocketAddr { } } -#[async_trait::async_trait] -impl FromUrl for uuid::Uuid { - async fn from_url(url: url::Url, _ip_version: IpVersion) -> Result { - let o = url.host_str().unwrap(); - let o = uuid::Uuid::parse_str(o).map_err(|e| TunnelError::InvalidAddr(e.to_string()))?; - Ok(o) - } -} - pub struct TunnelUrl { inner: url::Url, } @@ -284,12 +124,6 @@ pub fn generate_digest_from_str(str1: &str, str2: &str, digest: &mut [u8]) { } } -#[derive(Debug, Clone, Copy)] -struct IpSchemeAttributes { - protocol: Protocol, - port_offset: u16, -} - #[derive(Debug, Clone, Copy, PartialEq, Display, EnumString, IntoStaticStr, VariantArray)] #[strum(serialize_all = "lowercase")] pub enum IpScheme { @@ -308,42 +142,16 @@ pub enum IpScheme { } impl IpScheme { - const fn attributes(self) -> IpSchemeAttributes { - let (protocol, port_offset) = match self { - Self::Tcp => (Protocol::TCP, 0), - Self::Udp => (Protocol::UDP, 0), - #[cfg(feature = "wireguard")] - Self::Wg => (Protocol::UDP, 1), - #[cfg(feature = "quic")] - Self::Quic => (Protocol::UDP, 2), - #[cfg(feature = "websocket")] - Self::Ws => (Protocol::TCP, 1), - #[cfg(feature = "websocket")] - Self::Wss => (Protocol::TCP, 2), - #[cfg(feature = "faketcp")] - Self::FakeTcp => (Protocol::TCP, 3), - }; - IpSchemeAttributes { - protocol, - port_offset, - } - } - pub const fn protocol(self) -> Protocol { - self.attributes().protocol + pub fn port_offset(self) -> u16 { + let scheme: &'static str = self.into(); + easytier_core::connectivity::protocol::protocol_port_offset(scheme) + .expect("IpScheme must have core protocol metadata") } - pub const fn port_offset(self) -> u16 { - self.attributes().port_offset - } - - pub const fn default_port(self) -> u16 { - match self { - #[cfg(feature = "websocket")] - Self::Ws => 80, - #[cfg(feature = "websocket")] - Self::Wss => 443, - _ => 11010 + self.port_offset(), - } + pub fn default_port(self) -> u16 { + let scheme: &'static str = self.into(); + easytier_core::connectivity::protocol::protocol_default_port(scheme) + .expect("IpScheme must have core protocol metadata") } } @@ -376,49 +184,3 @@ impl TryFrom<&url::Url> for TunnelScheme { }) } } - -pub(crate) fn get_scheme_by_url(l: &url::Url) -> Result { - l.try_into() -} - -macro_rules! __matches_scheme__ { - ($url:expr, $( $pattern:pat_param )|+ ) => { - matches!($crate::tunnel::get_scheme_by_url(&$url), Ok($( $pattern )|+)) - }; -} - -pub(crate) use __matches_scheme__ as matches_scheme; - -pub fn get_protocol_by_url(l: &url::Url) -> Result { - let TunnelScheme::Ip(scheme) = l.try_into()? else { - return Err(Error::InvalidUrl(l.to_string())); - }; - Ok(scheme.protocol()) -} - -macro_rules! __matches_protocol__ { - ($url:expr, $( $pattern:pat_param )|+ ) => { - matches!($crate::tunnel::get_protocol_by_url($url), Ok($( $pattern )|+)) - }; -} - -pub(crate) use __matches_protocol__ as matches_protocol; - -#[cfg(test)] -mod tests { - use super::{IpScheme, TunnelScheme, matches_scheme}; - - #[test] - fn matches_scheme_accepts_owned_url() { - let url: url::Url = "udp://[2001:db8::1]:11010".parse().unwrap(); - - assert!(matches_scheme!(url, TunnelScheme::Ip(IpScheme::Udp))); - } - - #[test] - fn matches_scheme_accepts_borrowed_url() { - let url: url::Url = "udp://[2001:db8::1]:11010".parse().unwrap(); - - assert!(matches_scheme!(&url, TunnelScheme::Ip(IpScheme::Udp))); - } -} diff --git a/easytier/src/tunnel/mpsc.rs b/easytier/src/tunnel/mpsc.rs deleted file mode 100644 index e15231ae..00000000 --- a/easytier/src/tunnel/mpsc.rs +++ /dev/null @@ -1,256 +0,0 @@ -// this mod wrap tunnel to a mpsc tunnel, based on crossbeam_channel - -use std::{pin::Pin, time::Duration}; - -use anyhow::Context; -use tokio::time::timeout; - -use crate::proto::common::TunnelInfo; - -use super::{Tunnel, TunnelError, ZCPacketSink, ZCPacketStream, packet_def::ZCPacket}; - -use tokio::sync::mpsc::{Receiver, Sender, channel, error::TrySendError}; -use tokio_util::task::AbortOnDropHandle; -// use tachyonix::{channel, Receiver, Sender, TrySendError}; - -use futures::SinkExt; - -#[derive(Clone)] -pub struct MpscTunnelSender(Sender); - -impl MpscTunnelSender { - pub async fn send(&self, item: ZCPacket) -> Result<(), TunnelError> { - self.0.send(item).await.with_context(|| "send error")?; - Ok(()) - } - - pub fn try_send(&self, item: ZCPacket) -> Result<(), TunnelError> { - self.0.try_send(item).map_err(|e| match e { - TrySendError::Full(_) => TunnelError::BufferFull, - TrySendError::Closed(_) => TunnelError::Shutdown, - }) - } -} - -pub struct MpscTunnel { - tx: Option>, - - tunnel: T, - stream: Option>>, - - task: AbortOnDropHandle<()>, -} - -impl MpscTunnel { - pub fn new(tunnel: T, send_timeout: Option) -> Self { - let (tx, mut rx) = channel(32); - let (stream, mut sink) = tunnel.split(); - - let task = tokio::spawn(async move { - loop { - if let Err(e) = Self::forward_one_round(&mut rx, &mut sink, send_timeout).await { - tracing::error!(?e, "forward error"); - break; - } - } - rx.close(); - let close_ret = timeout(Duration::from_secs(5), sink.close()).await; - tracing::warn!(?close_ret, "mpsc close sink"); - }); - - Self { - tx: Some(tx), - tunnel, - stream: Some(stream), - task: AbortOnDropHandle::new(task), - } - } - - async fn forward_one_round( - rx: &mut Receiver, - sink: &mut Pin>, - send_timeout_ms: Option, - ) -> Result<(), TunnelError> { - let item = rx.recv().await.with_context(|| "recv error")?; - if let Some(timeout_ms) = send_timeout_ms { - Self::forward_one_round_with_timeout(rx, sink, item, timeout_ms).await - } else { - Self::forward_one_round_no_timeout(rx, sink, item).await - } - } - - async fn forward_one_round_no_timeout( - rx: &mut Receiver, - sink: &mut Pin>, - initial_item: ZCPacket, - ) -> Result<(), TunnelError> { - sink.feed(initial_item).await?; - - while let Ok(item) = rx.try_recv() { - if let Err(e) = sink.feed(item).await { - tracing::error!(?e, "feed error"); - return Err(e); - } - } - - sink.flush().await - } - - async fn forward_one_round_with_timeout( - rx: &mut Receiver, - sink: &mut Pin>, - initial_item: ZCPacket, - timeout_ms: Duration, - ) -> Result<(), TunnelError> { - match timeout(timeout_ms, async move { - Self::forward_one_round_no_timeout(rx, sink, initial_item).await - }) - .await - { - Ok(Ok(_)) => Ok(()), - Ok(Err(e)) => { - tracing::error!(?e, "forward error"); - Err(e) - } - Err(e) => { - tracing::error!(?e, "forward timeout"); - Err(e.into()) - } - } - } - - pub fn get_stream(&mut self) -> Pin> { - self.stream.take().unwrap() - } - - pub fn get_sink(&self) -> MpscTunnelSender { - MpscTunnelSender(self.tx.as_ref().unwrap().clone()) - } - - pub fn close(&mut self) { - self.tx.take(); - self.task.abort(); - } - - pub fn tunnel_info(&self) -> Option { - self.tunnel.info() - } -} - -#[cfg(test)] -mod tests { - use futures::StreamExt; - - use crate::tunnel::{ - TunnelConnector, TunnelListener, - ring::{RING_TUNNEL_CAP, create_ring_tunnel_pair}, - tcp::{TcpTunnelConnector, TcpTunnelListener}, - }; - - use super::*; - // test slow send lock in framed tunnel - #[tokio::test] - async fn mpsc_slow_receiver() { - let mut listener = TcpTunnelListener::new("tcp://127.0.0.1:11014".parse().unwrap()); - let mut connector = TcpTunnelConnector::new("tcp://127.0.0.1:11014".parse().unwrap()); - - listener.listen().await.unwrap(); - let t1 = tokio::spawn(async move { - let t = listener.accept().await.unwrap(); - let (mut stream, _sink) = t.split(); - let now = tokio::time::Instant::now(); - - let mut a_counter = 0; - let mut b_counter = 0; - - while let Some(Ok(msg)) = stream.next().await { - tokio::time::sleep(tokio::time::Duration::from_millis(100)).await; - if now.elapsed().as_secs() > 5 { - break; - } - - if msg.payload() == "hello".as_bytes() { - a_counter += 1; - } else if msg.payload() == "hello2".as_bytes() { - b_counter += 1; - } - } - - tracing::info!("t1 exit"); - assert_ne!(a_counter, 0); - assert_ne!(b_counter, 0); - }); - - let tunnel = connector.connect().await.unwrap(); - let mpsc_tunnel = MpscTunnel::new(tunnel, None); - - let sink1 = mpsc_tunnel.get_sink(); - let t2 = tokio::spawn(async move { - for i in 0..1000000 { - tokio::time::sleep(tokio::time::Duration::from_millis(50)).await; - let a = sink1 - .send(ZCPacket::new_with_payload("hello".as_bytes())) - .await; - if a.is_err() { - tracing::info!(?a, "t2 exit with err"); - break; - } - - if i % 5000 == 0 { - tracing::info!(i, "send2 1000"); - } - } - - tracing::info!("t2 exit"); - }); - - let sink2 = mpsc_tunnel.get_sink(); - let t3 = tokio::spawn(async move { - for i in 0..1000000 { - tokio::time::sleep(tokio::time::Duration::from_millis(100)).await; - let a = sink2 - .send(ZCPacket::new_with_payload("hello2".as_bytes())) - .await; - if a.is_err() { - tracing::info!(?a, "t3 exit with err"); - break; - } - - if i % 5000 == 0 { - tracing::info!(i, "send2 1000"); - } - } - - tracing::info!("t3 exit"); - }); - - let t4 = tokio::spawn(async move { - tokio::time::sleep(tokio::time::Duration::from_secs(5)).await; - tracing::info!("closing"); - drop(mpsc_tunnel); - tracing::info!("closed"); - }); - - let _ = tokio::join!(t1, t2, t3, t4); - } - - #[tokio::test] - async fn mpsc_slow_receiver_with_send_timeout() { - let (a, _b) = create_ring_tunnel_pair(); - let mpsc_tunnel = MpscTunnel::new(a, Some(Duration::from_secs(1))); - let s = mpsc_tunnel.get_sink(); - for _ in 0..RING_TUNNEL_CAP { - s.send(ZCPacket::new_with_payload(&[0; 1024])) - .await - .unwrap(); - } - tokio::time::sleep(Duration::from_millis(1500)).await; - let e = s.send(ZCPacket::new_with_payload(&[0; 1024])).await; - assert!(e.is_ok()); - - tokio::time::sleep(Duration::from_millis(1500)).await; - - let e = s.send(ZCPacket::new_with_payload(&[0; 1024])).await; - assert!(e.is_err()); - } -} diff --git a/easytier/src/tunnel/protocol.rs b/easytier/src/tunnel/protocol.rs new file mode 100644 index 00000000..1666cd4b --- /dev/null +++ b/easytier/src/tunnel/protocol.rs @@ -0,0 +1,481 @@ +use std::sync::Arc; + +use async_trait::async_trait; +use easytier_core::{ + connectivity::{ + protocol::{ + ClientProtocolUpgrader, CoreClientProtocolConfig, CoreClientProtocolUpgrader, + CoreServerProtocolConfig, CoreServerProtocolUpgrader, ServerProtocolAdmission, + ServerProtocolUpgrade, ServerProtocolUpgrader, + }, + transport::ConnectedTransport, + }, + socket::udp::UdpSession, + tunnel::Tunnel, +}; + +use crate::{common::global_ctx::ArcGlobalCtx, socket::tcp::RuntimeTcpSocket}; + +mod adapters; + +pub(crate) struct RuntimeClientProtocolUpgrader { + adapters: Vec, +} + +pub(crate) struct RuntimeServerProtocolUpgrader { + adapters: Vec, +} + +fn runtime_client_protocol_adapter(global_ctx: &ArcGlobalCtx) -> RuntimeClientProtocolUpgrader { + RuntimeClientProtocolUpgrader { + adapters: adapters::client_adapters(global_ctx), + } +} + +fn runtime_server_protocol_adapter(global_ctx: &ArcGlobalCtx) -> RuntimeServerProtocolUpgrader { + RuntimeServerProtocolUpgrader { + adapters: adapters::server_adapters(global_ctx), + } +} + +pub(crate) fn runtime_client_protocol_upgrader( + global_ctx: ArcGlobalCtx, +) -> Arc> { + Arc::new(CoreClientProtocolUpgrader::with_external( + CoreClientProtocolConfig { + unix: cfg!(unix), + faketcp: cfg!(feature = "faketcp"), + }, + Arc::new(runtime_client_protocol_adapter(&global_ctx)), + )) +} + +pub(crate) fn runtime_server_protocol_upgrader( + global_ctx: ArcGlobalCtx, +) -> Arc> { + Arc::new(CoreServerProtocolUpgrader::with_external( + CoreServerProtocolConfig { + unix: cfg!(unix), + faketcp: cfg!(feature = "faketcp"), + }, + Arc::new(runtime_server_protocol_adapter(&global_ctx)), + )) +} + +#[async_trait] +impl ClientProtocolUpgrader for RuntimeClientProtocolUpgrader { + fn supports_scheme(&self, scheme: &str) -> bool { + self.adapters + .iter() + .any(|adapter| adapter.supports_scheme(scheme)) + } + + fn connect_timeout(&self, scheme: &str) -> Option { + self.adapters + .iter() + .find(|adapter| adapter.supports_scheme(scheme)) + .and_then(|adapter| adapter.connect_timeout(scheme)) + } + + async fn upgrade_client( + &self, + connected: ConnectedTransport, + requested_url: url::Url, + ) -> anyhow::Result> { + let scheme = requested_url.scheme().to_owned(); + let adapter = self + .adapters + .iter() + .find(|adapter| adapter.supports_scheme(&scheme)) + .ok_or_else(|| anyhow::anyhow!("unsupported client protocol upgrader: {scheme}"))?; + adapter.upgrade_client(connected, requested_url).await + } +} + +#[async_trait] +impl ServerProtocolUpgrader for RuntimeServerProtocolUpgrader { + fn supports_scheme(&self, scheme: &str) -> bool { + self.adapters + .iter() + .any(|adapter| adapter.supports_scheme(scheme)) + } + + fn max_pending_tcp_upgrades(&self, scheme: &str) -> Option { + self.adapters + .iter() + .find(|adapter| adapter.supports_scheme(scheme)) + .and_then(|adapter| adapter.max_pending_tcp_upgrades(scheme)) + } + + async fn upgrade_tcp( + &self, + socket: RuntimeTcpSocket, + local_url: url::Url, + ) -> anyhow::Result { + let scheme = local_url.scheme().to_owned(); + let adapter = self + .adapters + .iter() + .find(|adapter| adapter.supports_scheme(&scheme)) + .ok_or_else(|| { + anyhow::anyhow!("unsupported native TCP server protocol upgrader: {scheme}") + })?; + adapter.upgrade_tcp(socket, local_url).await + } + + async fn upgrade_udp( + &self, + session: UdpSession, + local_url: url::Url, + admission: Option, + ) -> anyhow::Result { + let scheme = local_url.scheme().to_owned(); + let adapter = self + .adapters + .iter() + .find(|adapter| adapter.supports_scheme(&scheme)) + .ok_or_else(|| { + anyhow::anyhow!("unsupported native UDP server protocol upgrader: {scheme}") + })?; + adapter.upgrade_udp(session, local_url, admission).await + } + + async fn upgrade_byte_stream( + &self, + socket: RuntimeTcpSocket, + local_url: url::Url, + remote_url: Option, + ) -> anyhow::Result { + let scheme = local_url.scheme().to_owned(); + let adapter = self + .adapters + .iter() + .find(|adapter| adapter.supports_scheme(&scheme)) + .ok_or_else(|| { + anyhow::anyhow!("unsupported native byte-stream server protocol upgrader: {scheme}") + })?; + adapter + .upgrade_byte_stream(socket, local_url, remote_url) + .await + } +} + +#[cfg(test)] +mod tests { + use crate::common::global_ctx::tests::get_mock_global_ctx; + + use super::*; + + #[tokio::test] + async fn protocol_capabilities_follow_enabled_features() { + let global_ctx = get_mock_global_ctx(); + let external = runtime_client_protocol_adapter(&global_ctx); + + assert!(!external.supports_scheme("tcp")); + assert!(!external.supports_scheme("faketcp")); + assert_eq!(external.supports_scheme("ws"), cfg!(feature = "websocket")); + assert_eq!(external.supports_scheme("wss"), cfg!(feature = "websocket")); + assert_eq!(external.supports_scheme("wg"), cfg!(feature = "wireguard")); + assert_eq!(external.supports_scheme("quic"), cfg!(feature = "quic")); + + let upgrader = runtime_client_protocol_upgrader(global_ctx.clone()); + + assert!(upgrader.supports_scheme("tcp")); + assert!(upgrader.supports_scheme("udp")); + assert!(upgrader.supports_scheme("ring")); + assert_eq!(upgrader.supports_scheme("unix"), cfg!(unix)); + assert_eq!(upgrader.supports_scheme("ws"), cfg!(feature = "websocket")); + assert_eq!(upgrader.supports_scheme("wss"), cfg!(feature = "websocket")); + assert_eq!(upgrader.supports_scheme("wg"), cfg!(feature = "wireguard")); + assert_eq!(upgrader.supports_scheme("quic"), cfg!(feature = "quic")); + assert_eq!( + upgrader.supports_scheme("faketcp"), + cfg!(feature = "faketcp") + ); + + let server_external = runtime_server_protocol_adapter(&global_ctx); + assert!(!server_external.supports_scheme("tcp")); + assert!(!server_external.supports_scheme("udp")); + assert!(!server_external.supports_scheme("ring")); + assert_eq!( + server_external.supports_scheme("ws"), + cfg!(feature = "websocket") + ); + assert_eq!( + server_external.max_pending_tcp_upgrades("ws"), + cfg!(feature = "websocket").then_some(std::num::NonZeroUsize::MIN) + ); + assert_eq!( + server_external.supports_scheme("wg"), + cfg!(feature = "wireguard") + ); + assert_eq!( + server_external.supports_scheme("quic"), + cfg!(feature = "quic") + ); + + let server = runtime_server_protocol_upgrader(global_ctx); + assert!(server.supports_scheme("tcp")); + assert!(server.supports_scheme("udp")); + assert!(server.supports_scheme("ring")); + assert_eq!(server.supports_scheme("unix"), cfg!(unix)); + assert_eq!(server.supports_scheme("ws"), cfg!(feature = "websocket")); + assert_eq!(server.supports_scheme("wg"), cfg!(feature = "wireguard")); + assert_eq!(server.supports_scheme("quic"), cfg!(feature = "quic")); + } + + #[cfg(feature = "websocket")] + #[rstest::rstest] + #[case("ws")] + #[case("wss")] + #[tokio::test] + async fn runtime_websocket_upgraders_share_one_native_engine(#[case] scheme: &str) { + use easytier_core::packet::ZCPacket; + use futures::{SinkExt, StreamExt}; + + let listener = tokio::net::TcpListener::bind(("127.0.0.1", 0)) + .await + .unwrap(); + let addr = listener.local_addr().unwrap(); + let url: url::Url = format!("{scheme}://{addr}").parse().unwrap(); + let global_ctx = get_mock_global_ctx(); + let server = runtime_server_protocol_upgrader(global_ctx.clone()); + let client = runtime_client_protocol_upgrader(global_ctx); + + assert_eq!( + client.connect_timeout(scheme), + Some(crate::tunnel::websocket::CONNECT_TIMEOUT) + ); + + let server_url = url.clone(); + let server_task = tokio::spawn(async move { + let (socket, _) = listener.accept().await.unwrap(); + let ServerProtocolUpgrade::Tunnel(tunnel) = server + .upgrade_tcp(RuntimeTcpSocket::new(socket), server_url) + .await + .unwrap() + else { + panic!("WebSocket must upgrade directly to a tunnel"); + }; + crate::tunnel::common::tests::_tunnel_echo_server(tunnel, true).await; + }); + + tokio::time::timeout(std::time::Duration::from_secs(5), async { + let socket = tokio::net::TcpStream::connect(addr).await.unwrap(); + let tunnel = client + .upgrade_client(ConnectedTransport::Tcp(RuntimeTcpSocket::new(socket)), url) + .await + .unwrap(); + let (mut recv, mut send) = tunnel.split(); + send.send(ZCPacket::new_with_payload(b"runtime websocket seam")) + .await + .unwrap(); + let packet = recv.next().await.unwrap().unwrap(); + assert_eq!(packet.payload(), b"runtime websocket seam".as_slice()); + send.close().await.unwrap(); + server_task.await.unwrap(); + }) + .await + .unwrap(); + } + + #[cfg(feature = "websocket")] + #[tokio::test] + async fn runtime_websocket_upgraders_reject_ws_client_for_wss_server() { + let listener = tokio::net::TcpListener::bind(("127.0.0.1", 0)) + .await + .unwrap(); + let addr = listener.local_addr().unwrap(); + let server_url: url::Url = format!("wss://{addr}").parse().unwrap(); + let client_url: url::Url = format!("ws://{addr}").parse().unwrap(); + let global_ctx = get_mock_global_ctx(); + let server = runtime_server_protocol_upgrader(global_ctx.clone()); + let client = runtime_client_protocol_upgrader(global_ctx); + + let server_task = tokio::spawn(async move { + let (socket, _) = listener.accept().await.unwrap(); + assert!( + server + .upgrade_tcp(RuntimeTcpSocket::new(socket), server_url) + .await + .is_err() + ); + }); + + tokio::time::timeout(std::time::Duration::from_secs(5), async { + let socket = tokio::net::TcpStream::connect(addr).await.unwrap(); + assert!( + client + .upgrade_client( + ConnectedTransport::Tcp(RuntimeTcpSocket::new(socket)), + client_url, + ) + .await + .is_err() + ); + server_task.await.unwrap(); + }) + .await + .unwrap(); + } + + #[cfg(feature = "wireguard")] + #[rstest::rstest] + #[case("127.0.0.1:0")] + #[case("[::1]:0")] + #[tokio::test] + async fn runtime_wireguard_upgraders_consume_core_udp_sessions(#[case] bind_addr: &str) { + use crate::{ + common::netns::NetNS, host_runtime::native_host_runtime, + socket::udp::new_runtime_udp_session_listener, + }; + use easytier_core::{ + connectivity::transport::{UdpSessionMode, connect_udp}, + packet::ZCPacket, + socket::SocketListener, + socket::udp::{ + UdpBindOptions, UdpSessionAcceptKind, UdpSessionListenRequest, UdpSessionProtocol, + VirtualUdpSocket, + }, + }; + use futures::{SinkExt, StreamExt}; + + let bind_addr = bind_addr.parse().unwrap(); + let mut listener = new_runtime_udp_session_listener( + format!("wg://{bind_addr}").parse().unwrap(), + UdpSessionListenRequest::new( + UdpBindOptions::port_bound_listener(bind_addr).with_only_v6(bind_addr.is_ipv6()), + ), + UdpSessionAcceptKind::Classified(UdpSessionProtocol::WireGuard), + NetNS::new(None), + ); + listener.listen().await.unwrap(); + let remote_addr = listener.bound_socket().unwrap().local_addr().unwrap(); + let remote_url: url::Url = format!("wg://{remote_addr}").parse().unwrap(); + + let global_ctx = get_mock_global_ctx(); + let server = runtime_server_protocol_upgrader(global_ctx.clone()); + let client = runtime_client_protocol_upgrader(global_ctx); + let server_url = remote_url.clone(); + let server_task = tokio::spawn(async move { + let session = listener.accept().await.unwrap(); + let ServerProtocolUpgrade::Tunnel(tunnel) = + server.upgrade_udp(session, server_url, None).await.unwrap() + else { + panic!("WireGuard must upgrade directly to a tunnel"); + }; + crate::tunnel::common::tests::_tunnel_echo_server(tunnel, false).await; + }); + + tokio::time::timeout(std::time::Duration::from_secs(5), async { + let connected = connect_udp( + native_host_runtime(), + remote_addr, + Vec::new(), + UdpBindOptions::direct_connect(), + UdpSessionMode::Classified(UdpSessionProtocol::WireGuard), + ) + .await + .unwrap(); + let tunnel = client + .upgrade_client(ConnectedTransport::Udp(connected), remote_url) + .await + .unwrap(); + let (mut recv, mut send) = tunnel.split(); + send.send(ZCPacket::new_with_payload(b"runtime WireGuard seam")) + .await + .unwrap(); + let packet = recv.next().await.unwrap().unwrap(); + assert_eq!(packet.payload(), b"runtime WireGuard seam".as_slice()); + let _ = send.close().await; + server_task.abort(); + let _ = server_task.await; + }) + .await + .unwrap(); + } + + #[cfg(feature = "quic")] + #[rstest::rstest] + #[case("127.0.0.1:0")] + #[case("[::1]:0")] + #[tokio::test] + async fn runtime_quic_upgraders_consume_core_udp_sessions(#[case] bind_addr: &str) { + use crate::{ + common::netns::NetNS, host_runtime::native_host_runtime, + socket::udp::new_runtime_udp_session_listener, + }; + use easytier_core::{ + connectivity::{ + protocol::ServerProtocolAdmissionController, + transport::{UdpSessionMode, connect_udp}, + }, + packet::ZCPacket, + socket::SocketListener, + socket::udp::{ + UdpBindOptions, UdpSessionAcceptKind, UdpSessionListenRequest, UdpSessionProtocol, + VirtualUdpSocket, + }, + }; + use futures::{SinkExt, StreamExt}; + + let bind_addr = bind_addr.parse().unwrap(); + let mut listener = new_runtime_udp_session_listener( + format!("quic://{bind_addr}").parse().unwrap(), + UdpSessionListenRequest::new( + UdpBindOptions::port_bound_listener(bind_addr).with_only_v6(bind_addr.is_ipv6()), + ), + UdpSessionAcceptKind::Classified(UdpSessionProtocol::Quic), + NetNS::new(None), + ); + listener.listen().await.unwrap(); + let remote_addr = listener.bound_socket().unwrap().local_addr().unwrap(); + let remote_url: url::Url = format!("quic://{remote_addr}").parse().unwrap(); + + let global_ctx = get_mock_global_ctx(); + let server = runtime_server_protocol_upgrader(global_ctx.clone()); + let client = runtime_client_protocol_upgrader(global_ctx); + let server_url = remote_url.clone(); + let server_task = tokio::spawn(async move { + let session = listener.accept().await.unwrap(); + let admission = ServerProtocolAdmissionController::quic() + .try_admit() + .unwrap(); + let ServerProtocolUpgrade::Acceptor(mut accepted) = server + .upgrade_udp(session, server_url, Some(admission)) + .await + .unwrap() + else { + panic!("QUIC must keep accepting connections from its UDP session"); + }; + let tunnel = accepted.accept().await.unwrap(); + crate::tunnel::common::tests::_tunnel_echo_server(tunnel, false).await; + }); + + tokio::time::timeout(std::time::Duration::from_secs(5), async { + let connected = connect_udp( + native_host_runtime(), + remote_addr, + Vec::new(), + UdpBindOptions::direct_connect(), + UdpSessionMode::Classified(UdpSessionProtocol::Quic), + ) + .await + .unwrap(); + let tunnel = client + .upgrade_client(ConnectedTransport::Udp(connected), remote_url) + .await + .unwrap(); + let (mut recv, mut send) = tunnel.split(); + send.send(ZCPacket::new_with_payload(b"runtime QUIC seam")) + .await + .unwrap(); + let packet = recv.next().await.unwrap().unwrap(); + assert_eq!(packet.payload(), b"runtime QUIC seam".as_slice()); + let _ = send.close().await; + server_task.await.unwrap(); + }) + .await + .unwrap(); + } +} diff --git a/easytier/src/tunnel/protocol/adapters/mod.rs b/easytier/src/tunnel/protocol/adapters/mod.rs new file mode 100644 index 00000000..b1039728 --- /dev/null +++ b/easytier/src/tunnel/protocol/adapters/mod.rs @@ -0,0 +1,43 @@ +use std::sync::Arc; + +use easytier_core::connectivity::protocol::{ClientProtocolUpgrader, ServerProtocolUpgrader}; + +use crate::{common::global_ctx::ArcGlobalCtx, socket::tcp::RuntimeTcpSocket}; + +#[cfg(feature = "quic")] +mod quic; +#[cfg(feature = "websocket")] +mod websocket; +#[cfg(feature = "wireguard")] +mod wireguard; + +pub(super) type ClientAdapter = Arc>; +pub(super) type ServerAdapter = Arc>; + +pub(super) fn client_adapters(global_ctx: &ArcGlobalCtx) -> Vec { + let _ = global_ctx; + [ + #[cfg(feature = "websocket")] + websocket::client_adapter(global_ctx), + #[cfg(feature = "wireguard")] + wireguard::client_adapter(global_ctx), + #[cfg(feature = "quic")] + quic::client_adapter(global_ctx), + ] + .into_iter() + .collect() +} + +pub(super) fn server_adapters(global_ctx: &ArcGlobalCtx) -> Vec { + let _ = global_ctx; + [ + #[cfg(feature = "websocket")] + websocket::server_adapter(global_ctx), + #[cfg(feature = "wireguard")] + wireguard::server_adapter(global_ctx), + #[cfg(feature = "quic")] + quic::server_adapter(global_ctx), + ] + .into_iter() + .collect() +} diff --git a/easytier/src/tunnel/protocol/adapters/quic.rs b/easytier/src/tunnel/protocol/adapters/quic.rs new file mode 100644 index 00000000..9c1e6bc9 --- /dev/null +++ b/easytier/src/tunnel/protocol/adapters/quic.rs @@ -0,0 +1,88 @@ +use std::sync::Arc; + +use async_trait::async_trait; +use easytier_core::{ + connectivity::{ + protocol::{ + ClientProtocolUpgrader, ServerProtocolAdmission, ServerProtocolUpgrade, + ServerProtocolUpgrader, + }, + transport::ConnectedTransport, + }, + socket::udp::UdpSession, + tunnel::Tunnel, +}; + +use crate::{ + common::global_ctx::ArcGlobalCtx, + socket::tcp::RuntimeTcpSocket, + tunnel::quic::{QuicAcceptedSession, upgrade_connected}, +}; + +use super::{ClientAdapter, ServerAdapter}; + +#[derive(Default)] +struct QuicAdapter; + +pub(super) fn client_adapter(_global_ctx: &ArcGlobalCtx) -> ClientAdapter { + Arc::new(QuicAdapter) +} + +pub(super) fn server_adapter(_global_ctx: &ArcGlobalCtx) -> ServerAdapter { + Arc::new(QuicAdapter) +} + +#[async_trait] +impl ClientProtocolUpgrader for QuicAdapter { + fn supports_scheme(&self, scheme: &str) -> bool { + scheme == "quic" + } + + async fn upgrade_client( + &self, + connected: ConnectedTransport, + requested_url: url::Url, + ) -> anyhow::Result> { + let ConnectedTransport::Udp(session) = connected else { + anyhow::bail!("QUIC protocol requires a UDP session"); + }; + Ok(upgrade_connected(session, requested_url).await?) + } +} + +#[async_trait] +impl ServerProtocolUpgrader for QuicAdapter { + fn supports_scheme(&self, scheme: &str) -> bool { + scheme == "quic" + } + + async fn upgrade_tcp( + &self, + _socket: RuntimeTcpSocket, + _local_url: url::Url, + ) -> anyhow::Result { + anyhow::bail!("unsupported native TCP server protocol upgrader: quic") + } + + async fn upgrade_udp( + &self, + session: UdpSession, + local_url: url::Url, + admission: Option, + ) -> anyhow::Result { + let admission = + admission.ok_or_else(|| anyhow::anyhow!("QUIC server admission permit is missing"))?; + Ok(ServerProtocolUpgrade::Acceptor(Box::new( + QuicAcceptedSession::new(session, local_url, admission)?, + ))) + } + + async fn upgrade_byte_stream( + &self, + _socket: RuntimeTcpSocket, + _local_url: url::Url, + _remote_url: Option, + ) -> anyhow::Result { + anyhow::bail!("unsupported native byte-stream server protocol upgrader: quic") + } +} diff --git a/easytier/src/tunnel/protocol/adapters/websocket.rs b/easytier/src/tunnel/protocol/adapters/websocket.rs new file mode 100644 index 00000000..891f2b2b --- /dev/null +++ b/easytier/src/tunnel/protocol/adapters/websocket.rs @@ -0,0 +1,104 @@ +use std::{num::NonZeroUsize, sync::Arc, time::Duration}; + +use async_trait::async_trait; +use easytier_core::{ + connectivity::{ + protocol::{ + ClientProtocolUpgrader, ServerProtocolAdmission, ServerProtocolUpgrade, + ServerProtocolUpgrader, + }, + transport::ConnectedTransport, + }, + socket::udp::UdpSession, + tunnel::Tunnel, +}; + +use crate::{ + common::global_ctx::ArcGlobalCtx, + socket::tcp::RuntimeTcpSocket, + tunnel::websocket::{ + CONNECT_TIMEOUT, SERVER_HANDSHAKE_TIMEOUT, upgrade_accepted, upgrade_connected, + }, +}; + +use super::{ClientAdapter, ServerAdapter}; + +#[derive(Default)] +struct WebSocketAdapter; + +fn supports_scheme(scheme: &str) -> bool { + matches!(scheme, "ws" | "wss") +} + +pub(super) fn client_adapter(_global_ctx: &ArcGlobalCtx) -> ClientAdapter { + Arc::new(WebSocketAdapter) +} + +pub(super) fn server_adapter(_global_ctx: &ArcGlobalCtx) -> ServerAdapter { + Arc::new(WebSocketAdapter) +} + +#[async_trait] +impl ClientProtocolUpgrader for WebSocketAdapter { + fn supports_scheme(&self, scheme: &str) -> bool { + supports_scheme(scheme) + } + + fn connect_timeout(&self, scheme: &str) -> Option { + supports_scheme(scheme).then_some(CONNECT_TIMEOUT) + } + + async fn upgrade_client( + &self, + connected: ConnectedTransport, + requested_url: url::Url, + ) -> anyhow::Result> { + let ConnectedTransport::Tcp(socket) = connected else { + anyhow::bail!("WebSocket protocol requires a TCP transport"); + }; + Ok(upgrade_connected(socket, requested_url).await?) + } +} + +#[async_trait] +impl ServerProtocolUpgrader for WebSocketAdapter { + fn supports_scheme(&self, scheme: &str) -> bool { + supports_scheme(scheme) + } + + fn max_pending_tcp_upgrades(&self, scheme: &str) -> Option { + supports_scheme(scheme).then_some(NonZeroUsize::MIN) + } + + async fn upgrade_tcp( + &self, + socket: RuntimeTcpSocket, + local_url: url::Url, + ) -> anyhow::Result { + Ok(ServerProtocolUpgrade::Tunnel( + tokio::time::timeout( + SERVER_HANDSHAKE_TIMEOUT, + upgrade_accepted(socket, local_url), + ) + .await??, + )) + } + + async fn upgrade_udp( + &self, + _session: UdpSession, + _local_url: url::Url, + _admission: Option, + ) -> anyhow::Result { + anyhow::bail!("WebSocket protocol requires a TCP transport") + } + + async fn upgrade_byte_stream( + &self, + _socket: RuntimeTcpSocket, + _local_url: url::Url, + _remote_url: Option, + ) -> anyhow::Result { + anyhow::bail!("WebSocket protocol requires a TCP transport") + } +} diff --git a/easytier/src/tunnel/protocol/adapters/wireguard.rs b/easytier/src/tunnel/protocol/adapters/wireguard.rs new file mode 100644 index 00000000..a891994a --- /dev/null +++ b/easytier/src/tunnel/protocol/adapters/wireguard.rs @@ -0,0 +1,100 @@ +use std::sync::Arc; + +use async_trait::async_trait; +use easytier_core::{ + connectivity::{ + protocol::{ + ClientProtocolUpgrader, ServerProtocolAdmission, ServerProtocolUpgrade, + ServerProtocolUpgrader, + }, + transport::ConnectedTransport, + }, + socket::udp::UdpSession, + tunnel::Tunnel, +}; + +use crate::{ + common::global_ctx::ArcGlobalCtx, + socket::tcp::RuntimeTcpSocket, + tunnel::wireguard::{WgConfig, upgrade_accepted, upgrade_connected}, +}; + +use super::{ClientAdapter, ServerAdapter}; + +struct WireGuardAdapter { + config: WgConfig, +} + +impl WireGuardAdapter { + fn new(global_ctx: &ArcGlobalCtx) -> Self { + let identity = global_ctx.get_network_identity(); + Self { + config: WgConfig::new_from_network_identity( + &identity.network_name, + &identity.network_secret.unwrap_or_default(), + ), + } + } +} + +pub(super) fn client_adapter(global_ctx: &ArcGlobalCtx) -> ClientAdapter { + Arc::new(WireGuardAdapter::new(global_ctx)) +} + +pub(super) fn server_adapter(global_ctx: &ArcGlobalCtx) -> ServerAdapter { + Arc::new(WireGuardAdapter::new(global_ctx)) +} + +#[async_trait] +impl ClientProtocolUpgrader for WireGuardAdapter { + fn supports_scheme(&self, scheme: &str) -> bool { + scheme == "wg" + } + + async fn upgrade_client( + &self, + connected: ConnectedTransport, + requested_url: url::Url, + ) -> anyhow::Result> { + let ConnectedTransport::Udp(session) = connected else { + anyhow::bail!("WireGuard protocol requires a UDP session"); + }; + Ok(upgrade_connected(session, requested_url, self.config.clone()).await?) + } +} + +#[async_trait] +impl ServerProtocolUpgrader for WireGuardAdapter { + fn supports_scheme(&self, scheme: &str) -> bool { + scheme == "wg" + } + + async fn upgrade_tcp( + &self, + _socket: RuntimeTcpSocket, + _local_url: url::Url, + ) -> anyhow::Result { + anyhow::bail!("unsupported native TCP server protocol upgrader: wg") + } + + async fn upgrade_udp( + &self, + session: UdpSession, + _local_url: url::Url, + _admission: Option, + ) -> anyhow::Result { + Ok(ServerProtocolUpgrade::Tunnel(upgrade_accepted( + session, + self.config.clone(), + )?)) + } + + async fn upgrade_byte_stream( + &self, + _socket: RuntimeTcpSocket, + _local_url: url::Url, + _remote_url: Option, + ) -> anyhow::Result { + anyhow::bail!("unsupported native byte-stream server protocol upgrader: wg") + } +} diff --git a/easytier/src/tunnel/quic.rs b/easytier/src/tunnel/quic.rs index da68c66c..a4ba36b4 100644 --- a/easytier/src/tunnel/quic.rs +++ b/easytier/src/tunnel/quic.rs @@ -2,25 +2,36 @@ //! //! Checkout the `README.md` for guidance. -use super::{FromUrl, IpVersion, Tunnel, TunnelConnector, TunnelError, TunnelListener}; -use crate::common::global_ctx::ArcGlobalCtx; -use crate::tunnel::common::bind; -use crate::tunnel::{ - TunnelInfo, - common::{FramedReader, FramedWriter, TunnelWrapper}, -}; +use crate::proto::common::TunnelInfo; use anyhow::Context; -use derivative::Derivative; -use derive_more::{Deref, DerefMut}; -use parking_lot::RwLock; -use quinn::{ - ClientConfig, ConnectError, Connection, Endpoint, EndpointConfig, ServerConfig, - TransportConfig, congestion::BbrConfig, default_runtime, +use easytier_core::{ + connectivity::{ + protocol::{ServerProtocolAdmission, ServerTunnelAcceptor}, + transport::ConnectedUdpSession, + }, + socket::udp::UdpSession, + tunnel::{ + Tunnel, TunnelError, + framed::{FramedReader, FramedWriter}, + wrapper::TunnelWrapper, + }, +}; +use quinn::{ + AsyncUdpSocket, ClientConfig, Connecting, Connection, Endpoint, EndpointConfig, Incoming, + ServerConfig, TransportConfig, congestion::BbrConfig, default_runtime, }; -use std::net::{Ipv4Addr, Ipv6Addr}; -use std::sync::OnceLock; use std::{net::SocketAddr, sync::Arc, time::Duration}; -use tokio::net::UdpSocket; +use tokio::{ + sync::{ + OwnedSemaphorePermit, Semaphore, + mpsc::{Receiver, Sender, channel}, + }, + task::JoinSet, +}; +use tokio_util::task::AbortOnDropHandle; + +mod session_socket; +pub(crate) use session_socket::QuicUdpSessionSocket; // region config mod crypto { @@ -301,418 +312,11 @@ pub fn endpoint_config() -> EndpointConfig { } //endregion -//region rw pool -#[derive(Derivative)] -#[derivative(Default(bound = ""))] -#[derive(Debug, Deref, DerefMut)] -struct RwPoolInner { - #[deref] - #[deref_mut] - pool: Vec, - enabled: bool, -} - -#[derive(Debug)] -struct RwPool { - ephemeral: RwLock>, - persistent: RwLock>, - capacity: usize, -} - -impl RwPool { - fn new(capacity: usize) -> Self { - Self { - ephemeral: RwLock::new(RwPoolInner::default()), - persistent: RwLock::new(RwPoolInner::default()), - capacity, - } - } - - /// return the capacity of the ephemeral pool; - /// if `ephemeral` or `persistent` is None, read lock `self`'s pool - fn capacity( - &self, - ephemeral: Option<&RwPoolInner>, - persistent: Option<&RwPoolInner>, - ) -> usize { - let guard; - let ephemeral = if let Some(ephemeral) = ephemeral { - ephemeral - } else { - guard = self.ephemeral.read(); - &guard - }; - - let guard; - let persistent = if let Some(persistent) = persistent { - persistent - } else { - guard = self.persistent.read(); - &guard - }; - - (self.capacity * ephemeral.enabled as usize).saturating_sub(persistent.len()) - } - - fn is_full(&self) -> bool { - let pool = self.ephemeral.read(); - pool.len() >= self.capacity(Some(&pool), None) - } - - fn is_enabled(&self) -> bool { - self.ephemeral.read().enabled - } - - fn enable(&self) { - self.ephemeral.write().enabled = true; - self.resize(); - } - - fn disable(&self) { - self.ephemeral.write().enabled = false; - self.resize(); - } - - /// push an item to the persistent pool - fn push(&self, item: Item) { - self.persistent.write().push(item); - self.resize(); - } - - fn len(&self) -> usize { - let persistent_len = self.persistent.read().len(); - let ephemeral_len = self.ephemeral.read().len(); - persistent_len + ephemeral_len - } - - /// try to push an item to the ephemeral pool, return the item if full - fn try_push(&self, item: Item) -> Option { - let mut pool = self.ephemeral.write(); - if pool.len() < self.capacity(Some(&pool), None) { - pool.push(item); - return None; - } - Some(item) - } - - fn resize(&self) { - let resize = { - let pool = self.ephemeral.read(); - pool.capacity() != self.capacity(Some(&pool), None) - }; - if resize { - let mut pool = self.ephemeral.write(); - let capacity = self.capacity(Some(&pool), None); - pool.reserve_exact(capacity); - pool.truncate(capacity); - pool.shrink_to(capacity); - } - } - - fn with_iter(&self, f: F) -> R - where - F: FnOnce(&mut dyn Iterator) -> R, - { - let ephemeral = self.ephemeral.read(); - let persistent = self.persistent.read(); - f(&mut persistent.iter().chain(ephemeral.iter())) - } -} - -impl RwPool { - fn retain_endpoints(&self, mut keep: F) -> usize - where - F: FnMut(&Endpoint) -> bool, - { - let persistent_removed = { - let mut persistent = self.persistent.write(); - let before = persistent.len(); - persistent.retain(|endpoint| keep(endpoint)); - before - persistent.len() - }; - - let ephemeral_removed = { - let mut ephemeral = self.ephemeral.write(); - let before = ephemeral.len(); - ephemeral.retain(|endpoint| keep(endpoint)); - before - ephemeral.len() - }; - - let removed = persistent_removed + ephemeral_removed; - if removed > 0 { - self.resize(); - } - removed - } - - fn remove_by_local_addr(&self, local_addr: SocketAddr) -> usize { - self.retain_endpoints(|endpoint| endpoint.local_addr().ok() != Some(local_addr)) - } - - fn contains_local_addr(&self, local_addr: SocketAddr) -> bool { - self.persistent - .read() - .iter() - .any(|endpoint| endpoint.local_addr().ok() == Some(local_addr)) - || self - .ephemeral - .read() - .iter() - .any(|endpoint| endpoint.local_addr().ok() == Some(local_addr)) - } -} -//endregion - -//region endpoint manager -#[derive(Debug)] -pub struct QuicEndpointManager { - ipv4: RwPool, - ipv6: RwPool, - both: RwPool, -} - -static QUIC_ENDPOINT_MANAGER: OnceLock = OnceLock::new(); - -impl QuicEndpointManager { - fn try_create( - addr: SocketAddr, - dual_stack: bool, - socket_mark: Option, - ) -> Result { - let socket = bind::() - .addr(addr) - .only_v6(addr.is_ipv6() && !dual_stack) - .maybe_socket_mark(socket_mark) - .call()?; - let runtime = default_runtime().ok_or(TunnelError::InternalError( - "no async runtime found".to_owned(), - ))?; - let mut endpoint = Endpoint::new_with_abstract_socket( - endpoint_config(), - None, - runtime.wrap_udp_socket(socket.into_std()?)?, - runtime, - )?; - endpoint.set_default_client_config(client_config()); - Ok(endpoint) - } - - fn create( - &self, - socket_mark: Option, - mut selector: F, - ) -> Result<(&RwPool, Option), TunnelError> - where - F: FnMut(&QuicEndpointManager) -> (&RwPool, Option<(SocketAddr, bool)>), - { - loop { - let (pool, r) = selector(self); - let Some((addr, dual_stack)) = r else { - return Ok((pool, None)); - }; - - let endpoint = Self::try_create(addr, dual_stack, socket_mark); - if let Err(error) = endpoint.as_ref() - && dual_stack - { - tracing::warn!(?error, "create dual stack quic endpoint failed"); - self.both.disable(); - self.ipv4.enable(); - self.ipv6.enable(); - continue; - } - - return Ok((pool, Some(endpoint?))); - } - } -} - -impl QuicEndpointManager { - fn new(capacity: usize) -> Self { - let ipv4 = RwPool::new(capacity.div_ceil(2)); - let ipv6 = RwPool::new(capacity.div_ceil(2)); - let both = RwPool::new(capacity); - both.enable(); - Self { ipv4, ipv6, both } - } - - fn load(global_ctx: &ArcGlobalCtx) -> &Self { - let capacity = global_ctx - .config - .get_flags() - .multi_thread - .then(std::thread::available_parallelism) - .and_then(|r| r.ok()) - .map(|n| n.get()) - .unwrap_or(1); - - let mgr = QUIC_ENDPOINT_MANAGER.get(); - match mgr { - Some(mgr) => { - for pool in [&mgr.ipv4, &mgr.ipv6, &mgr.both] { - pool.resize(); - } - } - None => { - let _ = QUIC_ENDPOINT_MANAGER.set(Self::new(capacity)); - } - } - - QUIC_ENDPOINT_MANAGER.get().unwrap() - } - - fn client_pool(&self, ip_version: IpVersion) -> &RwPool { - let dual_stack = self.both.is_enabled(); - match ip_version { - IpVersion::V4 if !dual_stack => &self.ipv4, - _ => { - if dual_stack { - &self.both - } else { - &self.ipv6 - } - } - } - } - - /// Get a QUIC endpoint to be used as a server - /// - /// # Arguments - /// * `addr`: listen address - fn server(global_ctx: &ArcGlobalCtx, addr: SocketAddr) -> Result { - let mgr = Self::load(global_ctx); - let socket_mark = global_ctx.config.get_flags().socket_mark; - - let (pool, endpoint) = mgr.create(socket_mark, |mgr| { - let dual_stack = addr.ip() == Ipv6Addr::UNSPECIFIED && mgr.both.is_enabled(); - let pool = if addr.is_ipv4() { - &mgr.ipv4 - } else if dual_stack { - &mgr.both - } else { - &mgr.ipv6 - }; - (pool, Some((addr, dual_stack))) - })?; - - let endpoint = endpoint.expect("server endpoint creation should not return None"); - endpoint.set_server_config(Some(server_config())); - pool.push(endpoint.clone()); - - Ok(endpoint) - } - - fn client_endpoint( - &self, - ip_version: IpVersion, - socket_mark: Option, - ) -> Result { - let (pool, endpoint) = self.create(socket_mark, |mgr| { - let dual_stack = mgr.both.is_enabled(); - let (pool, addr) = match ip_version { - IpVersion::V4 if !dual_stack => (&mgr.ipv4, (Ipv4Addr::UNSPECIFIED, 0).into()), - _ => { - let pool = if dual_stack { &mgr.both } else { &mgr.ipv6 }; - (pool, (Ipv6Addr::UNSPECIFIED, 0).into()) - } - }; - if pool.is_full() { - (pool, None) - } else { - (pool, Some((addr, dual_stack))) - } - })?; - - if let Some(endpoint) = endpoint { - pool.try_push(endpoint); - } - - Ok(pool.with_iter(|iter| iter.min_by_key(|e| e.open_connections()).unwrap().clone())) - } - - fn remove_endpoint(&self, endpoint: &Endpoint) -> usize { - let Ok(local_addr) = endpoint.local_addr() else { - return 0; - }; - self.remove_endpoint_by_local_addr(local_addr) - } - - fn remove_endpoint_by_local_addr(&self, local_addr: SocketAddr) -> usize { - [&self.ipv4, &self.ipv6, &self.both] - .into_iter() - .map(|pool| pool.remove_by_local_addr(local_addr)) - .sum() - } - - fn contains_local_addr(&self, local_addr: SocketAddr) -> bool { - [&self.ipv4, &self.ipv6, &self.both] - .into_iter() - .any(|pool| pool.contains_local_addr(local_addr)) - } - - async fn connect( - global_ctx: &ArcGlobalCtx, - addr: SocketAddr, - ) -> Result<(Endpoint, Connection), TunnelError> { - let ip_version = if addr.ip().is_ipv4() { - IpVersion::V4 - } else { - IpVersion::V6 - }; - let socket_mark = global_ctx.config.get_flags().socket_mark; - Self::load(global_ctx) - .connect_with_ip_version(addr, ip_version, socket_mark) - .await - } - - async fn connect_with_ip_version( - &self, - addr: SocketAddr, - ip_version: IpVersion, - socket_mark: Option, - ) -> Result<(Endpoint, Connection), TunnelError> { - let max_endpoint_stopping_retries = self.client_pool(ip_version).len().saturating_add(1); - let mut endpoint_stopping_retries = 0; - - loop { - let endpoint = self.client_endpoint(ip_version, socket_mark)?; - let connecting = match endpoint.connect(addr, "localhost") { - Ok(connecting) => connecting, - Err(ConnectError::EndpointStopping) => { - let local_addr = endpoint.local_addr().ok(); - let removed = self.remove_endpoint(&endpoint); - endpoint_stopping_retries += 1; - tracing::warn!( - ?addr, - ?local_addr, - removed, - "removed stopped quic endpoint and retry connect" - ); - if endpoint_stopping_retries > max_endpoint_stopping_retries { - return Err(anyhow::Error::new(ConnectError::EndpointStopping) - .context(format!("failed to create connection to {}", addr)) - .into()); - } - continue; - } - Err(e) => { - return Err(anyhow::Error::new(e) - .context(format!("failed to create connection to {}", addr)) - .into()); - } - }; - let connection = connecting - .await - .with_context(|| format!("failed to connect to {}", addr))?; - - return Ok((endpoint, connection)); - } - } -} -//endregion +const QUIC_ACCEPT_COMPLETION_TIMEOUT: Duration = Duration::from_secs(10); struct ConnWrapper { conn: Connection, + _endpoint: Endpoint, } impl Drop for ConnWrapper { @@ -721,331 +325,349 @@ impl Drop for ConnWrapper { } } -pub struct QuicTunnelListener { - addr: url::Url, - global_ctx: ArcGlobalCtx, - endpoint: Option, +pub(crate) async fn upgrade_connected( + connected: ConnectedUdpSession, + remote_url: url::Url, +) -> Result, TunnelError> { + let socket = Arc::new(QuicUdpSessionSocket::new(connected)?); + let local_addr = socket.local_addr()?; + let remote_addr = socket.peer_addr(); + let runtime = default_runtime().ok_or(TunnelError::InternalError( + "no async runtime found".to_owned(), + ))?; + let mut endpoint = + Endpoint::new_with_abstract_socket(endpoint_config(), None, socket, runtime)?; + endpoint.set_default_client_config(client_config()); + let connecting = endpoint + .connect(remote_addr, "localhost") + .map_err(anyhow::Error::new) + .with_context(|| format!("failed to start connection to {remote_addr}"))?; + let connection = connecting + .await + .with_context(|| format!("failed to connect to {remote_addr}"))?; + let (write, read) = connection + .open_bi() + .await + .with_context(|| "open_bi failed")?; + let resolved_remote_addr = connection.remote_address(); + let connection = Arc::new(ConnWrapper { + conn: connection, + _endpoint: endpoint, + }); + let info = TunnelInfo { + tunnel_type: "quic".to_owned(), + local_addr: Some(super::build_url_from_socket_addr(&local_addr.to_string(), "quic").into()), + remote_addr: Some(remote_url.into()), + resolved_remote_addr: Some( + super::build_url_from_socket_addr(&resolved_remote_addr.to_string(), "quic").into(), + ), + }; + Ok(Box::new(TunnelWrapper::new( + FramedReader::new_with_associate_data(read, 4500, Some(Box::new(connection.clone()))), + FramedWriter::new_with_associate_data(write, Some(Box::new(connection))), + Some(info), + ))) } -impl QuicTunnelListener { - pub fn new(addr: url::Url, global_ctx: ArcGlobalCtx) -> Self { - QuicTunnelListener { - addr, - global_ctx, - endpoint: None, - } - } +struct PendingQuicSessionTunnel { + connecting: Connecting, + endpoint: Endpoint, + local_url: url::Url, + remote_addr: SocketAddr, + _handshake_permit: OwnedSemaphorePermit, +} - async fn do_accept(&self) -> Result, super::TunnelError> { - // accept a single connection - let conn = self - .endpoint - .as_ref() - .unwrap() - .accept() +async fn finish_quic_session_tunnel( + pending: PendingQuicSessionTunnel, +) -> Result, TunnelError> { + let PendingQuicSessionTunnel { + connecting, + endpoint, + local_url, + remote_addr, + _handshake_permit, + } = pending; + let connection = tokio::time::timeout(QUIC_ACCEPT_COMPLETION_TIMEOUT, connecting) + .await + .map_err(TunnelError::Timeout)? + .with_context(|| "accept connection failed")?; + let (write, read) = + tokio::time::timeout(QUIC_ACCEPT_COMPLETION_TIMEOUT, connection.accept_bi()) .await - .ok_or_else(|| anyhow::anyhow!("accept failed, no incoming"))?; - let conn = conn.await.with_context(|| "accept connection failed")?; - let remote_addr = conn.remote_address(); - let (w, r) = conn.accept_bi().await.with_context(|| "accept_bi failed")?; - - let arc_conn = Arc::new(ConnWrapper { conn }); - - let info = TunnelInfo { - tunnel_type: "quic".to_owned(), - local_addr: Some(self.local_url().into()), - remote_addr: Some( - super::build_url_from_socket_addr(&remote_addr.to_string(), "quic").into(), - ), - resolved_remote_addr: Some( - super::build_url_from_socket_addr(&remote_addr.to_string(), "quic").into(), - ), - }; - - Ok(Box::new(TunnelWrapper::new( - FramedReader::new_with_associate_data(r, 2000, Some(Box::new(arc_conn.clone()))), - FramedWriter::new_with_associate_data(w, Some(Box::new(arc_conn))), - Some(info), - ))) - } + .map_err(TunnelError::Timeout)? + .with_context(|| "accept_bi failed")?; + let connection = Arc::new(ConnWrapper { + conn: connection, + _endpoint: endpoint, + }); + let remote_url = super::build_url_from_socket_addr(&remote_addr.to_string(), "quic"); + let info = TunnelInfo { + tunnel_type: "quic".to_owned(), + local_addr: Some(local_url.into()), + remote_addr: Some(remote_url.clone().into()), + resolved_remote_addr: Some(remote_url.into()), + }; + Ok(Box::new(TunnelWrapper::new( + FramedReader::new_with_associate_data(read, 2000, Some(Box::new(connection.clone()))), + FramedWriter::new_with_associate_data(write, Some(Box::new(connection))), + Some(info), + ))) } -impl Drop for QuicTunnelListener { - fn drop(&mut self) { - let Some(endpoint) = &self.endpoint else { - return; - }; - let Ok(local_addr) = endpoint.local_addr() else { - return; - }; - QuicEndpointManager::load(&self.global_ctx).remove_endpoint_by_local_addr(local_addr); - } -} - -#[async_trait::async_trait] -impl TunnelListener for QuicTunnelListener { - async fn listen(&mut self) -> Result<(), TunnelError> { - let addr = SocketAddr::from_url(self.addr.clone(), IpVersion::Both).await?; - let endpoint = QuicEndpointManager::server(&self.global_ctx, addr)?; - self.addr - .set_port(Some(endpoint.local_addr()?.port())) - .unwrap(); - self.endpoint = Some(endpoint); - - Ok(()) - } - - async fn accept(&mut self) -> Result, super::TunnelError> { - loop { - match self.do_accept().await { - Ok(ret) => return Ok(ret), - Err(e) => { - tracing::warn!(?e, "accept fail"); - tokio::time::sleep(Duration::from_millis(1)).await; +async fn run_quic_accepted_session( + endpoint: Endpoint, + local_url: url::Url, + handshakes: Arc, + completed: Sender, TunnelError>>, +) { + let mut complete_tasks = JoinSet::new(); + let mut pending_incoming: Option = None; + loop { + tokio::select! { + Some(result) = complete_tasks.join_next(), if !complete_tasks.is_empty() => { + let result = match result { + Ok(result) => result, + Err(error) => Err(TunnelError::InternalError( + format!("quic accept task failed: {error}"), + )), + }; + if completed.send(result).await.is_err() { + break; + } + } + incoming = endpoint.accept(), if pending_incoming.is_none() => { + match incoming { + Some(incoming) => pending_incoming = Some(incoming), + None => break, + } + } + permit = handshakes.clone().acquire_owned(), if pending_incoming.is_some() => { + let Ok(handshake_permit) = permit else { + break; + }; + let incoming = pending_incoming.take().unwrap(); + let remote_addr = incoming.remote_address(); + match incoming.accept() { + Ok(connecting) => { + complete_tasks.spawn(finish_quic_session_tunnel( + PendingQuicSessionTunnel { + connecting, + endpoint: endpoint.clone(), + local_url: local_url.clone(), + remote_addr, + _handshake_permit: handshake_permit, + }, + )); + } + Err(error) => { + drop(handshake_permit); + if completed + .send(Err(anyhow::Error::new(error) + .context("quic accept connection failed") + .into())) + .await + .is_err() + { + break; + } + } } } } } +} - fn local_url(&self) -> url::Url { - self.addr.clone() +pub(crate) struct QuicAcceptedSession { + completed: Receiver, TunnelError>>, + _accept_task: AbortOnDropHandle<()>, +} + +impl QuicAcceptedSession { + pub(crate) fn new( + session: UdpSession, + local_url: url::Url, + admission: ServerProtocolAdmission, + ) -> Result { + let (active_session, handshake_slots) = admission.into_parts(); + Self::new_with_admission_parts(session, local_url, active_session, handshake_slots) } -} -pub struct QuicTunnelConnector { - addr: url::Url, - global_ctx: ArcGlobalCtx, - ip_version: IpVersion, - resolved_addr: Option, -} + fn new_with_admission_parts( + session: UdpSession, + local_url: url::Url, + active_session: OwnedSemaphorePermit, + handshakes: Arc, + ) -> Result { + let socket = Arc::new(QuicUdpSessionSocket::from_accepted( + session, + active_session, + )?); + let runtime = default_runtime().ok_or(TunnelError::InternalError( + "no async runtime found".to_owned(), + ))?; + let endpoint = Endpoint::new_with_abstract_socket( + endpoint_config(), + Some(server_config()), + socket, + runtime, + )?; + let (completed_tx, completed) = channel(100); + let accept_task = AbortOnDropHandle::new(tokio::spawn(run_quic_accepted_session( + endpoint, + local_url, + handshakes, + completed_tx, + ))); + Ok(Self { + completed, + _accept_task: accept_task, + }) + } -impl QuicTunnelConnector { - pub fn new(addr: url::Url, global_ctx: ArcGlobalCtx) -> Self { - QuicTunnelConnector { - addr, - global_ctx, - ip_version: IpVersion::Both, - resolved_addr: None, + pub(crate) async fn accept(&mut self) -> Result, TunnelError> { + while let Some(result) = self.completed.recv().await { + match result { + Ok(tunnel) => return Ok(tunnel), + Err(error) => { + tracing::warn!(?error, "QUIC session connection failed"); + } + } } + Err(TunnelError::Shutdown) } } #[async_trait::async_trait] -impl TunnelConnector for QuicTunnelConnector { - async fn connect(&mut self) -> Result, TunnelError> { - let addr = match self.resolved_addr { - Some(addr) => addr, - None => SocketAddr::from_url(self.addr.clone(), self.ip_version).await?, - }; - let (endpoint, connection) = QuicEndpointManager::connect(&self.global_ctx, addr).await?; - - let local_addr = endpoint.local_addr()?; - - let (w, r) = connection - .open_bi() - .await - .with_context(|| "open_bi failed")?; - - let info = TunnelInfo { - tunnel_type: "quic".to_owned(), - local_addr: Some( - super::build_url_from_socket_addr(&local_addr.to_string(), "quic").into(), - ), - remote_addr: Some(self.addr.clone().into()), - resolved_remote_addr: Some( - super::build_url_from_socket_addr(&connection.remote_address().to_string(), "quic") - .into(), - ), - }; - - let arc_conn = Arc::new(ConnWrapper { conn: connection }); - Ok(Box::new(TunnelWrapper::new( - FramedReader::new_with_associate_data(r, 4500, Some(Box::new(arc_conn.clone()))), - FramedWriter::new_with_associate_data(w, Some(Box::new(arc_conn))), - Some(info), - ))) - } - - fn remote_url(&self) -> url::Url { - self.addr.clone() - } - - fn set_ip_version(&mut self, ip_version: IpVersion) { - self.ip_version = ip_version; - } - - fn set_resolved_addr(&mut self, addr: SocketAddr) { - self.resolved_addr = Some(addr); +impl ServerTunnelAcceptor for QuicAcceptedSession { + async fn accept(&mut self) -> anyhow::Result> { + Ok(QuicAcceptedSession::accept(self).await?) } } #[cfg(test)] mod tests { - use crate::common::global_ctx::tests::get_mock_global_ctx_with_network; - use crate::tunnel::{ - TunnelConnector, - common::tests::{_tunnel_bench, _tunnel_pingpong}, + use std::{net::SocketAddr, sync::Arc, time::Duration}; + + use easytier_core::{ + connectivity::{ + protocol::ServerProtocolAdmissionController, + transport::{UdpSessionMode, connect_udp}, + }, + packet::ZCPacket, + socket::SocketListener, + socket::udp::{ + UdpBindOptions, UdpSessionAcceptKind, UdpSessionListenRequest, UdpSessionProtocol, + VirtualUdpSocket, + }, + }; + use futures::{SinkExt, StreamExt}; + + use crate::{ + common::netns::NetNS, host_runtime::native_host_runtime, + socket::udp::new_runtime_udp_session_listener, tunnel::common::tests::_tunnel_echo_server, }; - use std::sync::LazyLock; - use tokio::runtime::{Builder, Runtime}; use super::*; - // Shared runtime for all tests to avoid endpoint invalidation across runtimes - static RUNTIME: LazyLock = - LazyLock::new(|| Builder::new_multi_thread().enable_all().build().unwrap()); - - fn global_ctx() -> ArcGlobalCtx { - let identity = crate::common::config::NetworkIdentity::default(); - get_mock_global_ctx_with_network(Some(identity)) - } - - fn stopped_client_endpoint() -> (Endpoint, SocketAddr) { - let rt = Builder::new_current_thread().enable_all().build().unwrap(); - let endpoint = rt.block_on(async { - QuicEndpointManager::try_create((Ipv4Addr::UNSPECIFIED, 0).into(), false, None).unwrap() - }); - let local_addr = endpoint.local_addr().unwrap(); - drop(rt); - assert!(matches!( - endpoint.connect("127.0.0.1:1".parse().unwrap(), "localhost"), - Err(ConnectError::EndpointStopping) - )); - (endpoint, local_addr) - } - - #[test] - fn quic_pingpong() { - RUNTIME.block_on(quic_pingpong_impl()) - } - async fn quic_pingpong_impl() { - let listener = QuicTunnelListener::new("quic://[::]:21011".parse().unwrap(), global_ctx()); - let connector = - QuicTunnelConnector::new("quic://127.0.0.1:21011".parse().unwrap(), global_ctx()); - _tunnel_pingpong(listener, connector).await - } - - #[test] - fn quic_bench() { - RUNTIME.block_on(quic_bench_impl()) - } - async fn quic_bench_impl() { - let listener = QuicTunnelListener::new("quic://[::]:21012".parse().unwrap(), global_ctx()); - let connector = - QuicTunnelConnector::new("quic://127.0.0.1:21012".parse().unwrap(), global_ctx()); - _tunnel_bench(listener, connector).await - } - - #[test] - fn ipv6_pingpong() { - RUNTIME.block_on(ipv6_pingpong_impl()) - } - async fn ipv6_pingpong_impl() { - let listener = QuicTunnelListener::new("quic://[::1]:31015".parse().unwrap(), global_ctx()); - let connector = - QuicTunnelConnector::new("quic://[::1]:31015".parse().unwrap(), global_ctx()); - _tunnel_pingpong(listener, connector).await - } - - #[test] - fn ipv6_domain_pingpong() { - RUNTIME.block_on(ipv6_domain_pingpong_impl()) - } - async fn ipv6_domain_pingpong_impl() { - let listener = QuicTunnelListener::new("quic://[::1]:31016".parse().unwrap(), global_ctx()); - let mut connector = QuicTunnelConnector::new( - "quic://test.easytier.top:31016".parse().unwrap(), - global_ctx(), - ); - connector.set_ip_version(IpVersion::V6); - _tunnel_pingpong(listener, connector).await; - - let listener = - QuicTunnelListener::new("quic://127.0.0.1:31016".parse().unwrap(), global_ctx()); - let mut connector = QuicTunnelConnector::new( - "quic://test.easytier.top:31016".parse().unwrap(), - global_ctx(), - ); - connector.set_ip_version(IpVersion::V4); - _tunnel_pingpong(listener, connector).await; - } - - #[test] - fn alloc_port() { - RUNTIME.block_on(alloc_port_impl()) - } - async fn alloc_port_impl() { - // v4 - let mut listener = - QuicTunnelListener::new("quic://0.0.0.0:0".parse().unwrap(), global_ctx()); - listener.listen().await.unwrap(); - let port = listener.local_url().port().unwrap(); - assert!(port > 0); - - // v6 - let mut listener = QuicTunnelListener::new("quic://[::]:0".parse().unwrap(), global_ctx()); - listener.listen().await.unwrap(); - let port = listener.local_url().port().unwrap(); - assert!(port > 0); - } - - #[test] - fn listener_drop_removes_persistent_endpoint() { - RUNTIME.block_on(listener_drop_removes_persistent_endpoint_impl()) - } - async fn listener_drop_removes_persistent_endpoint_impl() { - let global_ctx = global_ctx(); - let endpoint_addr = { - let mut listener = - QuicTunnelListener::new("quic://127.0.0.1:0".parse().unwrap(), global_ctx.clone()); - listener.listen().await.unwrap(); - let endpoint_addr = listener.endpoint.as_ref().unwrap().local_addr().unwrap(); - assert!(QuicEndpointManager::load(&global_ctx).contains_local_addr(endpoint_addr)); - endpoint_addr - }; - - assert!(!QuicEndpointManager::load(&global_ctx).contains_local_addr(endpoint_addr)); - } - - #[test] - fn connect_removes_stopped_endpoints_and_retries() { - let (stopped_endpoint_a, stopped_addr_a) = stopped_client_endpoint(); - let (stopped_endpoint_b, stopped_addr_b) = stopped_client_endpoint(); - - RUNTIME.block_on(async move { - let mgr = QuicEndpointManager::new(2); - mgr.both.push(stopped_endpoint_a); - mgr.both.push(stopped_endpoint_b); - assert!(mgr.contains_local_addr(stopped_addr_a)); - assert!(mgr.contains_local_addr(stopped_addr_b)); - - let err = mgr - .connect_with_ip_version("127.0.0.1:0".parse().unwrap(), IpVersion::V4, None) - .await - .unwrap_err(); - let err = format!("{:?}", err); - assert!( - err.contains("invalid remote address"), - "unexpected error: {}", - err + #[tokio::test(flavor = "multi_thread")] + async fn accepted_udp_session_supports_multiple_quic_connections() { + tokio::time::timeout(Duration::from_secs(5), async { + let bind_addr: SocketAddr = "127.0.0.1:0".parse().unwrap(); + let mut listener = new_runtime_udp_session_listener( + format!("quic://{bind_addr}").parse().unwrap(), + UdpSessionListenRequest::new( + UdpBindOptions::port_bound_listener(bind_addr).with_only_v6(false), + ), + UdpSessionAcceptKind::Classified(UdpSessionProtocol::Quic), + NetNS::new(None), ); - assert!(!mgr.contains_local_addr(stopped_addr_a)); - assert!(!mgr.contains_local_addr(stopped_addr_b)); - }); - } + listener.listen().await.unwrap(); + let remote_addr = listener.bound_socket().unwrap().local_addr().unwrap(); + let local_url = listener.local_url(); - #[test] - fn invalid_peer_addr() { - RUNTIME.block_on(invalid_peer_addr_impl()) - } - async fn invalid_peer_addr_impl() { - let mut connector = - QuicTunnelConnector::new("quic://127.0.0.1:0".parse().unwrap(), global_ctx()); - let err = format!("{:?}", connector.connect().await.unwrap_err()); - assert!( - err.contains("invalid remote address"), - "unexpected error: {}", - err - ); + let connected = connect_udp( + native_host_runtime(), + remote_addr, + Vec::new(), + UdpBindOptions::direct_connect(), + UdpSessionMode::Classified(UdpSessionProtocol::Quic), + ) + .await + .unwrap(); + let socket = Arc::new(QuicUdpSessionSocket::new(connected).unwrap()); + let runtime = default_runtime().unwrap(); + let mut endpoint = + Endpoint::new_with_abstract_socket(endpoint_config(), None, socket, runtime) + .unwrap(); + endpoint.set_default_client_config(client_config()); + + let server_task = tokio::spawn(async move { + let session = listener.accept().await.unwrap(); + let admission = ServerProtocolAdmissionController::new(1, 2) + .try_admit() + .unwrap(); + let mut accepted = QuicAcceptedSession::new(session, local_url, admission).unwrap(); + let first = accepted.accept().await.unwrap(); + let second = accepted.accept().await.unwrap(); + assert!( + tokio::time::timeout(Duration::from_millis(50), listener.accept()) + .await + .is_err(), + "both QUIC connections must use the first accepted UDP session" + ); + (first, second) + }); + + let first_connection = endpoint + .connect(remote_addr, "localhost") + .unwrap() + .await + .unwrap(); + let (first_write, first_read) = first_connection.open_bi().await.unwrap(); + let mut first_send = FramedWriter::new(first_write); + first_send + .send(ZCPacket::new_with_payload(b"first QUIC connection")) + .await + .unwrap(); + let second_connection = endpoint + .connect(remote_addr, "localhost") + .unwrap() + .await + .unwrap(); + let (second_write, second_read) = second_connection.open_bi().await.unwrap(); + let mut second_send = FramedWriter::new(second_write); + second_send + .send(ZCPacket::new_with_payload(b"second QUIC connection ready")) + .await + .unwrap(); + let (first_server, second_server) = server_task.await.unwrap(); + + drop(first_send); + drop(first_read); + drop(first_server); + first_connection.close(0u32.into(), b"first connection done"); + + let echo_task = tokio::spawn(_tunnel_echo_server(second_server, false)); + let mut recv = FramedReader::new(second_read, 4500); + let ready = recv.next().await.unwrap().unwrap(); + assert_eq!(ready.payload(), b"second QUIC connection ready".as_slice()); + second_send + .send(ZCPacket::new_with_payload( + b"second QUIC connection after first closed", + )) + .await + .unwrap(); + let packet = recv.next().await.unwrap().unwrap(); + assert_eq!( + packet.payload(), + b"second QUIC connection after first closed".as_slice() + ); + let _ = second_send.close().await; + echo_task.await.unwrap(); + second_connection.close(0u32.into(), b"second connection done"); + endpoint.close(0u32.into(), b"test done"); + }) + .await + .unwrap(); } } diff --git a/easytier/src/tunnel/quic/session_socket.rs b/easytier/src/tunnel/quic/session_socket.rs new file mode 100644 index 00000000..06b6fdfb --- /dev/null +++ b/easytier/src/tunnel/quic/session_socket.rs @@ -0,0 +1,279 @@ +use std::{ + fmt, + future::Future, + io::{self, IoSliceMut}, + pin::Pin, + sync::{Arc, Mutex}, + task::{Context, Poll}, +}; + +use easytier_core::{ + connectivity::transport::ConnectedUdpSession, + socket::udp::{UdpSession, UdpSessionSocket}, +}; +use quinn::{ + AsyncUdpSocket, UdpPoller, + udp::{RecvMeta, Transmit}, +}; +use tokio::sync::mpsc::{self, Receiver, Sender, error::TrySendError}; +use tokio_util::task::AbortOnDropHandle; + +const DATAGRAM_QUEUE_CAPACITY: usize = 1024; + +type SendBatch = Vec>; +type WritableFuture = Pin> + Send>>; + +struct ReceivedDatagram { + payload: Vec, + dst_ip: Option, +} + +pub(crate) struct QuicUdpSessionSocket { + _session: Arc, + local_addr: std::net::SocketAddr, + peer_addr: std::net::SocketAddr, + incoming: Mutex>>, + outgoing: Sender, + _recv_task: AbortOnDropHandle<()>, + _send_task: AbortOnDropHandle<()>, + _session_guard: Box, +} + +impl fmt::Debug for QuicUdpSessionSocket { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter + .debug_struct("QuicUdpSessionSocket") + .field("local_addr", &self.local_addr) + .field("peer_addr", &self.peer_addr) + .finish_non_exhaustive() + } +} + +impl QuicUdpSessionSocket { + pub(crate) fn new(connected: ConnectedUdpSession) -> io::Result { + let (session, session_guard) = connected.into_parts(); + Self::from_session(Arc::new(session), session_guard) + } + + pub(crate) fn from_accepted(session: UdpSession, session_guard: T) -> io::Result + where + T: Send + Sync + 'static, + { + Self::from_session(Arc::new(session), Box::new(session_guard)) + } + + fn from_session( + session: Arc, + session_guard: Box, + ) -> io::Result { + let local_addr = session.local_addr()?; + let peer_addr = session.peer_addr()?; + let (incoming_tx, incoming) = mpsc::channel(DATAGRAM_QUEUE_CAPACITY); + let (outgoing, mut outgoing_rx) = mpsc::channel::(DATAGRAM_QUEUE_CAPACITY); + + let recv_session = session.clone(); + let recv_errors = incoming_tx.clone(); + let recv_task = AbortOnDropHandle::new(tokio::spawn(async move { + let mut buffer = vec![0; 64 * 1024]; + loop { + match recv_session.recv_with_meta(&mut buffer).await { + Ok((length, meta)) => { + let datagram = ReceivedDatagram { + payload: buffer[..length].to_vec(), + dst_ip: meta.dst_ip, + }; + if incoming_tx.send(Ok(datagram)).await.is_err() { + break; + } + } + Err(error) => { + let _ = incoming_tx.send(Err(error)).await; + break; + } + } + } + })); + + let send_session = session.clone(); + let send_task = AbortOnDropHandle::new(tokio::spawn(async move { + while let Some(batch) = outgoing_rx.recv().await { + for datagram in batch { + match send_session.send(&datagram).await { + Ok(length) if length == datagram.len() => {} + Ok(_) => { + let _ = recv_errors + .send(Err(io::Error::new( + io::ErrorKind::WriteZero, + "QUIC UDP session partially sent a datagram", + ))) + .await; + return; + } + Err(error) => { + let _ = recv_errors.send(Err(error)).await; + return; + } + } + } + } + })); + + Ok(Self { + _session: session, + local_addr, + peer_addr, + incoming: Mutex::new(incoming), + outgoing, + _recv_task: recv_task, + _send_task: send_task, + _session_guard: session_guard, + }) + } + + pub(crate) fn peer_addr(&self) -> std::net::SocketAddr { + self.peer_addr + } + + fn send_batch(&self, transmit: &Transmit<'_>) -> io::Result<()> { + if transmit.destination != self.peer_addr { + return Err(io::Error::new( + io::ErrorKind::AddrNotAvailable, + format!( + "QUIC UDP session is connected to {}, not {}", + self.peer_addr, transmit.destination + ), + )); + } + + let batch = match transmit.segment_size { + Some(0) => { + return Err(io::Error::new( + io::ErrorKind::InvalidInput, + "QUIC segment size cannot be zero", + )); + } + Some(segment_size) => transmit + .contents + .chunks(segment_size) + .map(<[u8]>::to_vec) + .collect(), + None => vec![transmit.contents.to_vec()], + }; + self.outgoing.try_send(batch).map_err(map_try_send_error) + } +} + +fn map_try_send_error(error: TrySendError) -> io::Error { + match error { + TrySendError::Full(_) => io::Error::new(io::ErrorKind::WouldBlock, "QUIC send queue full"), + TrySendError::Closed(_) => { + io::Error::new(io::ErrorKind::BrokenPipe, "QUIC UDP session closed") + } + } +} + +struct QuicUdpSessionPoller { + outgoing: Sender, + writable: Mutex>, +} + +impl fmt::Debug for QuicUdpSessionPoller { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter + .debug_struct("QuicUdpSessionPoller") + .finish_non_exhaustive() + } +} + +impl UdpPoller for QuicUdpSessionPoller { + fn poll_writable(self: Pin<&mut Self>, context: &mut Context<'_>) -> Poll> { + let mut writable = self.writable.lock().unwrap(); + if writable.is_none() { + let outgoing = self.outgoing.clone(); + *writable = Some(Box::pin(async move { + let permit = outgoing.reserve_owned().await.map_err(|_| { + io::Error::new(io::ErrorKind::BrokenPipe, "QUIC UDP session closed") + })?; + drop(permit); + Ok(()) + })); + } + + match writable.as_mut().unwrap().as_mut().poll(context) { + Poll::Ready(result) => { + *writable = None; + Poll::Ready(result) + } + Poll::Pending => Poll::Pending, + } + } +} + +impl AsyncUdpSocket for QuicUdpSessionSocket { + fn create_io_poller(self: Arc) -> Pin> { + Box::pin(QuicUdpSessionPoller { + outgoing: self.outgoing.clone(), + writable: Mutex::new(None), + }) + } + + fn try_send(&self, transmit: &Transmit<'_>) -> io::Result<()> { + self.send_batch(transmit) + } + + fn poll_recv( + &self, + context: &mut Context<'_>, + buffers: &mut [IoSliceMut<'_>], + meta: &mut [RecvMeta], + ) -> Poll> { + if buffers.is_empty() || meta.is_empty() { + return Poll::Ready(Err(io::Error::new( + io::ErrorKind::InvalidInput, + "QUIC UDP recv buffers are empty", + ))); + } + + let mut incoming = self.incoming.lock().unwrap(); + loop { + match Pin::new(&mut *incoming).poll_recv(context) { + Poll::Ready(Some(Ok(datagram))) => { + if buffers[0].len() < datagram.payload.len() { + tracing::debug!( + payload_len = datagram.payload.len(), + recv_buf_len = buffers[0].len(), + peer_addr = ?self.peer_addr, + "drop oversized QUIC UDP session datagram" + ); + continue; + } + buffers[0][..datagram.payload.len()].copy_from_slice(&datagram.payload); + meta[0] = RecvMeta { + addr: self.peer_addr, + len: datagram.payload.len(), + stride: datagram.payload.len(), + ecn: None, + dst_ip: datagram.dst_ip, + }; + return Poll::Ready(Ok(1)); + } + Poll::Ready(Some(Err(error))) => return Poll::Ready(Err(error)), + Poll::Ready(None) => { + return Poll::Ready(Err(io::Error::new( + io::ErrorKind::UnexpectedEof, + "QUIC UDP session closed", + ))); + } + Poll::Pending => return Poll::Pending, + } + } + } + + fn local_addr(&self) -> io::Result { + Ok(self.local_addr) + } + + fn may_fragment(&self) -> bool { + false + } +} diff --git a/easytier/src/tunnel/ring.rs b/easytier/src/tunnel/ring.rs deleted file mode 100644 index 101b67c6..00000000 --- a/easytier/src/tunnel/ring.rs +++ /dev/null @@ -1,391 +0,0 @@ -use async_ringbuf::{AsyncHeapCons, AsyncHeapProd, AsyncHeapRb, traits::*}; -use crossbeam::atomic::AtomicCell; -use std::{ - collections::HashMap, - fmt::Debug, - sync::Arc, - task::{Poll, ready}, -}; - -use async_trait::async_trait; -use futures::{Sink, SinkExt, Stream, StreamExt}; -use once_cell::sync::Lazy; - -use tokio::sync::mpsc::{UnboundedReceiver, UnboundedSender, unbounded_channel}; - -use uuid::Uuid; - -use crate::tunnel::{FromUrl, IpVersion, SinkError, SinkItem}; - -use super::{ - StreamItem, Tunnel, TunnelConnector, TunnelError, TunnelInfo, TunnelListener, - build_url_from_socket_addr, common::TunnelWrapper, -}; - -pub static RING_TUNNEL_CAP: usize = 128; -static RING_TUNNEL_RESERVED_CAP: usize = 4; - -type RingLock = parking_lot::Mutex<()>; - -type RingItem = SinkItem; - -pub struct RingTunnel { - id: Uuid, - - ring_cons_impl: AtomicCell>>, - ring_prod_impl: AtomicCell>>, -} - -impl RingTunnel { - fn id(&self) -> &Uuid { - &self.id - } - - pub fn new(cap: usize) -> Self { - let id = Uuid::new_v4(); - let ring_impl = AsyncHeapRb::new(std::cmp::max(RING_TUNNEL_RESERVED_CAP * 2, cap)); - let (ring_prod_impl, ring_cons_impl) = ring_impl.split(); - Self { - id, - ring_cons_impl: AtomicCell::new(Some(ring_cons_impl)), - ring_prod_impl: AtomicCell::new(Some(ring_prod_impl)), - } - } - - pub fn new_with_id(id: Uuid, cap: usize) -> Self { - let mut ret = Self::new(cap); - ret.id = id; - ret - } -} - -impl Debug for RingTunnel { - fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { - f.debug_struct("RingTunnel").field("id", &self.id).finish() - } -} - -pub struct RingStream { - id: Uuid, - ring_cons_impl: AsyncHeapCons, -} - -impl RingStream { - pub fn new(tunnel: Arc) -> Self { - Self { - id: tunnel.id, - ring_cons_impl: tunnel.ring_cons_impl.take().unwrap(), - } - } -} - -impl Stream for RingStream { - type Item = StreamItem; - - fn poll_next( - self: std::pin::Pin<&mut Self>, - cx: &mut std::task::Context<'_>, - ) -> Poll> { - let ret = ready!(self.get_mut().ring_cons_impl.poll_next_unpin(cx)); - match ret { - Some(item) => Poll::Ready(Some(Ok(item))), - None => Poll::Ready(None), - } - } -} - -impl Debug for RingStream { - fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { - f.debug_struct("RingStream") - .field("id", &self.id) - .field("len", &self.ring_cons_impl.base().occupied_len()) - .field("cap", &self.ring_cons_impl.base().capacity()) - .finish() - } -} - -pub struct RingSink { - id: Uuid, - ring_prod_impl: AsyncHeapProd, -} - -impl RingSink { - pub fn new(tunnel: Arc) -> Self { - Self { - id: tunnel.id, - ring_prod_impl: tunnel.ring_prod_impl.take().unwrap(), - } - } - - pub fn try_send(&mut self, item: RingItem) -> Result<(), RingItem> { - let base = self.ring_prod_impl.base(); - if base.occupied_len() >= base.capacity().get() - RING_TUNNEL_RESERVED_CAP { - return Err(item); - } - self.ring_prod_impl.try_push(item) - } - - pub fn force_send(&mut self, item: RingItem) -> Result<(), RingItem> { - self.ring_prod_impl.try_push(item) - } -} - -impl Sink for RingSink { - type Error = SinkError; - - fn poll_ready( - self: std::pin::Pin<&mut Self>, - cx: &mut std::task::Context<'_>, - ) -> std::task::Poll> { - let ret = ready!(self.get_mut().ring_prod_impl.poll_ready_unpin(cx)); - Poll::Ready(ret.map_err(|_| TunnelError::Shutdown)) - } - - fn start_send(self: std::pin::Pin<&mut Self>, item: SinkItem) -> Result<(), Self::Error> { - self.get_mut() - .ring_prod_impl - .start_send_unpin(item) - .map_err(|_| TunnelError::Shutdown) - } - - fn poll_flush( - self: std::pin::Pin<&mut Self>, - cx: &mut std::task::Context<'_>, - ) -> std::task::Poll> { - let ret = ready!(self.get_mut().ring_prod_impl.poll_flush_unpin(cx)); - Poll::Ready(ret.map_err(|_| TunnelError::Shutdown)) - } - - fn poll_close( - self: std::pin::Pin<&mut Self>, - cx: &mut std::task::Context<'_>, - ) -> std::task::Poll> { - let ret = ready!(self.get_mut().ring_prod_impl.poll_close_unpin(cx)); - Poll::Ready(ret.map_err(|_| TunnelError::Shutdown)) - } -} - -impl Debug for RingSink { - fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { - f.debug_struct("RingSink") - .field("id", &self.id) - .field("len", &self.ring_prod_impl.base().occupied_len()) - .field("cap", &self.ring_prod_impl.base().capacity()) - .finish() - } -} - -struct Connection { - client: Arc, - server: Arc, -} - -type ConnectionMap = HashMap>>; - -static CONNECTION_MAP: Lazy>> = - Lazy::new(|| Arc::new(std::sync::Mutex::new(HashMap::new()))); - -#[derive(Debug)] -pub struct RingTunnelListener { - listener_addr: url::Url, - conn_sender: UnboundedSender>, - conn_receiver: UnboundedReceiver>, - - key_in_conn_map: Option, -} - -impl RingTunnelListener { - pub fn new(key: url::Url) -> Self { - let (conn_sender, conn_receiver) = unbounded_channel(); - RingTunnelListener { - listener_addr: key, - conn_sender, - conn_receiver, - key_in_conn_map: None, - } - } -} - -fn get_tunnel_for_client(conn: Arc) -> impl Tunnel { - TunnelWrapper::new( - RingStream::new(conn.client.clone()), - RingSink::new(conn.server.clone()), - Some(TunnelInfo { - tunnel_type: "ring".to_owned(), - local_addr: Some(build_url_from_socket_addr(&conn.client.id.into(), "ring").into()), - remote_addr: Some(build_url_from_socket_addr(&conn.server.id.into(), "ring").into()), - resolved_remote_addr: Some( - build_url_from_socket_addr(&conn.server.id.into(), "ring").into(), - ), - }), - ) -} - -fn get_tunnel_for_server(conn: Arc) -> impl Tunnel { - TunnelWrapper::new( - RingStream::new(conn.server.clone()), - RingSink::new(conn.client.clone()), - Some(TunnelInfo { - tunnel_type: "ring".to_owned(), - local_addr: Some(build_url_from_socket_addr(&conn.server.id.into(), "ring").into()), - remote_addr: Some(build_url_from_socket_addr(&conn.client.id.into(), "ring").into()), - resolved_remote_addr: Some( - build_url_from_socket_addr(&conn.client.id.into(), "ring").into(), - ), - }), - ) -} - -impl RingTunnelListener { - async fn get_addr(&self) -> Result { - Uuid::from_url(self.listener_addr.clone(), IpVersion::Both).await - } -} - -#[async_trait] -impl TunnelListener for RingTunnelListener { - async fn listen(&mut self) -> Result<(), TunnelError> { - tracing::info!("listen new conn of key: {}", self.listener_addr); - let addr = self.get_addr().await?; - CONNECTION_MAP - .lock() - .unwrap() - .insert(addr, self.conn_sender.clone()); - self.key_in_conn_map = Some(addr); - Ok(()) - } - - async fn accept(&mut self) -> Result, TunnelError> { - tracing::info!("waiting accept new conn of key: {}", self.listener_addr); - let my_addr = self.get_addr().await?; - if let Some(conn) = self.conn_receiver.recv().await { - if conn.server.id == my_addr { - tracing::info!("accept new conn of key: {}", self.listener_addr); - return Ok(Box::new(get_tunnel_for_server(conn))); - } else { - tracing::error!(?conn.server.id, ?my_addr, "got new conn with wrong id"); - return Err(TunnelError::InternalError( - "accept got wrong ring server id".to_owned(), - )); - } - } - - return Err(TunnelError::InternalError( - "conn receiver stopped".to_owned(), - )); - } - - fn local_url(&self) -> url::Url { - self.listener_addr.clone() - } -} - -impl Drop for RingTunnelListener { - fn drop(&mut self) { - if let Some(addr) = self.key_in_conn_map { - CONNECTION_MAP.lock().unwrap().remove(&addr); - } - } -} - -pub struct RingTunnelConnector { - remote_addr: url::Url, -} - -impl RingTunnelConnector { - pub fn new(remote_addr: url::Url) -> Self { - RingTunnelConnector { remote_addr } - } -} - -#[async_trait] -impl TunnelConnector for RingTunnelConnector { - async fn connect(&mut self) -> Result, super::TunnelError> { - let remote_addr = Uuid::from_url(self.remote_addr.clone(), IpVersion::Both).await?; - let entry = CONNECTION_MAP - .lock() - .unwrap() - .get(&remote_addr) - .unwrap() - .clone(); - tracing::info!("connecting"); - let conn = Arc::new(Connection { - client: Arc::new(RingTunnel::new(RING_TUNNEL_CAP)), - server: Arc::new(RingTunnel::new_with_id(remote_addr, RING_TUNNEL_CAP)), - }); - entry - .send(conn.clone()) - .map_err(|_| TunnelError::InternalError("send conn to listner failed".to_owned()))?; - Ok(Box::new(get_tunnel_for_client(conn))) - } - - fn remote_url(&self) -> url::Url { - self.remote_addr.clone() - } -} - -pub fn create_ring_tunnel_pair() -> (Box, Box) { - let conn = Arc::new(Connection { - client: Arc::new(RingTunnel::new(RING_TUNNEL_CAP)), - server: Arc::new(RingTunnel::new(RING_TUNNEL_CAP)), - }); - ( - Box::new(get_tunnel_for_server(conn.clone())), - Box::new(get_tunnel_for_client(conn)), - ) -} - -#[cfg(test)] -mod tests { - use futures::StreamExt; - use tokio::time::timeout; - - use crate::tunnel::common::tests::{_tunnel_bench, _tunnel_pingpong}; - - use super::*; - - #[tokio::test] - async fn ring_pingpong() { - let id: url::Url = format!("ring://{}", Uuid::new_v4()).parse().unwrap(); - let listener = RingTunnelListener::new(id.clone()); - let connector = RingTunnelConnector::new(id.clone()); - _tunnel_pingpong(listener, connector).await - } - - #[tokio::test] - async fn ring_bench() { - let id: url::Url = format!("ring://{}", Uuid::new_v4()).parse().unwrap(); - let listener = RingTunnelListener::new(id.clone()); - let connector = RingTunnelConnector::new(id); - _tunnel_bench(listener, connector).await - } - - #[tokio::test] - async fn ring_close() { - let (stunnel, ctunnel) = create_ring_tunnel_pair(); - drop(stunnel); - - let mut stream = ctunnel.split().0; - let ret = stream.next().await; - assert!(ret.as_ref().is_none(), "expect none, got {:?}", ret); - } - - #[tokio::test] - async fn abort_ring_stream() { - let (_stunnel, ctunnel) = create_ring_tunnel_pair(); - let mut stream = ctunnel.split().0; - let task = tokio::spawn(async move { - let _ = stream.next().await; - }); - tokio::time::sleep(tokio::time::Duration::from_secs(1)).await; - task.abort(); - let _ = tokio::join!(task); - } - - #[tokio::test] - async fn ring_stream_recv_timeout() { - let (_stunnel, ctunnel) = create_ring_tunnel_pair(); - let mut stream = ctunnel.split().0; - let _ = timeout(tokio::time::Duration::from_millis(10), stream.next()).await; - } -} diff --git a/easytier/src/tunnel/tcp.rs b/easytier/src/tunnel/tcp.rs deleted file mode 100644 index 8e678170..00000000 --- a/easytier/src/tunnel/tcp.rs +++ /dev/null @@ -1,378 +0,0 @@ -use std::net::SocketAddr; - -use super::{FromUrl, TunnelInfo}; -use crate::tunnel::common::{apply_socket_mark, bind}; -use async_trait::async_trait; -use futures::stream::FuturesUnordered; -use tokio::net::{TcpListener, TcpSocket, TcpStream}; - -use super::{ - IpVersion, Tunnel, TunnelError, TunnelListener, - common::{FramedReader, FramedWriter, TunnelWrapper, wait_for_connect_futures}, -}; - -const TCP_MTU_BYTES: usize = 2000; - -#[derive(Debug)] -pub struct TcpTunnelListener { - addr: url::Url, - listener: Option, - socket_mark: Option, -} - -impl TcpTunnelListener { - pub fn new(addr: url::Url) -> Self { - TcpTunnelListener { - addr, - listener: None, - socket_mark: None, - } - } - - pub fn set_socket_mark(&mut self, socket_mark: Option) { - self.socket_mark = socket_mark; - } - - async fn do_accept(&self) -> Result, std::io::Error> { - let listener = self.listener.as_ref().unwrap(); - let (stream, _) = listener.accept().await?; - - if let Err(e) = stream.set_nodelay(true) { - tracing::warn!(?e, "set_nodelay fail in accept"); - } - - let info = TunnelInfo { - tunnel_type: "tcp".to_owned(), - local_addr: Some(self.local_url().into()), - remote_addr: Some( - super::build_url_from_socket_addr(&stream.peer_addr()?.to_string(), "tcp").into(), - ), - resolved_remote_addr: Some( - super::build_url_from_socket_addr(&stream.peer_addr()?.to_string(), "tcp").into(), - ), - }; - - let (r, w) = stream.into_split(); - Ok(Box::new(TunnelWrapper::new( - FramedReader::new(r, TCP_MTU_BYTES), - FramedWriter::new(w), - Some(info), - ))) - } -} - -#[async_trait] -impl TunnelListener for TcpTunnelListener { - async fn listen(&mut self) -> Result<(), TunnelError> { - self.listener = None; - - let addr = SocketAddr::from_url(self.addr.clone(), IpVersion::Both).await?; - let listener = bind::() - .addr(addr) - .only_v6(true) - .maybe_socket_mark(self.socket_mark) - .call()?; - - self.addr - .set_port(Some(listener.local_addr()?.port())) - .unwrap(); - self.listener = Some(listener); - - Ok(()) - } - - async fn accept(&mut self) -> Result, super::TunnelError> { - loop { - match self.do_accept().await { - Ok(ret) => return Ok(ret), - Err(e) => { - use std::io::ErrorKind::*; - if matches!( - e.kind(), - NotConnected | ConnectionAborted | ConnectionRefused | ConnectionReset - ) { - tracing::warn!(?e, "accept fail with retryable error: {:?}", e); - continue; - } - tracing::warn!(?e, "accept fail"); - return Err(e.into()); - } - } - } - } - - fn local_url(&self) -> url::Url { - self.addr.clone() - } -} - -fn get_tunnel_with_tcp_stream( - stream: TcpStream, - remote_url: url::Url, -) -> Result, super::TunnelError> { - if let Err(e) = stream.set_nodelay(true) { - tracing::warn!(?e, "set_nodelay fail in get_tunnel_with_tcp_stream"); - } - - let info = TunnelInfo { - tunnel_type: "tcp".to_owned(), - local_addr: Some( - super::build_url_from_socket_addr(&stream.local_addr()?.to_string(), "tcp").into(), - ), - remote_addr: Some(remote_url.into()), - resolved_remote_addr: Some( - super::build_url_from_socket_addr(&stream.peer_addr()?.to_string(), "tcp").into(), - ), - }; - - let (r, w) = stream.into_split(); - Ok(Box::new(TunnelWrapper::new( - FramedReader::new(r, TCP_MTU_BYTES), - FramedWriter::new(w), - Some(info), - ))) -} - -#[derive(Debug)] -pub struct TcpTunnelConnector { - addr: url::Url, - - bind_addrs: Vec, - ip_version: IpVersion, - resolved_addr: Option, - socket_mark: Option, -} - -impl TcpTunnelConnector { - pub fn new(addr: url::Url) -> Self { - TcpTunnelConnector { - addr, - bind_addrs: vec![], - ip_version: IpVersion::Both, - resolved_addr: None, - socket_mark: None, - } - } - - async fn connect_with_default_bind( - &self, - addr: SocketAddr, - ) -> Result, super::TunnelError> { - tracing::info!(url = ?self.addr, ?addr, "connect tcp start, bind addrs: {:?}", self.bind_addrs); - let stream = if self.socket_mark.is_some() { - // SO_MARK requires applying the option on the socket before - // connect, so go through TcpSocket rather than TcpStream::connect. - let socket = if addr.is_ipv4() { - TcpSocket::new_v4()? - } else { - TcpSocket::new_v6()? - }; - apply_socket_mark(&socket2::SockRef::from(&socket), self.socket_mark)?; - socket.connect(addr).await? - } else { - TcpStream::connect(addr).await? - }; - tracing::info!(url = ?self.addr, ?addr, "connect tcp succ"); - get_tunnel_with_tcp_stream(stream, self.addr.clone()) - } - - async fn connect_with_custom_bind( - &self, - addr: SocketAddr, - ) -> Result, super::TunnelError> { - let futures = FuturesUnordered::new(); - - for bind_addr in self.bind_addrs.iter() { - tracing::info!(?bind_addr, ?addr, "bind addr"); - match bind::() - .addr(*bind_addr) - .only_v6(true) - .maybe_socket_mark(self.socket_mark) - .call() - { - Ok(socket) => futures.push(socket.connect(addr)), - Err(error) => { - tracing::error!(?bind_addr, ?addr, ?error, "bind addr fail"); - continue; - } - } - } - - let ret = wait_for_connect_futures(futures).await; - get_tunnel_with_tcp_stream(ret?, self.addr.clone()) - } -} - -#[async_trait] -impl super::TunnelConnector for TcpTunnelConnector { - async fn connect(&mut self) -> Result, TunnelError> { - let addr = match self.resolved_addr { - Some(addr) => addr, - None => SocketAddr::from_url(self.addr.clone(), self.ip_version).await?, - }; - if self.bind_addrs.is_empty() { - self.connect_with_default_bind(addr).await - } else { - self.connect_with_custom_bind(addr).await - } - } - - fn remote_url(&self) -> url::Url { - self.addr.clone() - } - - fn set_bind_addrs(&mut self, addrs: Vec) { - self.bind_addrs = addrs; - } - - fn set_ip_version(&mut self, ip_version: IpVersion) { - self.ip_version = ip_version; - } - - fn set_resolved_addr(&mut self, addr: SocketAddr) { - self.resolved_addr = Some(addr); - } - - fn set_socket_mark(&mut self, socket_mark: Option) { - self.socket_mark = socket_mark; - } -} - -#[cfg(test)] -mod tests { - use crate::tunnel::{ - TunnelConnector, - common::tests::{_tunnel_bench, _tunnel_pingpong}, - }; - - use super::*; - - #[tokio::test] - async fn tcp_pingpong() { - let listener = TcpTunnelListener::new("tcp://0.0.0.0:31011".parse().unwrap()); - let connector = TcpTunnelConnector::new("tcp://127.0.0.1:31011".parse().unwrap()); - _tunnel_pingpong(listener, connector).await - } - - #[tokio::test] - async fn tcp_bench() { - let listener = TcpTunnelListener::new("tcp://0.0.0.0:31012".parse().unwrap()); - let connector = TcpTunnelConnector::new("tcp://127.0.0.1:31012".parse().unwrap()); - _tunnel_bench(listener, connector).await - } - - #[tokio::test] - async fn tcp_bench_with_bind() { - let listener = TcpTunnelListener::new("tcp://127.0.0.1:11013".parse().unwrap()); - let mut connector = TcpTunnelConnector::new("tcp://127.0.0.1:11013".parse().unwrap()); - connector.set_bind_addrs(vec!["127.0.0.1:0".parse().unwrap()]); - _tunnel_pingpong(listener, connector).await - } - - #[tokio::test] - #[should_panic] - async fn tcp_bench_with_bind_fail() { - let listener = TcpTunnelListener::new("tcp://127.0.0.1:11014".parse().unwrap()); - let mut connector = TcpTunnelConnector::new("tcp://127.0.0.1:11014".parse().unwrap()); - connector.set_bind_addrs(vec!["10.0.0.1:0".parse().unwrap()]); - _tunnel_pingpong(listener, connector).await - } - - #[tokio::test] - async fn bind_same_port() { - let mut listener = TcpTunnelListener::new("tcp://[::]:31014".parse().unwrap()); - let mut listener2 = TcpTunnelListener::new("tcp://0.0.0.0:31014".parse().unwrap()); - listener.listen().await.unwrap(); - listener2.listen().await.unwrap(); - } - - #[tokio::test] - async fn ipv6_pingpong() { - let listener = TcpTunnelListener::new("tcp://[::1]:31015".parse().unwrap()); - let connector = TcpTunnelConnector::new("tcp://[::1]:31015".parse().unwrap()); - _tunnel_pingpong(listener, connector).await - } - - #[tokio::test] - async fn ipv6_domain_pingpong() { - let listener = TcpTunnelListener::new("tcp://[::1]:31015".parse().unwrap()); - let mut connector = - TcpTunnelConnector::new("tcp://test.easytier.top:31015".parse().unwrap()); - connector.set_ip_version(IpVersion::V6); - _tunnel_pingpong(listener, connector).await; - - let listener = TcpTunnelListener::new("tcp://127.0.0.1:31015".parse().unwrap()); - let mut connector = - TcpTunnelConnector::new("tcp://test.easytier.top:31015".parse().unwrap()); - connector.set_ip_version(IpVersion::V4); - _tunnel_pingpong(listener, connector).await; - } - - #[tokio::test] - async fn connector_keeps_source_addr_and_reports_resolved_addr() { - let mut listener = TcpTunnelListener::new("tcp://127.0.0.1:0".parse().unwrap()); - listener.listen().await.unwrap(); - - let port = listener.local_url().port().unwrap(); - let source_url: url::Url = format!("tcp://localhost:{port}").parse().unwrap(); - let mut connector = TcpTunnelConnector::new(source_url.clone()); - connector.set_ip_version(IpVersion::V4); - - let accept_task = tokio::spawn(async move { listener.accept().await.unwrap() }); - let tunnel = connector.connect().await.unwrap(); - let accepted_tunnel = accept_task.await.unwrap(); - - let info = tunnel.info().unwrap(); - assert_eq!(info.remote_addr.unwrap().url, source_url.to_string()); - - let resolved_remote_addr: url::Url = info.resolved_remote_addr.unwrap().into(); - assert_eq!(resolved_remote_addr.host_str(), Some("127.0.0.1")); - assert_eq!(resolved_remote_addr.port(), Some(port)); - - let accepted_info = accepted_tunnel.info().unwrap(); - assert_eq!( - accepted_info.remote_addr, - accepted_info.resolved_remote_addr, - ); - } - - #[tokio::test] - async fn connector_uses_pre_resolved_addr_without_resolving_url() { - let mut listener = TcpTunnelListener::new("tcp://127.0.0.1:0".parse().unwrap()); - listener.listen().await.unwrap(); - - let port = listener.local_url().port().unwrap(); - let source_url: url::Url = format!("tcp://unresolvable.invalid:{port}") - .parse() - .unwrap(); - let resolved_addr: SocketAddr = format!("127.0.0.1:{port}").parse().unwrap(); - let mut connector = TcpTunnelConnector::new(source_url.clone()); - connector.set_resolved_addr(resolved_addr); - - let accept_task = tokio::spawn(async move { listener.accept().await.unwrap() }); - let tunnel = connector.connect().await.unwrap(); - let _accepted_tunnel = accept_task.await.unwrap(); - - let info = tunnel.info().unwrap(); - assert_eq!(info.remote_addr.unwrap().url, source_url.to_string()); - - let resolved_remote_addr: url::Url = info.resolved_remote_addr.unwrap().into(); - assert_eq!(resolved_remote_addr.host_str(), Some("127.0.0.1")); - assert_eq!(resolved_remote_addr.port(), Some(port)); - } - - #[tokio::test] - async fn test_alloc_port() { - // v4 - let mut listener = TcpTunnelListener::new("tcp://0.0.0.0:0".parse().unwrap()); - listener.listen().await.unwrap(); - let port = listener.local_url().port().unwrap(); - assert!(port > 0); - - // v6 - let mut listener = TcpTunnelListener::new("tcp://[::]:0".parse().unwrap()); - listener.listen().await.unwrap(); - let port = listener.local_url().port().unwrap(); - assert!(port > 0); - } -} diff --git a/easytier/src/tunnel/udp.rs b/easytier/src/tunnel/udp.rs deleted file mode 100644 index d473afb8..00000000 --- a/easytier/src/tunnel/udp.rs +++ /dev/null @@ -1,1517 +0,0 @@ -use std::{ - fmt::Debug, - net::{Ipv4Addr, Ipv6Addr, SocketAddr, SocketAddrV4, SocketAddrV6}, - sync::{Arc, Weak}, - time::Duration, -}; - -use anyhow::Context; -use async_trait::async_trait; -use bytes::BytesMut; -use dashmap::DashMap; -use futures::{StreamExt, stream::FuturesUnordered}; -use rand::{Rng, SeedableRng}; -use zerocopy::{AsBytes, FromBytes}; - -use tokio::{ - net::UdpSocket, - sync::mpsc::{ - Receiver, Sender, UnboundedReceiver, UnboundedSender, channel, unbounded_channel, - }, - task::JoinSet, -}; -use tokio_util::task::AbortOnDropHandle; -use tracing::{Instrument, instrument}; - -use super::{ - FromUrl, IpVersion, Tunnel, TunnelConnCounter, TunnelError, TunnelInfo, TunnelListener, - TunnelUrl, - common::wait_for_connect_futures, - packet_def::{UDP_TUNNEL_HEADER_SIZE, UDPTunnelHeader, V4HolePunchPacket, V6HolePunchPacket}, - ring::{RingSink, RingStream}, -}; -use crate::tunnel::common::bind; -use crate::{ - common::{join_joinset_background, shrink_dashmap}, - tunnel::{ - build_url_from_socket_addr, - common::{TunnelWrapper, reserve_buf}, - packet_def::{UdpPacketType, ZCPacket, ZCPacketType}, - ring::RingTunnel, - udp_src, - }, -}; - -pub const UDP_DATA_MTU: usize = 2000; - -type UdpCloseEventSender = UnboundedSender<(SocketAddr, Option)>; -type UdpCloseEventReceiver = UnboundedReceiver<(SocketAddr, Option)>; - -#[derive(Clone, Copy, Debug, PartialEq, Eq)] -pub struct PreferredIpv6Source { - pub ip: Ipv6Addr, - pub ifindex: u32, -} - -fn new_udp_packet(f: F, udp_body: Option<&[u8]>) -> ZCPacket -where - F: FnOnce(&mut UDPTunnelHeader), -{ - let mut buf = BytesMut::new(); - buf.resize( - UDP_TUNNEL_HEADER_SIZE + udp_body.as_ref().map(|v| v.len()).unwrap_or(0), - 0, - ); - buf[UDP_TUNNEL_HEADER_SIZE..].copy_from_slice(udp_body.unwrap()); - - let mut ret = ZCPacket::new_from_buf(buf, ZCPacketType::UDP); - let header = ret.mut_udp_tunnel_header().unwrap(); - f(header); - ret -} - -fn new_syn_packet(conn_id: u32, magic: u64) -> ZCPacket { - new_udp_packet( - |header| { - header.msg_type = UdpPacketType::Syn as u8; - header.conn_id.set(conn_id); - header.len.set(8); - }, - Some(&magic.to_le_bytes()), - ) -} - -fn new_sack_packet(conn_id: u32, magic: u64) -> ZCPacket { - new_udp_packet( - |header| { - header.msg_type = UdpPacketType::Sack as u8; - header.conn_id.set(conn_id); - header.len.set(8); - }, - Some(&magic.to_le_bytes()), - ) -} - -pub fn new_hole_punch_packet(tid: u32, buf_len: u16) -> ZCPacket { - // generate a 128 bytes vec with random data - let mut rng = rand::rngs::StdRng::from_entropy(); - let mut buf = vec![0u8; buf_len as usize]; - rng.fill(&mut buf[..]); - new_udp_packet( - |header| { - header.msg_type = UdpPacketType::HolePunch as u8; - header.conn_id.set(tid); - header.len.set(buf_len); - }, - Some(&buf), - ) -} - -pub fn new_v6_hole_punch_packet( - dst: &SocketAddrV6, - preferred_src: Option, -) -> ZCPacket { - // generate a 128 bytes vec with random data - let mut body = V6HolePunchPacket::default(); - body.dst_ipv6.copy_from_slice(&dst.ip().octets()); - body.dst_port.set(dst.port()); - if let Some(src) = preferred_src { - body.preferred_src_ipv6.copy_from_slice(&src.ip.octets()); - body.preferred_src_ifindex.set(src.ifindex); - } - new_udp_packet( - |header| { - header.msg_type = UdpPacketType::V6HolePunch as u8; - header.conn_id.set(dst.port() as u32); - header - .len - .set(std::mem::size_of::() as u16); - }, - Some(body.as_bytes()), - ) -} - -pub fn new_v4_hole_punch_packet(dst: &SocketAddrV4) -> ZCPacket { - let mut body = V4HolePunchPacket::default(); - body.dst_ipv4.copy_from_slice(&dst.ip().octets()); - body.dst_port.set(dst.port()); - new_udp_packet( - |header| { - header.msg_type = UdpPacketType::V4HolePunch as u8; - header.conn_id.set(dst.port() as u32); - header - .len - .set(std::mem::size_of::() as u16); - }, - Some(body.as_bytes()), - ) -} - -fn extract_dst_addr_from_v4_hole_punch_packet(buf: &[u8]) -> Option { - let body = V4HolePunchPacket::ref_from_prefix(buf)?; - let ip = Ipv4Addr::from(body.dst_ipv4); - Some(SocketAddrV4::new(ip, body.dst_port.get())) -} - -fn extract_v6_hole_punch_packet(buf: &[u8]) -> Option<(SocketAddrV6, Option)> { - let body = V6HolePunchPacket::ref_from_prefix(buf)?; - let ip = Ipv6Addr::from(body.dst_ipv6); - let preferred_src_ipv6 = Ipv6Addr::from(body.preferred_src_ipv6); - let preferred_src = (!preferred_src_ipv6.is_unspecified()).then_some(PreferredIpv6Source { - ip: preferred_src_ipv6, - ifindex: body.preferred_src_ifindex.get(), - }); - Some(( - SocketAddrV6::new(ip, body.dst_port.get(), 0, 0), - preferred_src, - )) -} - -fn is_stun_packet(b: &[u8]) -> bool { - // stun has following pattern: - // 1. first two bits are 0b00 - // 2. magic cookie between 32-64 bits: 0x2112A442 - b[4..8] == [0x21, 0x12, 0xA4, 0x42] && b[0] & 0xC0 == 0 -} - -pub async fn send_v6_hole_punch_packet( - listener_port: u16, - dst_addr: SocketAddrV6, - preferred_src: Option, -) -> Result<(), TunnelError> { - let local_socket = UdpSocket::bind("[::1]:0").await?; - let udp_packet = new_v6_hole_punch_packet(&dst_addr, preferred_src); - let remote_addr = format!("[::1]:{}", listener_port) - .parse::() - .unwrap(); - local_socket - .send_to(&udp_packet.into_bytes(), remote_addr) - .await?; - Ok(()) -} - -pub async fn send_v4_hole_punch_packet( - listener_port: u16, - dst_addr: SocketAddrV4, -) -> Result<(), TunnelError> { - let local_socket = UdpSocket::bind("127.0.0.1:0").await?; - let udp_packet = new_v4_hole_punch_packet(&dst_addr); - let remote_addr = format!("127.0.0.1:{}", listener_port) - .parse::() - .unwrap(); - local_socket - .send_to(&udp_packet.into_bytes(), remote_addr) - .await?; - Ok(()) -} - -async fn respond_stun_packet( - socket: Arc, - addr: SocketAddr, - req_buf: Vec, -) -> Result<(), anyhow::Error> { - use crate::common::stun_codec_ext::*; - use bytecodec::{DecodeExt as _, EncodeExt as _}; - use stun_codec::{ - Message, MessageClass, MessageDecoder, MessageEncoder, - rfc5389::{attributes::XorMappedAddress, methods::BINDING}, - }; - - let mut decoder = MessageDecoder::::new(); - let req_msg = decoder - .decode_from_bytes(&req_buf) - .map_err(|e| anyhow::anyhow!("stun decode error: {:?}", e))? - .map_err(|e| anyhow::anyhow!("stun decode broken message error: {:?}", e))?; - - let tid = req_msg.transaction_id(); - // we only respond easytier stun req, whose tid has 0xdeadbeef prefix - if tid.as_bytes()[0..4] != [0xde, 0xad, 0xbe, 0xef] { - anyhow::bail!("stun req tid not from easytier"); - } - - let mut resp_msg = Message::::new( - MessageClass::SuccessResponse, - BINDING, - // we discard the prefix, make sure our implementation is not compatible with other stun client - u32_to_tid(tid_to_u32(&tid)), - ); - resp_msg.add_attribute(Attribute::XorMappedAddress(XorMappedAddress::new(addr))); - - let mut encoder = MessageEncoder::new(); - let rsp_buf = encoder - .encode_into_bytes(resp_msg.clone()) - .map_err(|e| anyhow::anyhow!("stun encode error: {:?}", e))?; - - let change_req = req_msg - .get_attribute::() - .map(|r| r.ip() || r.port()) - .unwrap_or(false); - - if !change_req { - socket - .send_to(&rsp_buf, addr) - .await - .with_context(|| "send stun response error")?; - } else { - // send from a new udp socket - let socket = if addr.is_ipv4() { - UdpSocket::bind("0.0.0.0:0").await? - } else { - UdpSocket::bind("[::]:0").await? - }; - socket.send_to(&rsp_buf, addr).await?; - } - - tracing::debug!(?addr, ?req_msg, ?change_req, "udp respond stun packet done"); - Ok(()) -} - -fn get_zcpacket_from_buf(buf: BytesMut, allow_stun: bool) -> Result { - let dg_size = buf.len(); - if dg_size < UDP_TUNNEL_HEADER_SIZE { - return Err(TunnelError::InvalidPacket(format!( - "udp packet size too small: {:?}, packet: {:?}", - dg_size, buf - ))); - } - - if allow_stun && is_stun_packet(&buf[..UDP_TUNNEL_HEADER_SIZE]) { - return Ok(ZCPacket::new_from_buf(buf, ZCPacketType::UDP)); - } - - let zc_packet = ZCPacket::new_from_buf(buf, ZCPacketType::UDP); - let header = zc_packet.udp_tunnel_header().unwrap(); - let payload_len = header.len.get() as usize; - if payload_len != dg_size - UDP_TUNNEL_HEADER_SIZE { - return Err(TunnelError::InvalidPacket(format!( - "udp packet payload len not match: header len: {:?}, real len: {:?}", - payload_len, dg_size - ))); - } - - Ok(zc_packet) -} - -#[instrument] -async fn forward_from_ring_to_udp( - mut ring_recv: RingStream, - socket: &Arc, - addr: &SocketAddr, - conn_id: u32, -) -> Option { - tracing::debug!("udp forward from ring to udp"); - loop { - let buf = ring_recv.next().await?; - let packet = match buf { - Ok(v) => v, - Err(e) => { - return Some(e); - } - }; - - let mut packet = packet.convert_type(ZCPacketType::UDP); - let udp_payload_len = packet.udp_payload().len(); - let header = packet.mut_udp_tunnel_header().unwrap(); - header.conn_id.set(conn_id); - header.len.set(udp_payload_len as u16); - header.msg_type = UdpPacketType::Data as u8; - - let buf = packet.into_bytes(); - tracing::trace!(?udp_payload_len, ?buf, "udp forward from ring to udp"); - let ret = socket.send_to(&buf, &addr).await; - if ret.is_err() { - return Some(TunnelError::IOError(ret.unwrap_err())); - } else if ret.unwrap() == 0 { - return None; - } - } -} - -async fn udp_recv_from_socket_forward_task( - socket: &UdpSocket, - buf: &mut BytesMut, - allow_stun: bool, -) -> Result<(ZCPacket, SocketAddr), TunnelError> { - loop { - reserve_buf(buf, UDP_DATA_MTU, UDP_DATA_MTU * 4); - let (dg_size, addr) = match socket.recv_buf_from(buf).await { - Ok(v) => v, - Err(e) => { - tracing::error!(?e, "udp recv from socket error"); - return Err(e.into()); - } - }; - tracing::trace!( - "udp recv packet: {:?}, buf: {:?}, size: {}", - addr, - buf, - dg_size - ); - - let zc_packet = match get_zcpacket_from_buf(buf.split(), allow_stun) { - Ok(v) => v, - Err(e) => { - tracing::warn!(?e, "udp get zc packet from buf error"); - continue; - } - }; - - return Ok((zc_packet, addr)); - } -} - -struct UdpConnection { - socket: Arc, - conn_id: u32, - dst_addr: SocketAddr, - - ring_sender: RingSink, - forward_task: AbortOnDropHandle<()>, -} - -impl UdpConnection { - pub fn new( - socket: Arc, - conn_id: u32, - dst_addr: SocketAddr, - ring_sender: RingSink, - ring_recv: RingStream, - close_event_sender: UdpCloseEventSender, - ) -> Self { - let s = socket.clone(); - let forward_task = AbortOnDropHandle::new(tokio::spawn(async move { - let close_event_sender = close_event_sender; - let err = forward_from_ring_to_udp(ring_recv, &s, &dst_addr, conn_id).await; - if let Err(e) = close_event_sender.send((dst_addr, err)) { - tracing::error!(?e, "udp send close event error"); - } - })); - Self { - socket, - conn_id, - dst_addr, - ring_sender, - forward_task, - } - } - - pub fn handle_packet_from_remote(&mut self, zc_packet: ZCPacket) -> Result<(), TunnelError> { - let header = zc_packet.udp_tunnel_header().unwrap(); - let conn_id = header.conn_id.get(); - - if header.msg_type != UdpPacketType::Data as u8 { - return Err(TunnelError::InvalidPacket("not data packet".to_owned())); - } - - if self.conn_id != conn_id { - return Err(TunnelError::ConnIdNotMatch(self.conn_id, conn_id)); - } - - if zc_packet.is_lossy() { - if let Err(e) = self.ring_sender.try_send(zc_packet) { - tracing::trace!(?e, "ring sender full, drop lossy packet"); - } - } else if self.ring_sender.force_send(zc_packet).is_err() { - tracing::trace!("ring sender full, reject non-lossy packet"); - return Err(TunnelError::BufferFull); - } - - Ok(()) - } -} - -#[derive(Clone)] -struct UdpTunnelListenerData { - local_url: url::Url, - socket: Option>, - sock_map: Arc>, - conn_send: Sender>, - close_event_sender: UdpCloseEventSender, -} - -impl UdpTunnelListenerData { - pub fn new( - local_url: url::Url, - conn_send: Sender>, - close_event_sender: UdpCloseEventSender, - ) -> Self { - Self { - local_url, - socket: None, - sock_map: Arc::new(DashMap::new()), - conn_send, - close_event_sender, - } - } - - async fn handle_new_connect(self, remote_addr: SocketAddr, zc_packet: ZCPacket) { - let udp_payload = zc_packet.udp_payload(); - if udp_payload.len() != 8 { - tracing::warn!( - "udp syn packet payload len not match: {:?}, packet: {:?}", - udp_payload.len(), - zc_packet, - ); - return; - } - let magic = u64::from_le_bytes(udp_payload[..8].try_into().unwrap()); - let conn_id = zc_packet.udp_tunnel_header().unwrap().conn_id.get(); - - tracing::info!(?conn_id, ?remote_addr, "udp connection accept handling",); - let socket = self.socket.as_ref().unwrap().clone(); - - let sack_buf = new_sack_packet(conn_id, magic).into_bytes(); - if self - .sock_map - .get(&remote_addr) - .is_some_and(|conn| conn.conn_id == conn_id) - { - if let Err(e) = socket.send_to(&sack_buf, remote_addr).await { - tracing::error!(?e, "udp resend sack packet error"); - } - tracing::debug!(?conn_id, ?remote_addr, "udp duplicate syn, resent sack"); - return; - } - - let ring_for_send_udp = Arc::new(RingTunnel::new(128)); - let ring_for_recv_udp = Arc::new(RingTunnel::new(128)); - tracing::debug!( - ?ring_for_send_udp, - ?ring_for_recv_udp, - "udp build tunnel for listener" - ); - - let new_internal_conn = || { - UdpConnection::new( - socket.clone(), - conn_id, - remote_addr, - RingSink::new(ring_for_recv_udp.clone()), - RingStream::new(ring_for_send_udp.clone()), - self.close_event_sender.clone(), - ) - }; - let duplicate_syn = match self.sock_map.entry(remote_addr) { - dashmap::mapref::entry::Entry::Occupied(entry) if entry.get().conn_id == conn_id => { - true - } - dashmap::mapref::entry::Entry::Occupied(mut entry) => { - entry.insert(new_internal_conn()); - false - } - dashmap::mapref::entry::Entry::Vacant(entry) => { - entry.insert(new_internal_conn()); - false - } - }; - if duplicate_syn { - if let Err(e) = socket.send_to(&sack_buf, remote_addr).await { - tracing::error!(?e, "udp resend sack packet error"); - } - tracing::debug!(?conn_id, ?remote_addr, "udp duplicate syn, resent sack"); - return; - } - - if let Err(e) = socket.send_to(&sack_buf, remote_addr).await { - self.sock_map - .remove_if(&remote_addr, |_, conn| conn.conn_id == conn_id); - tracing::error!(?e, "udp send sack packet error"); - return; - } - - let conn = Box::new(TunnelWrapper::new( - Box::new(RingStream::new(ring_for_recv_udp)), - Box::new(RingSink::new(ring_for_send_udp)), - Some(TunnelInfo { - tunnel_type: "udp".to_owned(), - local_addr: Some(self.local_url.clone().into()), - remote_addr: Some( - build_url_from_socket_addr(&remote_addr.to_string(), "udp").into(), - ), - resolved_remote_addr: Some( - build_url_from_socket_addr(&remote_addr.to_string(), "udp").into(), - ), - }), - )); - - tracing::info!(info = ?conn.info().unwrap().remote_addr, "udp connection accept done"); - - if let Err(e) = self.conn_send.send(conn).await { - tracing::warn!(?e, "udp send conn to accept channel error"); - } - } - - fn do_forward_one_packet_to_conn(&self, zc_packet: ZCPacket, addr: SocketAddr) { - let header = zc_packet.udp_tunnel_header().unwrap(); - if header.msg_type == UdpPacketType::Syn as u8 { - tokio::spawn(Self::handle_new_connect(self.clone(), addr, zc_packet)); - } else if is_stun_packet(header.as_bytes()) { - // ignore stun packet - tracing::debug!("udp forward packet ignore stun packet"); - let socket = self.socket.as_ref().unwrap().clone(); - tokio::spawn(async move { - let ret = respond_stun_packet(socket, addr, zc_packet.inner().to_vec()).await; - if let Err(e) = ret { - tracing::error!(?e, "udp respond stun packet error"); - } - }); - } else if header.msg_type == UdpPacketType::V4HolePunch as u8 { - if !addr.ip().is_loopback() { - tracing::warn!(?addr, "v4 hole punch packet should be from loopback"); - return; - } - if !addr.ip().is_ipv4() { - tracing::warn!(?addr, "v4 hole punch packet should be sent from ipv4"); - return; - } - let Some(dst_addr) = - extract_dst_addr_from_v4_hole_punch_packet(zc_packet.udp_payload()) - else { - tracing::warn!("invalid v4 hole punch packet"); - return; - }; - let socket = self.socket.as_ref().unwrap().clone(); - let udp_packet = new_hole_punch_packet(1, 32); - if let Err(e) = socket.try_send_to(&udp_packet.into_bytes(), SocketAddr::V4(dst_addr)) { - tracing::error!(?e, "udp send hole punch packet error"); - } - tracing::debug!(?dst_addr, "udp forward packet send hole punch packet"); - } else if header.msg_type == UdpPacketType::V6HolePunch as u8 { - if !addr.ip().is_loopback() { - tracing::warn!(?addr, "v6 hole punch packet should be from loopback"); - return; - } - if !addr.ip().is_ipv6() { - tracing::warn!(?addr, "v6 hole punch packet should be sent from ipv6"); - return; - } - let Some((dst_addr, preferred_src)) = - extract_v6_hole_punch_packet(zc_packet.udp_payload()) - else { - tracing::warn!("invalid v6 hole punch packet"); - return; - }; - let socket = self.socket.as_ref().unwrap().clone(); - let udp_packet = new_hole_punch_packet(1, 32); - let udp_packet = udp_packet.into_bytes(); - let sent_with_src = if let Some(src) = preferred_src { - match udp_src::send_to_with_src_ipv6( - &socket, - src.ip, - src.ifindex, - dst_addr, - &udp_packet, - ) { - Ok(ret) => { - tracing::debug!( - ?src, - ?dst_addr, - ?ret, - "udp forward packet send hole punch packet with preferred ipv6 source" - ); - true - } - Err(e) => { - tracing::debug!( - ?src, - ?dst_addr, - ?e, - "udp forward packet preferred ipv6 source failed, falling back" - ); - false - } - } - } else { - false - }; - if !sent_with_src - && let Err(e) = socket.try_send_to(&udp_packet, SocketAddr::V6(dst_addr)) - { - tracing::error!(?e, "udp send hole punch packet error"); - } - tracing::debug!( - ?dst_addr, - ?preferred_src, - "udp forward packet send hole punch packet" - ); - } else if header.msg_type != UdpPacketType::HolePunch as u8 { - let Some(mut conn) = self.sock_map.get_mut(&addr) else { - tracing::trace!(?header, "udp forward packet error, connection not found"); - return; - }; - if let Err(e) = conn.handle_packet_from_remote(zc_packet) { - tracing::trace!(?e, "udp forward packet error"); - } - } else { - tracing::trace!(?header, "udp forward packet ignore hole punch packet"); - } - } - - async fn do_forward_task(self) { - let socket = self.socket.as_ref().unwrap().clone(); - let mut buf = BytesMut::new(); - loop { - match udp_recv_from_socket_forward_task(&socket, &mut buf, true).await { - Ok((zc_packet, addr)) => self.do_forward_one_packet_to_conn(zc_packet, addr), - Err(e) => { - tracing::error!(?e, "udp recv packet error"); - break; - } - } - } - } -} - -pub struct UdpTunnelListener { - addr: url::Url, - socket: Option>, - - conn_recv: Receiver>, - data: UdpTunnelListenerData, - forward_tasks: Arc>>, - close_event_recv: Option, - socket_mark: Option, -} - -impl UdpTunnelListener { - pub fn new(addr: url::Url) -> Self { - let (close_event_send, close_event_recv) = unbounded_channel(); - let (conn_send, conn_recv) = channel(100); - Self { - addr: addr.clone(), - socket: None, - conn_recv, - data: UdpTunnelListenerData::new(addr, conn_send, close_event_send), - forward_tasks: Arc::new(std::sync::Mutex::new(JoinSet::new())), - close_event_recv: Some(close_event_recv), - socket_mark: None, - } - } - - pub fn set_socket_mark(&mut self, socket_mark: Option) { - self.socket_mark = socket_mark; - } - - pub fn new_with_socket(addr: url::Url, socket: Arc) -> Self { - let mut listener = Self::new(addr); - listener.socket = Some(socket); - listener - } - - pub fn get_socket(&self) -> Option> { - self.socket.clone() - } -} - -#[async_trait] -impl TunnelListener for UdpTunnelListener { - async fn listen(&mut self) -> Result<(), TunnelError> { - if self.socket.is_none() { - let addr = SocketAddr::from_url(self.addr.clone(), IpVersion::Both).await?; - let tunnel_url: TunnelUrl = self.addr.clone().into(); - self.socket = Some(Arc::new( - bind() - .addr(addr) - .only_v6(true) - .maybe_dev(tunnel_url.bind_dev()) - .maybe_socket_mark(self.socket_mark) - .call()?, - )); - } - self.data.socket = self.socket.clone(); - - self.addr - .set_port(Some(self.socket.as_ref().unwrap().local_addr()?.port())) - .unwrap(); - - self.forward_tasks - .lock() - .unwrap() - .spawn(self.data.clone().do_forward_task()); - - let sock_map = Arc::downgrade(&self.data.sock_map.clone()); - let mut close_recv = self.close_event_recv.take().unwrap(); - self.forward_tasks.lock().unwrap().spawn(async move { - while let Some((dst_addr, err)) = close_recv.recv().await { - if let Some(err) = err { - tracing::error!(?err, "udp close event error"); - } - if let Some(sock_map) = sock_map.upgrade() { - sock_map.remove(&dst_addr); - shrink_dashmap(&sock_map, None); - } - } - }); - - join_joinset_background(self.forward_tasks.clone(), "UdpTunnelListener".to_owned()); - - Ok(()) - } - - async fn accept(&mut self) -> Result, super::TunnelError> { - tracing::info!("start udp accept: {:?}", self.addr); - if let Some(conn) = self.conn_recv.recv().await { - return Ok(conn); - } - return Err(super::TunnelError::InternalError( - "udp accept error".to_owned(), - )); - } - - fn local_url(&self) -> url::Url { - self.addr.clone() - } - - fn get_conn_counter(&self) -> Arc> { - struct UdpTunnelConnCounter { - sock_map: Weak>, - } - - impl TunnelConnCounter for UdpTunnelConnCounter { - fn get(&self) -> Option { - self.sock_map.upgrade().map(|x| x.len() as u32) - } - } - - impl Debug for UdpTunnelConnCounter { - fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { - f.debug_struct("UdpTunnelConnCounter") - .field("sock_map_len", &self.get()) - .finish() - } - } - - Arc::new(Box::new(UdpTunnelConnCounter { - sock_map: Arc::downgrade(&self.data.sock_map.clone()), - })) - } -} - -#[derive(Debug)] -pub struct UdpTunnelConnector { - addr: url::Url, - bind_addrs: Vec, - ip_version: IpVersion, - resolved_addr: Option, - socket_mark: Option, -} - -impl UdpTunnelConnector { - pub fn new(addr: url::Url) -> Self { - Self { - addr, - bind_addrs: vec![], - ip_version: IpVersion::Both, - resolved_addr: None, - socket_mark: None, - } - } - - fn should_resend_syn_to_hole_punch_source( - recv_addr: SocketAddr, - expected_addr: SocketAddr, - ) -> bool { - recv_addr == expected_addr - } - - async fn wait_sack( - socket: &UdpSocket, - addr: SocketAddr, - conn_id: u32, - magic: u64, - ) -> Result { - let mut buf = BytesMut::new(); - buf.reserve(UDP_DATA_MTU); - - let (usize, recv_addr) = tokio::time::timeout( - tokio::time::Duration::from_secs(3), - socket.recv_buf_from(&mut buf), - ) - .await??; - let zc_packet = get_zcpacket_from_buf(buf.split(), false)?; - let header = zc_packet.udp_tunnel_header().unwrap(); - if header.msg_type == UdpPacketType::HolePunch as u8 { - tracing::debug!(?recv_addr, ?addr, "udp wait sack got hole punch packet"); - if Self::should_resend_syn_to_hole_punch_source(recv_addr, addr) { - let udp_packet = new_syn_packet(conn_id, magic).into_bytes(); - match socket.send_to(&udp_packet, recv_addr).await { - Ok(ret) => { - tracing::debug!(?recv_addr, ?ret, "udp send syn to hole punch source") - } - Err(e) => { - tracing::debug!(?recv_addr, ?e, "udp send syn to hole punch source failed") - } - } - } else { - tracing::debug!( - ?recv_addr, - ?addr, - "ignore hole punch packet from unexpected source" - ); - } - return Err(TunnelError::InvalidPacket( - "got hole punch packet while waiting for sack".to_owned(), - )); - } - if recv_addr != addr { - tracing::warn!(?recv_addr, ?addr, ?usize, "udp wait sack addr not match"); - } - - if header.conn_id.get() != conn_id { - return Err(super::TunnelError::ConnIdNotMatch( - header.conn_id.get(), - conn_id, - )); - } - - if header.msg_type != UdpPacketType::Sack as u8 { - return Err(TunnelError::InvalidPacket("not sack packet".to_owned())); - } - - let payload = zc_packet.udp_payload(); - if payload.len() != 8 { - return Err(TunnelError::InvalidPacket( - "udp sack packet payload len not match".to_owned(), - )); - } - - let sack_magic = u64::from_le_bytes(payload[..8].try_into().unwrap()); - if sack_magic != magic { - return Err(TunnelError::InvalidPacket( - "udp sack magic not match".to_owned(), - )); - } - - Ok(recv_addr) - } - - async fn wait_sack_loop( - socket: &UdpSocket, - addr: SocketAddr, - conn_id: u32, - magic: u64, - ) -> Result { - loop { - let ret = Self::wait_sack(socket, addr, conn_id, magic).await; - if ret.is_err() { - tracing::debug!(?ret, "udp wait sack error"); - continue; - } else { - return ret; - } - } - } - - async fn build_tunnel( - &self, - socket: Arc, - dst_addr: SocketAddr, - conn_id: u32, - ) -> Result, super::TunnelError> { - let ring_for_send_udp = Arc::new(RingTunnel::new(128)); - let ring_for_recv_udp = Arc::new(RingTunnel::new(128)); - tracing::debug!( - ?ring_for_send_udp, - ?ring_for_recv_udp, - "udp build tunnel for connector" - ); - - let (close_event_sender, mut close_event_recv) = unbounded_channel(); - - let ring_recv = RingStream::new(ring_for_send_udp.clone()); - let ring_sender = RingSink::new(ring_for_recv_udp.clone()); - let mut udp_conn = UdpConnection::new( - socket.clone(), - conn_id, - dst_addr, - ring_sender, - ring_recv, - close_event_sender, - ); - - let socket_clone = socket.clone(); - - let recv_loop = async move { - let mut buf = BytesMut::new(); - loop { - match udp_recv_from_socket_forward_task(&socket_clone, &mut buf, false).await { - Ok((zc_packet, addr)) => { - tracing::trace!(?addr, "connector udp forward task done"); - if let Err(e) = udp_conn.handle_packet_from_remote(zc_packet) { - tracing::trace!(?e, ?addr, "udp forward packet error"); - } - } - Err(e) => { - tracing::trace!(?e, "udp forward task error"); - break; - } - } - } - }; - tokio::spawn( - async move { - tokio::select! { - _ = close_event_recv.recv() => { - tracing::debug!("connector udp close event"); - } - _ = recv_loop => { - tracing::debug!("connector udp forward task done"); - } - } - } - .instrument(tracing::info_span!( - "udp forward from udp to ring", - ?conn_id, - ?dst_addr, - )), - ); - - Ok(Box::new(TunnelWrapper::new( - Box::new(RingStream::new(ring_for_recv_udp)), - Box::new(RingSink::new(ring_for_send_udp)), - Some(TunnelInfo { - tunnel_type: "udp".to_owned(), - local_addr: Some( - build_url_from_socket_addr(&socket.local_addr()?.to_string(), "udp").into(), - ), - remote_addr: Some(self.addr.clone().into()), - resolved_remote_addr: Some( - build_url_from_socket_addr(&dst_addr.to_string(), "udp").into(), - ), - }), - ))) - } - - pub async fn try_connect_with_socket( - &self, - socket: Arc, - addr: SocketAddr, - ) -> Result, super::TunnelError> { - tracing::warn!("udp connect: {:?}", self.addr); - - #[cfg(target_os = "windows")] - crate::arch::windows::disable_connection_reset(socket.as_ref())?; - - // send syn - let conn_id = rand::random(); - let magic = rand::random(); - let udp_packet = new_syn_packet(conn_id, magic).into_bytes(); - let ret = socket.send_to(&udp_packet, &addr).await?; - tracing::warn!(?udp_packet, ?ret, "udp send syn"); - let resend_task = AbortOnDropHandle::new(tokio::spawn({ - let socket = socket.clone(); - let udp_packet = udp_packet.clone(); - let resend_addr = addr; - async move { - loop { - tokio::time::sleep(Duration::from_millis(200)).await; - match socket.send_to(&udp_packet, &resend_addr).await { - Ok(ret) => tracing::trace!(?ret, ?resend_addr, "udp resend syn"), - Err(e) => { - tracing::debug!(?e, ?resend_addr, "udp resend syn failed"); - break; - } - } - } - } - })); - - // wait sack - let recv_addr = tokio::time::timeout( - tokio::time::Duration::from_secs(3), - Self::wait_sack_loop(&socket, addr, conn_id, magic), - ) - .await??; - drop(resend_task); - - if recv_addr != addr { - tracing::debug!(?recv_addr, ?addr, "udp connect addr not match"); - } - - self.build_tunnel(socket, recv_addr, conn_id).await - } - - async fn connect_with_default_bind( - &self, - addr: SocketAddr, - ) -> Result, super::TunnelError> { - // Route through bind() so socket_mark is applied consistently for - // both the None (no-op) and Some(_) paths. - let bind_addr: SocketAddr = if addr.is_ipv4() { - "0.0.0.0:0".parse().unwrap() - } else { - "[::]:0".parse().unwrap() - }; - let socket = bind::() - .addr(bind_addr) - .only_v6(true) - .maybe_socket_mark(self.socket_mark) - .call()?; - - return self.try_connect_with_socket(Arc::new(socket), addr).await; - } - - async fn connect_with_custom_bind( - &self, - addr: SocketAddr, - ) -> Result, super::TunnelError> { - let futures = FuturesUnordered::new(); - - for bind_addr in self.bind_addrs.iter() { - tracing::info!(?bind_addr, ?addr, "bind addr"); - match bind() - .addr(*bind_addr) - .only_v6(true) - .maybe_socket_mark(self.socket_mark) - .call() - { - Ok(socket) => futures.push(self.try_connect_with_socket(Arc::new(socket), addr)), - Err(error) => { - tracing::error!(?error, ?bind_addr, ?addr, "bind addr fail"); - continue; - } - } - } - wait_for_connect_futures(futures).await - } -} - -#[async_trait] -impl super::TunnelConnector for UdpTunnelConnector { - async fn connect(&mut self) -> Result, TunnelError> { - let addr = match self.resolved_addr { - Some(addr) => addr, - None => SocketAddr::from_url(self.addr.clone(), self.ip_version).await?, - }; - if self.bind_addrs.is_empty() || addr.is_ipv6() { - self.connect_with_default_bind(addr).await - } else { - self.connect_with_custom_bind(addr).await - } - } - - fn remote_url(&self) -> url::Url { - self.addr.clone() - } - - fn set_bind_addrs(&mut self, addrs: Vec) { - self.bind_addrs = addrs; - } - - fn set_ip_version(&mut self, ip_version: IpVersion) { - self.ip_version = ip_version; - } - - fn set_resolved_addr(&mut self, addr: SocketAddr) { - self.resolved_addr = Some(addr); - } - - fn set_socket_mark(&mut self, socket_mark: Option) { - self.socket_mark = socket_mark; - } -} - -#[cfg(test)] -mod tests { - use std::{net::IpAddr, time::Duration}; - - use futures::SinkExt; - use tokio::time::timeout; - - use super::*; - use crate::{ - common::global_ctx::tests::get_mock_global_ctx, - tunnel::{ - TunnelConnector, - common::{ - get_interface_name_by_ip, - tests::{_tunnel_bench, _tunnel_echo_server, _tunnel_pingpong, wait_for_condition}, - }, - packet_def::PacketType, - }, - }; - - fn new_udp_data_packet(conn_id: u32, packet_type: PacketType) -> ZCPacket { - let mut packet = ZCPacket::new_with_payload(b"udp-data").convert_type(ZCPacketType::UDP); - packet.fill_peer_manager_hdr(1, 2, packet_type as u8); - let udp_payload_len = packet.udp_payload().len(); - let header = packet.mut_udp_tunnel_header().unwrap(); - header.conn_id.set(conn_id); - header.msg_type = UdpPacketType::Data as u8; - header.len.set(udp_payload_len as u16); - packet - } - - fn assert_sync_packet_handler(_: fn(&mut UdpConnection, ZCPacket) -> Result<(), TunnelError>) {} - - #[test] - fn hole_punch_source_must_match_connect_addr_before_syn_resend() { - let expected_addr: SocketAddr = "198.51.100.10:11010".parse().unwrap(); - let same_port_different_ip: SocketAddr = "198.51.100.11:11010".parse().unwrap(); - let same_ip_different_port: SocketAddr = "198.51.100.10:11011".parse().unwrap(); - - assert!(UdpTunnelConnector::should_resend_syn_to_hole_punch_source( - expected_addr, - expected_addr - )); - assert!(!UdpTunnelConnector::should_resend_syn_to_hole_punch_source( - same_port_different_ip, - expected_addr - )); - assert!(!UdpTunnelConnector::should_resend_syn_to_hole_punch_source( - same_ip_different_port, - expected_addr - )); - } - - #[tokio::test] - async fn udp_pingpong() { - let listener = UdpTunnelListener::new("udp://0.0.0.0:5556".parse().unwrap()); - let connector = UdpTunnelConnector::new("udp://127.0.0.1:5556".parse().unwrap()); - _tunnel_pingpong(listener, connector).await; - } - - #[tokio::test] - async fn udp_connection_handler_uses_sync_nonblocking_ring_delivery() { - assert_sync_packet_handler(UdpConnection::handle_packet_from_remote); - - let socket = Arc::new(UdpSocket::bind("127.0.0.1:0").await.unwrap()); - let dst_addr = "127.0.0.1:1".parse().unwrap(); - let ring_for_send_udp = Arc::new(RingTunnel::new(8)); - let ring_for_recv_udp = Arc::new(RingTunnel::new(8)); - let (close_event_sender, _close_event_recv) = tokio::sync::mpsc::unbounded_channel(); - let mut conn = UdpConnection::new( - socket, - 7, - dst_addr, - RingSink::new(ring_for_recv_udp), - RingStream::new(ring_for_send_udp), - close_event_sender, - ); - - for _ in 0..16 { - conn.handle_packet_from_remote(new_udp_data_packet(7, PacketType::Data)) - .unwrap(); - } - - let mut got_buffer_full = false; - for _ in 0..16 { - match conn.handle_packet_from_remote(new_udp_data_packet(7, PacketType::Ping)) { - Ok(()) => {} - Err(TunnelError::BufferFull) => { - got_buffer_full = true; - break; - } - Err(e) => panic!("unexpected error: {e:?}"), - } - } - assert!(got_buffer_full); - } - - #[tokio::test] - async fn udp_bench() { - let listener = UdpTunnelListener::new("udp://0.0.0.0:5555".parse().unwrap()); - let connector = UdpTunnelConnector::new("udp://127.0.0.1:5555".parse().unwrap()); - _tunnel_bench(listener, connector).await - } - - #[tokio::test] - async fn udp_bench_with_bind() { - let listener = UdpTunnelListener::new("udp://127.0.0.1:5554".parse().unwrap()); - let mut connector = UdpTunnelConnector::new("udp://127.0.0.1:5554".parse().unwrap()); - connector.set_bind_addrs(vec!["127.0.0.1:0".parse().unwrap()]); - _tunnel_pingpong(listener, connector).await - } - - #[tokio::test] - #[should_panic] - async fn udp_bench_with_bind_fail() { - let listener = UdpTunnelListener::new("udp://127.0.0.1:5553".parse().unwrap()); - let mut connector = UdpTunnelConnector::new("udp://127.0.0.1:5553".parse().unwrap()); - connector.set_bind_addrs(vec!["10.0.0.1:0".parse().unwrap()]); - _tunnel_pingpong(listener, connector).await - } - - async fn send_random_data_to_socket(remote_url: url::Url) { - let socket = UdpSocket::bind("0.0.0.0:0").await.unwrap(); - socket - .connect(format!( - "{}:{}", - remote_url.host().unwrap(), - remote_url.port().unwrap() - )) - .await - .unwrap(); - - // get a random 100-len buf - loop { - let mut buf = vec![0u8; 100]; - rand::thread_rng().fill(&mut buf[..]); - socket.send(&buf).await.unwrap(); - tokio::time::sleep(tokio::time::Duration::from_millis(50)).await; - } - } - - #[tokio::test] - async fn udp_multiple_conns() { - let mut listener = UdpTunnelListener::new("udp://0.0.0.0:5557".parse().unwrap()); - listener.listen().await.unwrap(); - - let _lis = tokio::spawn(async move { - loop { - let ret = listener.accept().await.unwrap(); - assert_eq!( - ret.info() - .unwrap() - .local_addr - .unwrap_or_default() - .to_string(), - listener.local_url().to_string() - ); - tokio::spawn(async move { _tunnel_echo_server(ret, false).await }); - } - }); - - let mut connector1 = UdpTunnelConnector::new("udp://127.0.0.1:5557".parse().unwrap()); - let mut connector2 = UdpTunnelConnector::new("udp://127.0.0.1:5557".parse().unwrap()); - - let t1 = connector1.connect().await.unwrap(); - let t2 = connector2.connect().await.unwrap(); - - tokio::spawn(timeout( - Duration::from_secs(2), - send_random_data_to_socket(t1.info().unwrap().local_addr.unwrap().into()), - )); - tokio::spawn(timeout( - Duration::from_secs(2), - send_random_data_to_socket(t1.info().unwrap().remote_addr.unwrap().into()), - )); - tokio::spawn(timeout( - Duration::from_secs(2), - send_random_data_to_socket(t2.info().unwrap().remote_addr.unwrap().into()), - )); - - let sender1 = tokio::spawn(async move { - let (mut stream, mut sink) = t1.split(); - - for i in 0..10 { - sink.send(ZCPacket::new_with_payload("hello1".as_bytes())) - .await - .unwrap(); - let recv = stream.next().await.unwrap().unwrap(); - println!("t1 recv: {:?}, {:?}", recv, i); - assert_eq!(recv.payload(), "hello1".as_bytes()); - tokio::time::sleep(tokio::time::Duration::from_millis(100)).await; - } - }); - - let sender2 = tokio::spawn(async move { - let (mut stream, mut sink) = t2.split(); - - for i in 0..10 { - sink.send(ZCPacket::new_with_payload("hello2".as_bytes())) - .await - .unwrap(); - let recv = stream.next().await.unwrap().unwrap(); - println!("t2 recv: {:?}, {:?}", recv, i); - assert_eq!(recv.payload(), "hello2".as_bytes()); - tokio::time::sleep(tokio::time::Duration::from_millis(100)).await; - } - }); - - let _ = tokio::join!(sender1, sender2); - } - - #[tokio::test] - async fn bind_multi_ip_to_same_dev() { - let global_ctx = get_mock_global_ctx(); - let ips = global_ctx - .get_ip_collector() - .collect_ip_addrs() - .await - .interface_ipv4s; - if ips.is_empty() { - return; - } - let bind_dev = get_interface_name_by_ip(&IpAddr::V4(ips[0].into())); - - for ip in ips { - println!("bind to ip: {}, {:?}", ip, bind_dev); - let addr = SocketAddr::from_url( - format!("udp://{}:11111", ip).parse().unwrap(), - IpVersion::Both, - ) - .await - .unwrap(); - let _ = bind::() - .addr(addr) - .maybe_dev(bind_dev.clone()) - .only_v6(true) - .call() - .unwrap(); - } - } - - #[tokio::test] - async fn bind_same_port() { - println!("{}", "[::]:8888".parse::().unwrap()); - let mut listener = UdpTunnelListener::new("udp://[::]:31014".parse().unwrap()); - let mut listener2 = UdpTunnelListener::new("udp://0.0.0.0:31014".parse().unwrap()); - listener.listen().await.unwrap(); - listener2.listen().await.unwrap(); - } - - #[tokio::test] - async fn ipv6_pingpong() { - let listener = UdpTunnelListener::new("udp://[::1]:31015".parse().unwrap()); - let connector = UdpTunnelConnector::new("udp://[::1]:31015".parse().unwrap()); - _tunnel_pingpong(listener, connector).await - } - - #[tokio::test] - async fn ipv6_domain_pingpong() { - let listener = UdpTunnelListener::new("udp://[::1]:31016".parse().unwrap()); - let mut connector = - UdpTunnelConnector::new("udp://test.easytier.top:31016".parse().unwrap()); - connector.set_ip_version(IpVersion::V6); - _tunnel_pingpong(listener, connector).await; - - let listener = UdpTunnelListener::new("udp://127.0.0.1:31016".parse().unwrap()); - let mut connector = - UdpTunnelConnector::new("udp://test.easytier.top:31016".parse().unwrap()); - connector.set_ip_version(IpVersion::V4); - _tunnel_pingpong(listener, connector).await; - } - - #[tokio::test] - async fn test_alloc_port() { - // v4 - let mut listener = UdpTunnelListener::new("udp://0.0.0.0:0".parse().unwrap()); - listener.listen().await.unwrap(); - let port = listener.local_url().port().unwrap(); - assert!(port > 0); - - // v6 - let mut listener = UdpTunnelListener::new("udp://[::]:0".parse().unwrap()); - listener.listen().await.unwrap(); - let port = listener.local_url().port().unwrap(); - assert!(port > 0); - } - - #[tokio::test] - async fn test_conn_counter() { - let mut listener = UdpTunnelListener::new("udp://0.0.0.0:5556".parse().unwrap()); - let mut connector = UdpTunnelConnector::new("udp://127.0.0.1:5556".parse().unwrap()); - tokio::spawn(async move { - tokio::time::sleep(tokio::time::Duration::from_secs(1)).await; - let _c1 = connector.connect().await.unwrap(); - let _c2 = connector.connect().await.unwrap(); - }); - - let conn_counter = listener.get_conn_counter(); - - listener.listen().await.unwrap(); - let c1 = listener.accept().await.unwrap(); - assert_eq!(conn_counter.get(), Some(1)); - let c2 = listener.accept().await.unwrap(); - assert_eq!(conn_counter.get(), Some(2)); - - drop(c2); - wait_for_condition( - || async { conn_counter.get() == Some(1) }, - Duration::from_secs(1), - ) - .await; - - drop(c1); - wait_for_condition( - || async { conn_counter.get().unwrap_or(0) == 0 }, - Duration::from_secs(1), - ) - .await; - } - - #[test] - fn v6_hole_punch_packet_preserves_preferred_source_ifindex() { - let dst_addr = "[2001:db8::1]:10001".parse::().unwrap(); - let preferred_src = PreferredIpv6Source { - ip: "2001:db8::2".parse().unwrap(), - ifindex: 42, - }; - - let packet = new_v6_hole_punch_packet(&dst_addr, Some(preferred_src)); - let (parsed_dst_addr, parsed_preferred_src) = - extract_v6_hole_punch_packet(packet.udp_payload()).unwrap(); - - assert_eq!(parsed_dst_addr, dst_addr); - assert_eq!(parsed_preferred_src, Some(preferred_src)); - } - - #[tokio::test] - async fn test_v6_hole_punch_packet() { - let mut lis = UdpTunnelListener::new("udp://[::]:0".parse().unwrap()); - lis.listen().await.unwrap(); - - // a socket to receive forwarded hole punch packets - let socket = Arc::new(UdpSocket::bind("[::]:0").await.unwrap()); - let socket_clone = socket.clone(); - let t = tokio::spawn(async move { - let mut buf = BytesMut::new(); - buf.resize(128, 0); - socket_clone.recv_from(&mut buf).await.unwrap(); - }); - - tracing::info!("lis local addr: {:?}", lis.local_url()); - tracing::info!("socket local addr: {:?}", socket.local_addr().unwrap()); - - tokio::time::sleep(tokio::time::Duration::from_millis(100)).await; - - // a socket to send v6 hole punch packets - send_v6_hole_punch_packet( - lis.local_url().port().unwrap(), - match socket.local_addr().unwrap() { - std::net::SocketAddr::V6(addr_v6) => addr_v6, - _ => panic!("Expected an IPv6 address"), - }, - None, - ) - .await - .unwrap(); - - tokio::time::timeout(tokio::time::Duration::from_secs(2), t) - .await - .expect("Timeout waiting for v6 hole punch packet") - .unwrap(); - } - - #[tokio::test] - async fn test_v4_hole_punch_packet() { - let mut lis = UdpTunnelListener::new("udp://0.0.0.0:0".parse().unwrap()); - lis.listen().await.unwrap(); - - let socket = Arc::new(UdpSocket::bind("127.0.0.1:0").await.unwrap()); - let socket_clone = socket.clone(); - let t = tokio::spawn(async move { - let mut buf = BytesMut::new(); - buf.resize(128, 0); - socket_clone.recv_from(&mut buf).await.unwrap(); - }); - - tokio::time::sleep(tokio::time::Duration::from_millis(100)).await; - - send_v4_hole_punch_packet( - lis.local_url().port().unwrap(), - match socket.local_addr().unwrap() { - std::net::SocketAddr::V4(addr_v4) => addr_v4, - _ => panic!("Expected an IPv4 address"), - }, - ) - .await - .unwrap(); - - tokio::time::timeout(tokio::time::Duration::from_secs(2), t) - .await - .expect("Timeout waiting for v4 hole punch packet") - .unwrap(); - } -} diff --git a/easytier/src/tunnel/udp_src.rs b/easytier/src/tunnel/udp_src.rs deleted file mode 100644 index 71f10916..00000000 --- a/easytier/src/tunnel/udp_src.rs +++ /dev/null @@ -1,210 +0,0 @@ -use std::{ - io, - net::{Ipv6Addr, SocketAddrV6}, -}; - -use tokio::net::UdpSocket; - -#[cfg(unix)] -pub(crate) fn send_to_with_src_ipv6( - socket: &UdpSocket, - src_ip: Ipv6Addr, - src_ifindex: u32, - dst_addr: SocketAddrV6, - buf: &[u8], -) -> io::Result { - #[cfg(target_env = "ohos")] - { - let _ = (socket, src_ip, src_ifindex, dst_addr, buf); - return Err(io::Error::new( - io::ErrorKind::Unsupported, - "sending UDP with a selected IPv6 source is not supported on OHOS", - )); - } - - #[cfg(not(target_env = "ohos"))] - { - use std::{mem, os::fd::AsRawFd, ptr}; - - use nix::libc; - - #[repr(align(8))] - struct ControlBuffer([u8; 128]); - - #[cfg(target_os = "android")] - let ipi6_ifindex: libc::c_int = i32::try_from(src_ifindex).map_err(|_| { - io::Error::new( - io::ErrorKind::InvalidInput, - "IPv6 source interface index is out of range", - ) - })?; - #[cfg(not(target_os = "android"))] - let ipi6_ifindex: libc::c_uint = src_ifindex; - - let pktinfo = libc::in6_pktinfo { - ipi6_addr: libc::in6_addr { - s6_addr: src_ip.octets(), - }, - ipi6_ifindex, - }; - let mut iov = libc::iovec { - iov_base: buf.as_ptr() as *mut libc::c_void, - iov_len: buf.len(), - }; - let dst_addr = socket2::SockAddr::from(std::net::SocketAddr::V6(dst_addr)); - let control_len = unsafe { - libc::CMSG_SPACE(mem::size_of::() as libc::c_uint) as usize - }; - let mut control = ControlBuffer([0u8; 128]); - if control_len > control.0.len() { - return Err(io::Error::new( - io::ErrorKind::InvalidInput, - "IPv6 packet info control buffer is too small", - )); - } - - let mut msg = unsafe { mem::zeroed::() }; - msg.msg_name = dst_addr.as_ptr() as *mut libc::c_void; - msg.msg_namelen = dst_addr.len() as _; - msg.msg_iov = &mut iov; - msg.msg_iovlen = 1; - msg.msg_control = control.0.as_mut_ptr() as *mut libc::c_void; - msg.msg_controllen = control_len as _; - msg.msg_flags = 0; - - unsafe { - let cmsg = libc::CMSG_FIRSTHDR(&msg); - if cmsg.is_null() { - return Err(io::Error::new( - io::ErrorKind::InvalidInput, - "IPv6 packet info control buffer is invalid", - )); - } - (*cmsg).cmsg_level = libc::IPPROTO_IPV6; - (*cmsg).cmsg_type = libc::IPV6_PKTINFO; - (*cmsg).cmsg_len = - libc::CMSG_LEN(mem::size_of::() as libc::c_uint) as _; - ptr::write(libc::CMSG_DATA(cmsg) as *mut libc::in6_pktinfo, pktinfo); - - let ret = libc::sendmsg(socket.as_raw_fd(), &msg, 0); - if ret < 0 { - Err(io::Error::last_os_error()) - } else { - Ok(ret as usize) - } - } - } -} - -#[cfg(windows)] -pub(crate) fn send_to_with_src_ipv6( - socket: &UdpSocket, - src_ip: Ipv6Addr, - src_ifindex: u32, - dst_addr: SocketAddrV6, - buf: &[u8], -) -> io::Result { - use std::{mem, os::windows::io::AsRawSocket, ptr}; - - use windows::{ - Win32::Networking::WinSock::{ - CMSGHDR, IN6_ADDR, IN6_ADDR_0, IN6_PKTINFO, IPPROTO_IPV6, IPV6_PKTINFO, SOCKET, - SOCKET_ERROR, WSABUF, WSAGetLastError, WSAMSG, WSASendMsg, - }, - core::PSTR, - }; - - fn cmsghdr_align(length: usize) -> usize { - (length + mem::align_of::() - 1) & !(mem::align_of::() - 1) - } - - fn cmsgdata_align(length: usize) -> usize { - (length + mem::align_of::() - 1) & !(mem::align_of::() - 1) - } - - fn cmsg_len(length: usize) -> usize { - cmsgdata_align(mem::size_of::()) + length - } - - fn cmsg_space(length: usize) -> usize { - cmsgdata_align(mem::size_of::() + cmsghdr_align(length)) - } - - fn cmsg_data(cmsg: *mut CMSGHDR) -> *mut u8 { - (cmsg as usize + cmsgdata_align(mem::size_of::())) as *mut u8 - } - - #[repr(align(8))] - struct ControlBuffer([u8; 128]); - - let dst = socket2::SockAddr::from(std::net::SocketAddr::V6(dst_addr)); - let mut data = WSABUF { - len: buf.len() as u32, - buf: PSTR(buf.as_ptr() as *mut u8), - }; - let control_len = cmsg_space(mem::size_of::()); - let mut control = ControlBuffer([0u8; 128]); - if control_len > control.0.len() { - return Err(io::Error::new( - io::ErrorKind::InvalidInput, - "IPv6 packet info control buffer is too small", - )); - } - let mut msg = WSAMSG { - name: dst.as_ptr() as *mut _, - namelen: dst.len(), - lpBuffers: &mut data, - dwBufferCount: 1, - Control: WSABUF { - len: control_len as u32, - buf: PSTR(control.0.as_mut_ptr()), - }, - dwFlags: 0, - }; - - let pktinfo = IN6_PKTINFO { - ipi6_addr: IN6_ADDR { - u: IN6_ADDR_0 { - Byte: src_ip.octets(), - }, - }, - ipi6_ifindex: src_ifindex, - }; - - unsafe { - let cmsg = control.0.as_mut_ptr() as *mut CMSGHDR; - (*cmsg).cmsg_level = IPPROTO_IPV6.0; - (*cmsg).cmsg_type = IPV6_PKTINFO; - (*cmsg).cmsg_len = cmsg_len(mem::size_of::()); - ptr::write(cmsg_data(cmsg) as *mut IN6_PKTINFO, pktinfo); - msg.Control.len = control_len as u32; - - let mut sent = 0; - let ret = WSASendMsg( - SOCKET(socket.as_raw_socket() as usize), - &msg, - 0, - Some(&mut sent), - None, - None, - ); - if ret == SOCKET_ERROR { - return Err(io::Error::from_raw_os_error(WSAGetLastError().0)); - } - Ok(sent as usize) - } -} - -#[cfg(not(any(unix, windows)))] -pub(crate) fn send_to_with_src_ipv6( - _socket: &UdpSocket, - _src_ip: Ipv6Addr, - _src_ifindex: u32, - _dst_addr: SocketAddrV6, - _buf: &[u8], -) -> io::Result { - Err(io::Error::new( - io::ErrorKind::Unsupported, - "sending UDP with a selected IPv6 source is not supported on this platform", - )) -} diff --git a/easytier/src/tunnel/unix.rs b/easytier/src/tunnel/unix.rs deleted file mode 100644 index 7cba3f2a..00000000 --- a/easytier/src/tunnel/unix.rs +++ /dev/null @@ -1,218 +0,0 @@ -use std::path::Path; - -use async_trait::async_trait; -use tokio::net::{UnixListener, UnixStream, unix::SocketAddr}; - -use super::TunnelInfo; - -use super::{ - IpVersion, Tunnel, TunnelError, TunnelListener, - common::{FramedReader, FramedWriter, TunnelWrapper}, -}; - -const MAX_PACKET_SIZE: usize = 4096; - -fn url_from_unix_socket_addr(addr: SocketAddr) -> Option { - addr.as_pathname() - .and_then(|p| p.to_str()) - .and_then(|s| format!("unix://{}", s).parse().ok()) -} - -#[derive(Debug)] -pub struct UnixSocketTunnelListener { - addr: url::Url, - listener: Option, - unlink_on_drop: bool, -} - -impl UnixSocketTunnelListener { - pub fn new(addr: url::Url) -> Self { - UnixSocketTunnelListener { - addr, - listener: None, - unlink_on_drop: true, - } - } - - async fn do_accept(&self) -> Result, std::io::Error> { - let listener = self.listener.as_ref().unwrap(); - let (stream, _) = listener.accept().await?; - - let remote_addr = stream.peer_addr().ok().and_then(url_from_unix_socket_addr); - - let info = TunnelInfo { - tunnel_type: "unix".to_owned(), - local_addr: Some(self.local_url().into()), - remote_addr: remote_addr.clone().map(Into::into), - resolved_remote_addr: remote_addr.map(Into::into), - }; - - let (r, w) = stream.into_split(); - Ok(Box::new(TunnelWrapper::new( - FramedReader::new(r, MAX_PACKET_SIZE), - FramedWriter::new(w), - Some(info), - ))) - } - - fn set_unlink_on_drop(&mut self, unlink: bool) { - self.unlink_on_drop = unlink; - } -} - -#[async_trait] -impl TunnelListener for UnixSocketTunnelListener { - async fn listen(&mut self) -> Result<(), TunnelError> { - self.listener = None; - let path_str = self.addr.path(); - let path = Path::new(path_str); - - let listener = UnixListener::bind(path)?; - self.listener = Some(listener); - Ok(()) - } - - async fn accept(&mut self) -> Result, super::TunnelError> { - loop { - match self.do_accept().await { - Ok(ret) => return Ok(ret), - Err(e) => { - use std::io::ErrorKind::*; - if matches!( - e.kind(), - NotConnected | ConnectionAborted | ConnectionRefused | ConnectionReset - ) { - tracing::warn!(?e, "accept fail with retryable error: {:?}", e); - continue; - } - tracing::warn!(?e, "accept fail"); - return Err(e.into()); - } - } - } - } - - fn local_url(&self) -> url::Url { - self.addr.clone() - } -} - -#[derive(Debug)] -pub struct UnixSocketTunnelConnector { - addr: url::Url, -} - -impl UnixSocketTunnelConnector { - pub fn new(addr: url::Url) -> Self { - UnixSocketTunnelConnector { addr } - } -} - -#[async_trait] -impl super::TunnelConnector for UnixSocketTunnelConnector { - async fn connect(&mut self) -> Result, super::TunnelError> { - let path_str = self.addr.path(); - let path = Path::new(path_str); - tracing::info!(url = ?self.addr, "connect unix socket start"); - let stream = UnixStream::connect(path).await?; - tracing::info!(url = ?self.addr, "connect unix socket succ"); - - let local_addr = stream.local_addr().ok().and_then(url_from_unix_socket_addr); - - let info = TunnelInfo { - tunnel_type: "unix".to_owned(), - local_addr: local_addr.map(Into::into), - remote_addr: Some(self.addr.clone().into()), - resolved_remote_addr: Some(self.addr.clone().into()), - }; - - let (r, w) = stream.into_split(); - Ok(Box::new(TunnelWrapper::new( - FramedReader::new(r, MAX_PACKET_SIZE), - FramedWriter::new(w), - Some(info), - ))) - } - - fn remote_url(&self) -> url::Url { - self.addr.clone() - } - - fn set_ip_version(&mut self, _ip_version: IpVersion) { - // IP version is not applicable to UNIX sockets - } -} - -impl Drop for UnixSocketTunnelListener { - fn drop(&mut self) { - if self.unlink_on_drop { - let _ = std::fs::remove_file(self.addr.path()); - } - } -} - -#[cfg(test)] -mod tests { - use crate::tunnel::common::tests::{_tunnel_bench, _tunnel_pingpong}; - - use super::*; - - #[tokio::test] - async fn unix_socket_pingpong() { - let listener = - UnixSocketTunnelListener::new("unix:///tmp/easytier-test.sock".parse().unwrap()); - let connector = - UnixSocketTunnelConnector::new("unix:///tmp/easytier-test.sock".parse().unwrap()); - _tunnel_pingpong(listener, connector).await - } - - #[tokio::test] - async fn unix_socket_bench() { - let listener = - UnixSocketTunnelListener::new("unix:///tmp/easytier-test-bench.sock".parse().unwrap()); - let connector = - UnixSocketTunnelConnector::new("unix:///tmp/easytier-test-bench.sock".parse().unwrap()); - _tunnel_bench(listener, connector).await - } - - #[tokio::test] - async fn unlink_on_drop() { - let listener = - UnixSocketTunnelListener::new("unix:///tmp/easytier-test-exists.sock".parse().unwrap()); - let connector = UnixSocketTunnelConnector::new( - "unix:///tmp/easytier-test-exists.sock".parse().unwrap(), - ); - _tunnel_pingpong(listener, connector).await; - - let mut listener = - UnixSocketTunnelListener::new("unix:///tmp/easytier-test-exists.sock".parse().unwrap()); - listener.set_unlink_on_drop(false); - let connector = UnixSocketTunnelConnector::new( - "unix:///tmp/easytier-test-exists.sock".parse().unwrap(), - ); - _tunnel_pingpong(listener, connector).await; - - let mut listener = - UnixSocketTunnelListener::new("unix:///tmp/easytier-test-exists.sock".parse().unwrap()); - let result = listener.listen().await; - assert!( - matches!(result, Err(TunnelError::IOError(err)) if err.kind() == std::io::ErrorKind::AddrInUse) - ) - } - - #[tokio::test] - async fn bind_file_exists() { - use std::fs; - - let path = "/tmp/easytier-test-exists.sock"; - fs::File::create(path).unwrap(); - let mut listener = - UnixSocketTunnelListener::new("unix:///tmp/easytier-test-exists.sock".parse().unwrap()); - let result = listener.listen().await; - - fs::remove_file(path).unwrap(); - assert!( - matches!(result, Err(TunnelError::IOError(err)) if err.kind() == std::io::ErrorKind::AddrInUse) - ) - } -} diff --git a/easytier/src/tunnel/websocket.rs b/easytier/src/tunnel/websocket.rs index d7d1e6ae..7cbd128d 100644 --- a/easytier/src/tunnel/websocket.rs +++ b/easytier/src/tunnel/websocket.rs @@ -1,81 +1,250 @@ -use super::{ - FromUrl, IpVersion, Tunnel, TunnelConnector, TunnelError, TunnelListener, - common::{TunnelWrapper, wait_for_connect_futures}, - insecure_tls::{get_insecure_tls_cert, init_crypto_provider}, - packet_def::{ZCPacket, ZCPacketType}, -}; +use super::FromUrl; use crate::tunnel::common::bind; -use crate::{proto::common::TunnelInfo, tunnel::insecure_tls::get_insecure_tls_client_config}; -use anyhow::Context; +use crate::{proto::common::TunnelInfo, socket::tcp::RuntimeTcpSocket}; +use anyhow::Context as _; use bytes::BytesMut; +use cidr::IpCidr; +use easytier_core::{ + packet::{ZCPacket, ZCPacketType}, + socket::tcp::VirtualTcpSocket, + tunnel::{IpVersion, Tunnel, TunnelError, wrapper::TunnelWrapper}, +}; use forwarded_header_value::ForwardedHeaderValue; -use futures::{SinkExt, StreamExt, stream::FuturesUnordered}; -use pnet::ipnetwork::IpNetwork; +use futures::{SinkExt, StreamExt}; use std::{ - net::SocketAddr, + net::{IpAddr, SocketAddr}, sync::{Arc, LazyLock}, time::Duration, }; -use tokio::{ - net::{TcpListener, TcpSocket, TcpStream}, - time::timeout, -}; +use tokio::{net::TcpListener, time::timeout}; use tokio_rustls::TlsAcceptor; use tokio_util::either::Either; use tokio_websockets::{ClientBuilder, Limits, MaybeTlsStream, Message, ServerBuilder}; -use zerocopy::AsBytes; +use zerocopy::AsBytes as _; -fn is_wss(addr: &url::Url) -> Result { - match addr.scheme() { +pub(crate) const CONNECT_TIMEOUT: Duration = Duration::from_secs(20); +pub(crate) const SERVER_HANDSHAKE_TIMEOUT: Duration = Duration::from_secs(3); + +static TRUSTED_PROXIES: LazyLock> = LazyLock::new(|| { + [ + "127.0.0.0/8", + "10.0.0.0/8", + "172.16.0.0/12", + "192.168.0.0/16", + "::1/128", + "fc00::/7", + ] + .into_iter() + .map(|cidr| cidr.parse().unwrap()) + .collect() +}); + +fn trusted_proxy_contains(ip: IpAddr) -> bool { + TRUSTED_PROXIES.iter().any(|cidr| match (cidr, ip) { + (IpCidr::V4(cidr), IpAddr::V4(ip)) => cidr.contains(&ip), + (IpCidr::V6(cidr), IpAddr::V6(ip)) => cidr.contains(&ip), + _ => false, + }) +} + +fn websocket_error(error: impl std::fmt::Display) -> TunnelError { + TunnelError::ProtocolError(format!("websocket error: {error}")) +} + +fn is_wss(url: &url::Url) -> Result { + match url.scheme() { "ws" => Ok(false), "wss" => Ok(true), - _ => Err(TunnelError::InvalidProtocol(addr.scheme().to_string())), + scheme => Err(TunnelError::InvalidProtocol(scheme.to_owned())), } } -async fn sink_from_zc_packet(msg: ZCPacket) -> Result { - Ok(Message::binary(msg.tunnel_payload_bytes().freeze())) +async fn sink_from_zc_packet(packet: ZCPacket) -> Result { + Ok(Message::binary(packet.tunnel_payload_bytes().freeze())) } async fn map_from_ws_message( - msg: Result, + message: Result, ) -> Option> { - if let Err(e) = msg { - tracing::error!(?e, "recv from websocket error"); - return Some(Err(TunnelError::WebSocketError(e))); - } - - let msg = msg.unwrap(); - if msg.is_close() { + let message = match message { + Ok(message) => message, + Err(error) => { + tracing::error!(?error, "recv from websocket error"); + return Some(Err(websocket_error(error))); + } + }; + if message.is_close() { tracing::warn!("recv close message from websocket"); return None; } - - if !msg.is_binary() { - let msg = format!("{:?}", msg); - tracing::error!(?msg, "Invalid packet"); - return Some(Err(TunnelError::InvalidPacket(msg))); + if !message.is_binary() { + let message = format!("{message:?}"); + tracing::error!(?message, "Invalid packet"); + return Some(Err(TunnelError::InvalidPacket(message))); } - Some(Ok(ZCPacket::new_from_buf( - BytesMut::from(msg.into_payload().as_bytes()), + BytesMut::from(message.into_payload().as_bytes()), ZCPacketType::DummyTunnel, ))) } -static TRUSTED_PROXIES: LazyLock> = LazyLock::new(|| { - [ - "127.0.0.0/8", // IPV4 Loopback - "10.0.0.0/8", // IPV4 Private Networks - "172.16.0.0/12", - "192.168.0.0/16", - "::1/128", // IPV6 Loopback - "fc00::/7", // IPV6 Private network - ] - .into_iter() - .map(|s| s.parse().unwrap()) - .collect() -}); +#[derive(Debug)] +struct SkipServerVerification(Arc); + +impl SkipServerVerification { + fn new(provider: Arc) -> Arc { + Arc::new(Self(provider)) + } +} + +impl rustls::client::danger::ServerCertVerifier for SkipServerVerification { + fn verify_server_cert( + &self, + _end_entity: &rustls::pki_types::CertificateDer<'_>, + _intermediates: &[rustls::pki_types::CertificateDer<'_>], + _server_name: &rustls::pki_types::ServerName<'_>, + _ocsp: &[u8], + _now: rustls::pki_types::UnixTime, + ) -> Result { + Ok(rustls::client::danger::ServerCertVerified::assertion()) + } + + fn verify_tls12_signature( + &self, + message: &[u8], + cert: &rustls::pki_types::CertificateDer<'_>, + dss: &rustls::DigitallySignedStruct, + ) -> Result { + rustls::crypto::verify_tls12_signature( + message, + cert, + dss, + &self.0.signature_verification_algorithms, + ) + } + + fn verify_tls13_signature( + &self, + message: &[u8], + cert: &rustls::pki_types::CertificateDer<'_>, + dss: &rustls::DigitallySignedStruct, + ) -> Result { + rustls::crypto::verify_tls13_signature( + message, + cert, + dss, + &self.0.signature_verification_algorithms, + ) + } + + fn supported_verify_schemes(&self) -> Vec { + self.0.signature_verification_algorithms.supported_schemes() + } +} + +fn init_crypto_provider() { + let _ = + rustls::crypto::CryptoProvider::install_default(rustls::crypto::ring::default_provider()); +} + +fn get_insecure_tls_client_config() -> rustls::ClientConfig { + init_crypto_provider(); + let provider = rustls::crypto::CryptoProvider::get_default().unwrap(); + let mut config = rustls::ClientConfig::builder() + .dangerous() + .with_custom_certificate_verifier(SkipServerVerification::new(provider.clone())) + .with_no_client_auth(); + config.enable_sni = true; + config.enable_early_data = false; + config +} + +fn get_insecure_tls_cert<'a>() -> ( + Vec>, + rustls::pki_types::PrivateKeyDer<'a>, +) { + let cert = rcgen::generate_simple_self_signed(vec!["localhost".into()]).unwrap(); + let cert_der = cert.serialize_der().unwrap(); + let private_key = cert.serialize_private_key_der(); + let private_key = rustls::pki_types::PrivatePkcs8KeyDer::from(private_key); + (vec![cert_der.into()], private_key.into()) +} + +pub(crate) async fn upgrade_accepted( + stream: S, + local_url: url::Url, +) -> Result, TunnelError> +where + S: VirtualTcpSocket, +{ + let peer_addr = stream.peer_addr()?; + let mut remote_url = socket_url(local_url.scheme(), peer_addr); + let stream = if is_wss(&local_url)? { + init_crypto_provider(); + let (certificates, private_key) = get_insecure_tls_cert(); + let config = rustls::ServerConfig::builder() + .with_no_client_auth() + .with_single_cert(certificates, private_key) + .with_context(|| "Failed to create server config")?; + Either::Left(TlsAcceptor::from(Arc::new(config)).accept(stream).await?) + } else { + Either::Right(stream) + }; + + let (request, stream) = ServerBuilder::new() + .limits(Limits::unlimited()) + .max_headers(128) + .accept(stream) + .await + .map_err(websocket_error)?; + + if trusted_proxy_contains(peer_addr.ip()) + && let Some(forwarded) = request + .headers() + .get("Forwarded") + .and_then(|value| value.to_str().ok()) + .and_then(|value| ForwardedHeaderValue::from_forwarded(value).ok()) + .or_else(|| { + request + .headers() + .get("X-Forwarded-For") + .and_then(|value| value.to_str().ok()) + .and_then(|value| ForwardedHeaderValue::from_x_forwarded_for(value).ok()) + }) + && let Some(ip) = forwarded.remotest_forwarded_for_ip() + { + remote_url + .set_host(Some(&ip.to_string())) + .map_err(|_| TunnelError::InvalidAddr(format!("invalid forwarded ip {ip}")))?; + remote_url + .query_pairs_mut() + .append_pair("proxy", &peer_addr.to_string()); + } + + let (write, read) = stream.split(); + let remote_url: crate::proto::common::Url = remote_url.into(); + let info = TunnelInfo { + tunnel_type: local_url.scheme().to_owned(), + local_addr: Some(local_url.into()), + remote_addr: Some(remote_url.clone()), + resolved_remote_addr: Some(remote_url), + }; + Ok(Box::new(TunnelWrapper::new( + read.filter_map(map_from_ws_message), + write + .sink_map_err(websocket_error) + .with(sink_from_zc_packet::), + Some(info), + ))) +} + +fn socket_url(scheme: &str, addr: SocketAddr) -> url::Url { + let mut url = url::Url::parse(&format!("{scheme}://0.0.0.0")) + .expect("WebSocket transport scheme should be a valid URL scheme"); + url.set_ip_host(addr.ip()).unwrap(); + url.set_port(Some(addr.port())).unwrap(); + url +} #[derive(Debug)] pub struct WsTunnelListener { @@ -97,77 +266,7 @@ impl WsTunnelListener { self.socket_mark = socket_mark; } - async fn try_accept(&self, stream: TcpStream) -> Result, TunnelError> { - let peer_addr = stream.peer_addr()?; - let mut remote_addr = - super::build_url_from_socket_addr(&peer_addr.to_string(), self.addr.scheme()); - - let stream = if is_wss(&self.addr)? { - init_crypto_provider(); - let (certs, key) = get_insecure_tls_cert(); - let config = rustls::ServerConfig::builder() - .with_no_client_auth() - .with_single_cert(certs, key) - .with_context(|| "Failed to create server config")?; - - let stream = TlsAcceptor::from(Arc::new(config)).accept(stream).await?; - Either::Left(stream) - } else { - Either::Right(stream) - }; - - let (request, stream) = ServerBuilder::new() - .limits(Limits::unlimited()) - .max_headers(128) - .accept(stream) - .await?; - - if TRUSTED_PROXIES - .iter() - .any(|net| net.contains(peer_addr.ip())) - && let Some(forwarded) = request - .headers() - .get("Forwarded") - .and_then(|f| f.to_str().ok()) - .and_then(|f| ForwardedHeaderValue::from_forwarded(f).ok()) - .or_else(|| { - request - .headers() - .get("X-Forwarded-For") - .and_then(|f| f.to_str().ok()) - .and_then(|f| ForwardedHeaderValue::from_x_forwarded_for(f).ok()) - }) - && let Some(ip) = forwarded.remotest_forwarded_for_ip() - { - remote_addr - .set_host(Some(&ip.to_string())) - .map_err(|_| TunnelError::InvalidAddr(format!("invalid forwarded ip {}", ip)))?; - remote_addr - .query_pairs_mut() - .append_pair("proxy", &peer_addr.to_string()); - } - - let (write, read) = stream.split(); - let remote_addr: crate::proto::common::Url = remote_addr.into(); - - let info = TunnelInfo { - tunnel_type: self.addr.scheme().to_owned(), - local_addr: Some(self.local_url().into()), - remote_addr: Some(remote_addr.clone()), - resolved_remote_addr: Some(remote_addr), - }; - - Ok(Box::new(TunnelWrapper::new( - read.filter_map(map_from_ws_message), - write.with(sink_from_zc_packet), - Some(info), - ))) - } -} - -#[async_trait::async_trait] -impl TunnelListener for WsTunnelListener { - async fn listen(&mut self) -> Result<(), TunnelError> { + async fn listen_tunnel(&mut self) -> Result<(), TunnelError> { self.listener = None; let addr = SocketAddr::from_url(self.addr.clone(), IpVersion::Both).await?; @@ -185,13 +284,18 @@ impl TunnelListener for WsTunnelListener { Ok(()) } - async fn accept(&mut self) -> Result, super::TunnelError> { + async fn accept_tunnel(&mut self) -> Result, TunnelError> { loop { let listener = self.listener.as_ref().unwrap(); // only fail on tcp accept error let (stream, _) = listener.accept().await?; stream.set_nodelay(true).unwrap(); - match timeout(Duration::from_secs(3), self.try_accept(stream)).await { + match timeout( + SERVER_HANDSHAKE_TIMEOUT, + upgrade_accepted(RuntimeTcpSocket::new(stream), self.addr.clone()), + ) + .await + { Ok(Ok(tunnel)) => return Ok(tunnel), e => { tracing::error!(?e, ?self, "Failed to accept ws/wss tunnel"); @@ -200,217 +304,82 @@ impl TunnelListener for WsTunnelListener { } } } +} + +#[async_trait::async_trait] +impl easytier_core::socket::SocketListener for WsTunnelListener { + type Accepted = Box; + + async fn listen(&mut self) -> anyhow::Result<()> { + Ok(self.listen_tunnel().await?) + } + + async fn accept(&mut self) -> anyhow::Result { + Ok(self.accept_tunnel().await?) + } fn local_url(&self) -> url::Url { self.addr.clone() } } -pub struct WsTunnelConnector { - addr: url::Url, - ip_version: IpVersion, - resolved_addr: Option, +pub(crate) async fn upgrade_connected( + stream: S, + remote_url: url::Url, +) -> Result, TunnelError> +where + S: VirtualTcpSocket, +{ + let is_wss = is_wss(&remote_url)?; + let local_addr = stream.local_addr()?; + let resolved_remote_addr = stream.peer_addr()?; + let info = TunnelInfo { + tunnel_type: remote_url.scheme().to_owned(), + local_addr: Some( + super::build_url_from_socket_addr(&local_addr.to_string(), remote_url.scheme()).into(), + ), + remote_addr: Some(remote_url.clone().into()), + resolved_remote_addr: Some( + super::build_url_from_socket_addr( + &resolved_remote_addr.to_string(), + remote_url.scheme(), + ) + .into(), + ), + }; - bind_addrs: Vec, - socket_mark: Option, -} + let client = ClientBuilder::from_uri(http::Uri::try_from(remote_url.to_string()).unwrap()) + .max_headers(128); + let stream: MaybeTlsStream = if is_wss { + init_crypto_provider(); + let tls = tokio_rustls::TlsConnector::from(Arc::new(get_insecure_tls_client_config())); + let sni = remote_url.domain().unwrap_or("localhost").to_owned(); + let server_name = rustls::pki_types::ServerName::try_from(sni) + .map_err(|_| TunnelError::InvalidProtocol("Invalid SNI".to_owned()))?; + MaybeTlsStream::Rustls(tls.connect(server_name, stream).await?) + } else { + MaybeTlsStream::Plain(stream) + }; -impl WsTunnelConnector { - pub fn new(addr: url::Url) -> Self { - WsTunnelConnector { - addr, - ip_version: IpVersion::Both, - resolved_addr: None, - - bind_addrs: vec![], - socket_mark: None, - } - } - - async fn connect_with( - addr: url::Url, - socket_addr: SocketAddr, - tcp_socket: TcpSocket, - ) -> Result, TunnelError> { - let is_wss = is_wss(&addr)?; - let stream = tcp_socket.connect(socket_addr).await?; - if let Err(error) = stream.set_nodelay(true) { - tracing::warn!(?error, "set_nodelay fail in ws connect"); - } - - let info = TunnelInfo { - tunnel_type: addr.scheme().to_owned(), - local_addr: Some( - super::build_url_from_socket_addr( - &stream.local_addr()?.to_string(), - addr.scheme().to_string().as_str(), - ) - .into(), - ), - remote_addr: Some(addr.clone().into()), - resolved_remote_addr: Some( - super::build_url_from_socket_addr(&socket_addr.to_string(), addr.scheme()).into(), - ), - }; - - let c = ClientBuilder::from_uri(http::Uri::try_from(addr.to_string()).unwrap()) - .max_headers(128); - let stream: MaybeTlsStream = if is_wss { - init_crypto_provider(); - let tls_conn = - tokio_rustls::TlsConnector::from(Arc::new(get_insecure_tls_client_config())); - // Modify SNI logic: use "localhost" as SNI for url without domain to avoid IP blocking. - let sni = match addr.domain() { - None => "localhost".to_string(), - Some(domain) => domain.to_string(), - }; - let server_name = rustls::pki_types::ServerName::try_from(sni) - .map_err(|_| TunnelError::InvalidProtocol("Invalid SNI".to_string()))?; - let stream = tls_conn.connect(server_name, stream).await?; - MaybeTlsStream::Rustls(stream) - } else { - MaybeTlsStream::Plain(stream) - }; - - let (client, _) = c.connect_on(stream).await?; - let (write, read) = client.split(); - let read = read.filter_map(map_from_ws_message); - let write = write.with(sink_from_zc_packet); - Ok(Box::new(TunnelWrapper::new(read, write, Some(info)))) - } - - async fn connect_with_default_bind( - &self, - addr: SocketAddr, - ) -> Result, super::TunnelError> { - let socket = if addr.is_ipv4() { - TcpSocket::new_v4()? - } else { - TcpSocket::new_v6()? - }; - crate::tunnel::common::apply_socket_mark( - &socket2::SockRef::from(&socket), - self.socket_mark, - )?; - Self::connect_with(self.addr.clone(), addr, socket).await - } - - async fn connect_with_custom_bind( - &self, - addr: SocketAddr, - ) -> Result, super::TunnelError> { - let futures = FuturesUnordered::new(); - - for bind_addr in self.bind_addrs.iter() { - tracing::info!(?bind_addr, ?addr, "bind addr"); - match bind() - .addr(*bind_addr) - .only_v6(true) - .maybe_socket_mark(self.socket_mark) - .call() - { - Ok(socket) => futures.push(Self::connect_with(self.addr.clone(), addr, socket)), - Err(error) => { - tracing::error!(?bind_addr, ?addr, ?error, "bind addr fail"); - continue; - } - } - } - - wait_for_connect_futures(futures).await - } -} - -#[async_trait::async_trait] -impl TunnelConnector for WsTunnelConnector { - async fn connect(&mut self) -> Result, TunnelError> { - let addr = match self.resolved_addr { - Some(addr) => addr, - None => SocketAddr::from_url(self.addr.clone(), self.ip_version).await?, - }; - if self.bind_addrs.is_empty() || addr.is_ipv6() { - self.connect_with_default_bind(addr).await - } else { - self.connect_with_custom_bind(addr).await - } - } - - fn remote_url(&self) -> url::Url { - self.addr.clone() - } - - fn set_ip_version(&mut self, ip_version: IpVersion) { - self.ip_version = ip_version; - } - - fn set_bind_addrs(&mut self, addrs: Vec) { - self.bind_addrs = addrs; - } - - fn set_resolved_addr(&mut self, addr: SocketAddr) { - self.resolved_addr = Some(addr); - } - - fn set_socket_mark(&mut self, socket_mark: Option) { - self.socket_mark = socket_mark; - } + let (client, _) = client.connect_on(stream).await.map_err(websocket_error)?; + let (write, read) = client.split(); + Ok(Box::new(TunnelWrapper::new( + read.filter_map(map_from_ws_message), + write + .sink_map_err(websocket_error) + .with(sink_from_zc_packet::), + Some(info), + ))) } #[cfg(test)] pub mod tests { use super::*; - use crate::tunnel::common::tests::_tunnel_pingpong; - use tokio::io::{AsyncReadExt, AsyncWriteExt}; - - #[rstest::rstest] - #[tokio::test] - #[serial_test::serial] - async fn ws_pingpong(#[values("ws", "wss")] proto: &str) { - let listener = WsTunnelListener::new(format!("{}://0.0.0.0:25556", proto).parse().unwrap()); - let connector = - WsTunnelConnector::new(format!("{}://127.0.0.1:25556", proto).parse().unwrap()); - _tunnel_pingpong(listener, connector).await - } - - #[rstest::rstest] - #[tokio::test] - #[serial_test::serial] - async fn ws_pingpong_bind(#[values("ws", "wss")] proto: &str) { - let listener = WsTunnelListener::new(format!("{}://0.0.0.0:25557", proto).parse().unwrap()); - let mut connector = - WsTunnelConnector::new(format!("{}://127.0.0.1:25557", proto).parse().unwrap()); - connector.set_bind_addrs(vec!["127.0.0.1:0".parse().unwrap()]); - _tunnel_pingpong(listener, connector).await - } - - // TODO: tokio-websockets cannot correctly handle close, benchmark case is disabled - // #[rstest::rstest] - // #[tokio::test] - // #[serial_test::serial] - // async fn ws_bench(#[values("ws", "wss")] proto: &str) { - // enable_log(); - // let listener = WSTunnelListener::new(format!("{}://0.0.0.0:25557", proto).parse().unwrap()); - // let connector = - // WSTunnelConnector::new(format!("{}://127.0.0.1:25557", proto).parse().unwrap()); - // _tunnel_bench(listener, connector).await - // } - - #[tokio::test] - async fn ws_accept_wss() { - let mut listener = WsTunnelListener::new("wss://0.0.0.0:25558".parse().unwrap()); - listener.listen().await.unwrap(); - let j = tokio::spawn(async move { - let _ = listener.accept().await; - }); - - let mut connector = WsTunnelConnector::new("ws://127.0.0.1:25558".parse().unwrap()); - connector.connect().await.unwrap_err(); - - let mut connector = WsTunnelConnector::new("wss://127.0.0.1:25558".parse().unwrap()); - connector.connect().await.unwrap(); - - j.abort(); - } + use easytier_core::socket::SocketListener; + use tokio::{ + io::{AsyncReadExt, AsyncWriteExt}, + net::TcpSocket, + }; #[tokio::test] async fn ws_forwarded() { diff --git a/easytier/src/tunnel/wireguard.rs b/easytier/src/tunnel/wireguard.rs index 7efc81f6..8268da09 100644 --- a/easytier/src/tunnel/wireguard.rs +++ b/easytier/src/tunnel/wireguard.rs @@ -8,22 +8,12 @@ use std::{ use quanta::Instant; -use super::{ - FromUrl, IpVersion, Tunnel, TunnelError, TunnelInfo, TunnelListener, TunnelUrl, ZCPacketSink, - ZCPacketStream, - common::wait_for_connect_futures, - generate_digest_from_str, - packet_def::{PEER_MANAGER_HEADER_SIZE, ZCPacketType}, - ring::create_ring_tunnel_pair, -}; -use crate::tunnel::common::{BindDev, bind}; +use super::FromUrl; use crate::{ - common::shrink_dashmap, - tunnel::{ - build_url_from_socket_addr, - common::TunnelWrapper, - packet_def::{WG_TUNNEL_HEADER_SIZE, ZCPacket}, - }, + common::{netns::NetNS, shrink_dashmap}, + proto::common::TunnelInfo, + socket::udp::{RuntimeUdpSessionSocketListener, new_runtime_udp_session_listener}, + tunnel::{TunnelUrl, build_url_from_socket_addr}, }; use anyhow::Context; use async_recursion::async_recursion; @@ -35,21 +25,29 @@ use boringtun::{ use bytes::BytesMut; use crossbeam::atomic::AtomicCell; use dashmap::DashMap; -use futures::{SinkExt, StreamExt, stream::FuturesUnordered}; +use easytier_core::tunnel::ring::create_ring_tunnel_pair; +use easytier_core::tunnel::{IpVersion, Tunnel, TunnelError, ZCPacketSink, ZCPacketStream}; +use easytier_core::{ + connectivity::transport::ConnectedUdpSession, + packet::{PEER_MANAGER_HEADER_SIZE, WG_TUNNEL_HEADER_SIZE, ZCPacket, ZCPacketType}, + socket::udp::{ + UdpBindOptions, UdpSession, UdpSessionAcceptKind, UdpSessionListenRequest, + UdpSessionProtocol, UdpSessionSocket, + }, + tunnel::wrapper::TunnelWrapper, +}; +use futures::{SinkExt, StreamExt}; use rand::RngCore; use tokio::{ - net::UdpSocket, sync::{Mutex, mpsc::unbounded_channel}, task::JoinSet, }; const MAX_PACKET: usize = 2048; -#[derive(Debug, Clone)] +#[derive(Debug, Clone, Copy, PartialEq, Eq)] enum WgType { - // used by easytier peer, need remove/add ip header for in/out wg msg InternalUse, - // used by wireguard peer, keep original ip header ExternalUse, } @@ -57,31 +55,16 @@ enum WgType { pub struct WgConfig { my_secret_key: StaticSecret, my_public_key: PublicKey, - peer_secret_key: StaticSecret, peer_public_key: PublicKey, - wg_type: WgType, } impl WgConfig { pub fn new_from_network_identity(network_name: &str, network_secret: &str) -> Self { - let mut my_sec = [0u8; 32]; - generate_digest_from_str(network_name, network_secret, &mut my_sec); - - let my_secret_key = StaticSecret::from(my_sec); - let my_public_key = PublicKey::from(&my_secret_key); - let peer_secret_key = StaticSecret::from(my_sec); - let peer_public_key = my_public_key; - - WgConfig { - my_secret_key, - my_public_key, - peer_secret_key, - peer_public_key, - - wg_type: WgType::InternalUse, - } + let mut secret = [0u8; 32]; + super::generate_digest_from_str(network_name, network_secret, &mut secret); + Self::new_internal(secret, secret) } pub fn new_for_portal(server_key_seed: &str, client_key_seed: &str) -> Self { @@ -92,11 +75,24 @@ impl WgConfig { my_public_key: server_cfg.my_public_key, peer_secret_key: client_cfg.my_secret_key, peer_public_key: client_cfg.my_public_key, - wg_type: WgType::ExternalUse, } } + pub fn new_internal(my_secret_key: [u8; 32], peer_secret_key: [u8; 32]) -> Self { + let my_secret_key = StaticSecret::from(my_secret_key); + let my_public_key = PublicKey::from(&my_secret_key); + let peer_secret_key = StaticSecret::from(peer_secret_key); + let peer_public_key = PublicKey::from(&peer_secret_key); + Self { + my_secret_key, + my_public_key, + peer_secret_key, + peer_public_key, + wg_type: WgType::InternalUse, + } + } + pub fn my_secret_key(&self) -> &[u8] { self.my_secret_key.as_bytes() } @@ -112,14 +108,49 @@ impl WgConfig { pub fn peer_public_key(&self) -> &[u8] { self.peer_public_key.as_bytes() } + + pub fn is_internal(&self) -> bool { + self.wg_type == WgType::InternalUse + } +} + +#[cfg(test)] +mod config_tests { + use super::*; + + #[test] + fn network_identity_produces_matching_internal_key_pairs() { + let config = WgConfig::new_from_network_identity("network", "secret"); + + assert!(config.is_internal()); + assert_eq!(config.my_secret_key(), config.peer_secret_key()); + assert_eq!(config.my_public_key(), config.peer_public_key()); + assert_eq!( + PublicKey::from(&StaticSecret::from( + <[u8; 32]>::try_from(config.my_secret_key()).unwrap() + )) + .as_bytes(), + config.my_public_key() + ); + } + + #[test] + fn portal_uses_distinct_external_key_pairs() { + let config = WgConfig::new_for_portal("server-seed", "client-seed"); + + assert!(!config.is_internal()); + assert_ne!(config.my_secret_key(), config.peer_secret_key()); + assert_ne!(config.my_public_key(), config.peer_public_key()); + } } #[derive(Clone)] struct WgPeerData { - udp: Arc, // only for send + session: Arc, endpoint: SocketAddr, tunn: Arc>, - wg_type: WgType, + internal_use: bool, + access_time: Arc>, stopped: Arc, } @@ -127,7 +158,7 @@ impl Debug for WgPeerData { fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result { f.debug_struct("WgPeerData") .field("endpoint", &self.endpoint) - .field("local", &self.udp.local_addr()) + .field("local", &self.session.local_addr()) .finish() } } @@ -137,7 +168,7 @@ impl WgPeerData { async fn handle_one_packet_from_me(&self, zc_packet: ZCPacket) -> Result<(), anyhow::Error> { let mut send_buf = vec![0u8; MAX_PACKET]; - let packet = if matches!(self.wg_type, WgType::InternalUse) { + let packet = if self.internal_use { let mut zc_packet = zc_packet.convert_type(ZCPacketType::WG); Self::fill_ip_header(&mut zc_packet); zc_packet.into_bytes() @@ -159,8 +190,8 @@ impl WgPeerData { match encapsulate_result { TunnResult::WriteToNetwork(packet) => { - self.udp - .send_to(packet, self.endpoint) + self.session + .send(packet) .await .context("Failed to send encrypted IP packet to WireGuard endpoint.")?; tracing::debug!( @@ -192,6 +223,7 @@ impl WgPeerData { mut sink: S, recv_buf: &[u8], ) { + self.access_time.store(Instant::now()); let mut send_buf = vec![0u8; MAX_PACKET]; let data = recv_buf; let decapsulate_result = { @@ -203,7 +235,7 @@ impl WgPeerData { match decapsulate_result { TunnResult::WriteToNetwork(packet) => { - match self.udp.send_to(packet, self.endpoint).await { + match self.session.send(packet).await { Ok(_) => {} Err(e) => { tracing::error!( @@ -218,7 +250,7 @@ impl WgPeerData { let mut send_buf = vec![0u8; MAX_PACKET]; match peer.decapsulate(None, &[], &mut send_buf) { TunnResult::WriteToNetwork(packet) => { - match self.udp.send_to(packet, self.endpoint).await { + match self.session.send(packet).await { Ok(_) => {} Err(e) => { tracing::error!( @@ -242,7 +274,7 @@ impl WgPeerData { packet.len() ); let mut b = BytesMut::new(); - if matches!(self.wg_type, WgType::InternalUse) { + if self.internal_use { b.resize(WG_TUNNEL_HEADER_SIZE, 0); b.extend_from_slice(self.remove_ip_header(packet, packet[0] >> 4 == 4)); } else { @@ -273,7 +305,7 @@ impl WgPeerData { "Sending routine packet of {} bytes to WireGuard endpoint", packet.len() ); - match self.udp.send_to(packet, self.endpoint).await { + match self.session.send(packet).await { Ok(_) => {} Err(e) => { tracing::error!( @@ -343,7 +375,8 @@ impl WgPeerData { struct WgPeer { tunn: Option>, - udp: Arc, // only for send + _session_guard: Box, + session: Arc, config: WgConfig, endpoint: SocketAddr, @@ -352,22 +385,28 @@ struct WgPeer { data: Option, tasks: JoinSet<()>, - access_time: AtomicCell, + access_time: Arc>, } impl WgPeer { - fn new(udp: Arc, config: WgConfig, endpoint: SocketAddr) -> Self { + fn new( + session_guard: Box, + session: Arc, + config: WgConfig, + endpoint: SocketAddr, + ) -> Self { WgPeer { tunn: Some(Mutex::new(Tunn::new( - config.my_secret_key.clone(), - config.peer_public_key, + StaticSecret::from(<[u8; 32]>::try_from(config.my_secret_key()).unwrap()), + PublicKey::from(<[u8; 32]>::try_from(config.peer_public_key()).unwrap()), None, None, rand::thread_rng().next_u32(), None, ))), - udp, + _session_guard: session_guard, + session, config, endpoint, sink: std::sync::Mutex::new(None), @@ -375,7 +414,7 @@ impl WgPeer { data: None, tasks: JoinSet::new(), - access_time: AtomicCell::new(Instant::now()), + access_time: Arc::new(AtomicCell::new(Instant::now())), } } @@ -390,26 +429,17 @@ impl WgPeer { .store(true, std::sync::atomic::Ordering::Relaxed); } - async fn handle_packet_from_peer(&self, packet: &[u8]) { - self.access_time.store(Instant::now()); - tracing::trace!("Received {} bytes from peer", packet.len()); - let data = self.data.as_ref().unwrap(); - // TODO: improve this - let mut sink = self.sink.lock().unwrap().take().unwrap(); - data.handle_one_packet_from_peer(&mut sink, packet).await; - self.sink.lock().unwrap().replace(sink); - } - fn start_and_get_tunnel(&mut self) -> Box { let (stunnel, ctunnel) = create_ring_tunnel_pair(); let (stream, sink) = stunnel.split(); let data = WgPeerData { - udp: self.udp.clone(), + session: self.session.clone(), endpoint: self.endpoint, tunn: Arc::new(self.tunn.take().unwrap()), - wg_type: self.config.wg_type.clone(), + internal_use: self.config.is_internal(), + access_time: self.access_time.clone(), stopped: Arc::new(AtomicBool::new(false)), }; @@ -450,8 +480,29 @@ impl WgPeer { handshake_init.into() } - fn udp_socket(&self) -> Arc { - self.udp.clone() + fn spawn_session_recv_task(&mut self, first_packet: Option>) { + let session = self.session.clone(); + let data = self.data.as_ref().unwrap().clone(); + let mut sink = self.sink.lock().unwrap().take().unwrap(); + self.tasks.spawn(async move { + if let Some(packet) = first_packet { + data.handle_one_packet_from_peer(&mut sink, &packet).await; + } + + let mut buf = vec![0u8; MAX_PACKET]; + loop { + let n = match session.recv(&mut buf).await { + Ok(n) => n, + Err(e) => { + tracing::error!("Failed to receive wg packet: {}", e); + data.stopped + .store(true, std::sync::atomic::Ordering::Relaxed); + break; + } + }; + data.handle_one_packet_from_peer(&mut sink, &buf[..n]).await; + } + }); } } @@ -460,16 +511,26 @@ type ConnReceiver = tokio::sync::mpsc::UnboundedReceiver>; pub struct WgTunnelListener { addr: url::Url, + session_listener: Option>, + socket_mark: Option, config: WgConfig, - udp: Option>, conn_recv: ConnReceiver, conn_send: Option, wg_peer_map: Arc>>, tasks: JoinSet<()>, - socket_mark: Option, +} + +impl Debug for WgTunnelListener { + fn fmt(&self, formatter: &mut Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("WgTunnelListener") + .field("addr", &self.addr) + .field("listening", &self.session_listener.is_some()) + .finish() + } } impl WgTunnelListener { @@ -477,16 +538,16 @@ impl WgTunnelListener { let (conn_send, conn_recv) = unbounded_channel(); WgTunnelListener { addr, + session_listener: None, + socket_mark: None, config, - udp: None, conn_recv, conn_send: Some(conn_send), wg_peer_map: Arc::new(DashMap::new()), tasks: JoinSet::new(), - socket_mark: None, } } @@ -494,12 +555,8 @@ impl WgTunnelListener { self.socket_mark = socket_mark; } - fn get_udp_socket(&self) -> Arc { - self.udp.as_ref().unwrap().clone() - } - - async fn handle_udp_incoming( - socket: Arc, + async fn accept_udp_sessions( + session_listener: Arc, config: WgConfig, conn_sender: ConnSender, peer_map: Arc>>, @@ -517,429 +574,287 @@ impl WgTunnelListener { } }); - let mut buf = vec![0u8; MAX_PACKET]; loop { - let Ok((n, addr)) = socket.recv_from(&mut buf).await else { - tracing::error!("Failed to receive from UDP socket"); - break; + let session = match session_listener.accept_session().await { + Ok(session) => Arc::new(session) as Arc, + Err(e) => { + tracing::error!("Failed to accept wg udp session: {}", e); + break; + } + }; + let addr = match session.peer_addr() { + Ok(addr) => addr, + Err(e) => { + tracing::error!("Failed to get wg session peer addr: {}", e); + continue; + } + }; + if peer_map.contains_key(&addr) { + continue; + } + let local_addr = match session.local_addr() { + Ok(addr) => addr, + Err(e) => { + tracing::error!("Failed to get wg session local addr: {}", e); + continue; + } }; - let data = &buf[..n]; - tracing::trace!(?n, ?addr, "Received bytes from peer"); - - if !peer_map.contains_key(&addr) { - tracing::info!("New peer: {}", addr); - let mut wg = WgPeer::new(socket.clone(), config.clone(), addr); - let (stream, sink) = wg.start_and_get_tunnel().split(); - let tunnel = Box::new(TunnelWrapper::new( - stream, - sink, - Some(TunnelInfo { - tunnel_type: "wg".to_owned(), - local_addr: Some( - build_url_from_socket_addr( - &socket.local_addr().unwrap().to_string(), - "wg", - ) - .into(), - ), - remote_addr: Some( - build_url_from_socket_addr(&addr.to_string(), "wg").into(), - ), - resolved_remote_addr: Some( - build_url_from_socket_addr(&addr.to_string(), "wg").into(), - ), - }), - )); - if let Err(e) = conn_sender.send(tunnel) { - tracing::error!("Failed to send tunnel to conn_sender: {}", e); - } - peer_map.insert(addr, Arc::new(wg)); + tracing::info!("New peer: {}", addr); + let mut wg = WgPeer::new( + Box::new(session_listener.clone()), + session, + config.clone(), + addr, + ); + let (stream, sink) = wg.start_and_get_tunnel().split(); + wg.spawn_session_recv_task(None); + let tunnel = Box::new(TunnelWrapper::new( + stream, + sink, + Some(TunnelInfo { + tunnel_type: "wg".to_owned(), + local_addr: Some( + build_url_from_socket_addr(&local_addr.to_string(), "wg").into(), + ), + remote_addr: Some(build_url_from_socket_addr(&addr.to_string(), "wg").into()), + resolved_remote_addr: Some( + build_url_from_socket_addr(&addr.to_string(), "wg").into(), + ), + }), + )); + if let Err(e) = conn_sender.send(tunnel) { + tracing::error!("Failed to send tunnel to conn_sender: {}", e); + break; } - - let peer = peer_map.get(&addr).unwrap().clone(); - peer.handle_packet_from_peer(data).await; + peer_map.insert(addr, Arc::new(wg)); } } -} -#[async_trait] -impl TunnelListener for WgTunnelListener { - async fn listen(&mut self) -> Result<(), TunnelError> { - let addr = SocketAddr::from_url(self.addr.clone(), IpVersion::Both).await?; - let tunnel_url: TunnelUrl = self.addr.clone().into(); - self.udp = Some(Arc::new( - bind() - .addr(addr) - .only_v6(true) - .maybe_dev(tunnel_url.bind_dev()) - .maybe_socket_mark(self.socket_mark) - .call()?, - )); - self.addr - .set_port(Some(self.udp.as_ref().unwrap().local_addr()?.port())) - .unwrap(); + async fn listen_tunnel(&mut self) -> Result<(), TunnelError> { + if self.session_listener.is_some() { + return Ok(()); + } - self.tasks.spawn(Self::handle_udp_incoming( - self.get_udp_socket(), + let local_addr = SocketAddr::from_url(self.addr.clone(), IpVersion::Both).await?; + let bind = UdpBindOptions::port_bound_listener(local_addr) + .with_socket_mark(self.socket_mark) + .with_bind_device(TunnelUrl::from(self.addr.clone()).bind_dev()) + .with_only_v6(true); + let mut session_listener = new_runtime_udp_session_listener( + self.addr.clone(), + UdpSessionListenRequest::new(bind), + UdpSessionAcceptKind::Classified(UdpSessionProtocol::WireGuard), + NetNS::new(None), + ); + easytier_core::socket::SocketListener::listen(&mut session_listener).await?; + let session_listener = Arc::new(session_listener); + + self.tasks.spawn(Self::accept_udp_sessions( + session_listener.clone(), self.config.clone(), self.conn_send.take().unwrap(), self.wg_peer_map.clone(), )); + self.session_listener = Some(session_listener); Ok(()) } - async fn accept(&mut self) -> Result, super::TunnelError> { + async fn accept_tunnel(&mut self) -> Result, TunnelError> { if let Some(tunnel) = self.conn_recv.recv().await { tracing::info!(?tunnel, "Accepted tunnel"); return Ok(tunnel); } Err(TunnelError::Shutdown) } - - fn local_url(&self) -> url::Url { - self.addr.clone() - } -} - -#[derive(Clone)] -pub struct WgTunnelConnector { - addr: url::Url, - config: WgConfig, - udp: Option>, - - bind_addrs: Vec, - ip_version: IpVersion, - resolved_addr: Option, - socket_mark: Option, -} - -impl Debug for WgTunnelConnector { - fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result { - f.debug_struct("WgTunnelConnector") - .field("addr", &self.addr) - .field("udp", &self.udp) - .finish() - } -} - -impl WgTunnelConnector { - pub fn new(addr: url::Url, config: WgConfig) -> Self { - WgTunnelConnector { - addr, - config, - udp: None, - bind_addrs: vec![], - ip_version: IpVersion::Both, - resolved_addr: None, - socket_mark: None, - } - } - - #[tracing::instrument(skip(config))] - async fn connect_with_socket( - addr_url: url::Url, - config: WgConfig, - udp: UdpSocket, - addr: SocketAddr, - ) -> Result, super::TunnelError> { - tracing::warn!("wg connect: {:?}", addr); - let local_addr = udp - .local_addr() - .with_context(|| "Failed to get local addr")? - .to_string(); - - let mut wg_peer = WgPeer::new(Arc::new(udp), config.clone(), addr); - let udp = wg_peer.udp_socket(); - - // do handshake here so we will return after receive first packet - let handshake = wg_peer.create_handshake_init().await; - udp.send_to(&handshake, addr).await?; - let mut buf = [0u8; MAX_PACKET]; - let (n, recv_addr) = match udp.recv_from(&mut buf).await { - Ok(ret) => ret, - Err(e) => { - tracing::error!("Failed to receive handshake response: {}", e); - return Err(TunnelError::IOError(e)); - } - }; - - if recv_addr != addr { - tracing::warn!(?recv_addr, "Received packet from changed address"); - } - - let tunnel = wg_peer.start_and_get_tunnel(); - let data = wg_peer.data.as_ref().unwrap().clone(); - let mut sink = wg_peer.sink.lock().unwrap().take().unwrap(); - wg_peer.tasks.spawn(async move { - data.handle_one_packet_from_peer(&mut sink, &buf[..n]).await; - loop { - let mut buf = vec![0u8; MAX_PACKET]; - let (n, _) = match udp.recv_from(&mut buf).await { - Ok(ret) => ret, - Err(e) => { - tracing::error!("Failed to receive wg packet: {}", e); - break; - } - }; - data.handle_one_packet_from_peer(&mut sink, &buf[..n]).await; - } - }); - - let (stream, sink) = tunnel.split(); - let ret = Box::new(TunnelWrapper::new_with_associate_data( - stream, - sink, - Some(TunnelInfo { - tunnel_type: "wg".to_owned(), - local_addr: Some(super::build_url_from_socket_addr(&local_addr, "wg").into()), - remote_addr: Some(addr_url.into()), - resolved_remote_addr: Some( - super::build_url_from_socket_addr(&addr.to_string(), "wg").into(), - ), - }), - Some(Box::new(wg_peer)), - )); - - Ok(ret) - } - - async fn connect_with_ipv6(&self, addr: SocketAddr) -> Result, TunnelError> { - let socket = bind() - .addr("[::]:0".parse().unwrap()) - .dev(BindDev::Disabled) - .only_v6(true) - .maybe_socket_mark(self.socket_mark) - .call()?; - Self::connect_with_socket(self.addr.clone(), self.config.clone(), socket, addr).await - } } #[async_trait] -impl super::TunnelConnector for WgTunnelConnector { - #[tracing::instrument] - async fn connect(&mut self) -> Result, TunnelError> { - let addr = match self.resolved_addr { - Some(addr) => addr, - None => SocketAddr::from_url(self.addr.clone(), self.ip_version).await?, - }; +impl easytier_core::socket::SocketListener for WgTunnelListener { + type Accepted = Box; - if addr.is_ipv6() { - return self.connect_with_ipv6(addr).await; + async fn listen(&mut self) -> anyhow::Result<()> { + Ok(self.listen_tunnel().await?) + } + + async fn accept(&mut self) -> anyhow::Result { + Ok(self.accept_tunnel().await?) + } + + fn local_url(&self) -> url::Url { + self.session_listener + .as_ref() + .map(|listener| easytier_core::socket::SocketListener::local_url(listener.as_ref())) + .unwrap_or_else(|| self.addr.clone()) + } +} + +pub(crate) async fn upgrade_connected( + connected: ConnectedUdpSession, + addr_url: url::Url, + config: WgConfig, +) -> Result, TunnelError> { + let (session, session_guard) = connected.into_parts(); + let session = Arc::new(session) as Arc; + let addr = session.peer_addr()?; + let local_addr = session + .local_addr() + .with_context(|| "Failed to get local addr")? + .to_string(); + + let mut wg_peer = WgPeer::new(session_guard, session.clone(), config.clone(), addr); + + // do handshake here so we will return after receive first packet + let handshake = wg_peer.create_handshake_init().await; + session.send(&handshake).await?; + let mut buf = [0u8; MAX_PACKET]; + let n = match session.recv(&mut buf).await { + Ok(ret) => ret, + Err(e) => { + tracing::error!("Failed to receive handshake response: {}", e); + return Err(TunnelError::IOError(e)); } + }; - let bind_addrs = if self.bind_addrs.is_empty() { - vec!["0.0.0.0:0".parse().unwrap()] - } else { - self.bind_addrs.clone() - }; - let futures = FuturesUnordered::new(); - for bind_addr in bind_addrs.into_iter() { - tracing::info!(?bind_addr, ?addr, "bind addr"); - match bind() - .addr(bind_addr) - .only_v6(true) - .maybe_socket_mark(self.socket_mark) - .call() - { - Ok(socket) => futures.push(Self::connect_with_socket( - self.addr.clone(), - self.config.clone(), - socket, - addr, - )), - Err(error) => { - tracing::error!(?error, ?bind_addr, ?addr, "bind addr fail"); - continue; - } - } - } + let tunnel = wg_peer.start_and_get_tunnel(); + wg_peer.spawn_session_recv_task(Some(buf[..n].to_vec())); - wait_for_connect_futures(futures).await - } + let (stream, sink) = tunnel.split(); + let ret = Box::new(TunnelWrapper::new_with_associate_data( + stream, + sink, + Some(TunnelInfo { + tunnel_type: "wg".to_owned(), + local_addr: Some(super::build_url_from_socket_addr(&local_addr, "wg").into()), + remote_addr: Some(addr_url.into()), + resolved_remote_addr: Some( + super::build_url_from_socket_addr(&addr.to_string(), "wg").into(), + ), + }), + Some(Box::new(wg_peer)), + )); - fn remote_url(&self) -> url::Url { - self.addr.clone() - } + Ok(ret) +} - fn set_bind_addrs(&mut self, addrs: Vec) { - self.bind_addrs = addrs; - } +pub(crate) fn upgrade_accepted( + session: UdpSession, + config: WgConfig, +) -> Result, TunnelError> { + let session = Arc::new(session) as Arc; + let remote_addr = session.peer_addr()?; + let local_addr = session.local_addr()?; + let mut wg_peer = WgPeer::new(Box::new(()), session, config, remote_addr); + let tunnel = wg_peer.start_and_get_tunnel(); + wg_peer.spawn_session_recv_task(None); - fn set_ip_version(&mut self, ip_version: IpVersion) { - self.ip_version = ip_version; - } - - fn set_resolved_addr(&mut self, addr: SocketAddr) { - self.resolved_addr = Some(addr); - } - - fn set_socket_mark(&mut self, socket_mark: Option) { - self.socket_mark = socket_mark; - } + let (stream, sink) = tunnel.split(); + let remote_url = build_url_from_socket_addr(&remote_addr.to_string(), "wg"); + Ok(Box::new(TunnelWrapper::new_with_associate_data( + stream, + sink, + Some(TunnelInfo { + tunnel_type: "wg".to_owned(), + local_addr: Some(build_url_from_socket_addr(&local_addr.to_string(), "wg").into()), + remote_addr: Some(remote_url.clone().into()), + resolved_remote_addr: Some(remote_url.into()), + }), + Some(Box::new(wg_peer)), + ))) } #[cfg(test)] pub mod tests { use super::*; - use crate::tunnel::{ - TunnelConnector, - common::tests::{_tunnel_bench, _tunnel_pingpong}, + use crate::{ + common::global_ctx::tests::get_mock_global_ctx, host_runtime::native_host_runtime, + tunnel::protocol::runtime_client_protocol_upgrader, + }; + use easytier_core::{ + connectivity::transport::{ConnectedTransport, UdpSessionMode, connect_udp}, + socket::SocketListener, + socket::udp::{UdpBindOptions, UdpSessionProtocol}, }; - use boringtun::*; - pub fn create_wg_config() -> (WgConfig, WgConfig) { - let my_secret_key = x25519::StaticSecret::random_from_rng(rand::thread_rng()); - let my_public_key = x25519::PublicKey::from(&my_secret_key); - - let their_secret_key = x25519::StaticSecret::random_from_rng(rand::thread_rng()); - let their_public_key = x25519::PublicKey::from(&their_secret_key); - - let server_cfg = WgConfig { - my_secret_key: my_secret_key.clone(), - my_public_key, - peer_secret_key: their_secret_key.clone(), - peer_public_key: their_public_key, - wg_type: WgType::InternalUse, - }; - - let client_cfg = WgConfig { - my_secret_key: their_secret_key, - my_public_key: their_public_key, - peer_secret_key: my_secret_key, - peer_public_key: my_public_key, - wg_type: WgType::InternalUse, - }; - - (server_cfg, client_cfg) - } - - #[tokio::test] - async fn wg_pingpong() { - let (server_cfg, client_cfg) = create_wg_config(); - let listener = WgTunnelListener::new("wg://0.0.0.0:5599".parse().unwrap(), server_cfg); - let connector = WgTunnelConnector::new("wg://127.0.0.1:5599".parse().unwrap(), client_cfg); - _tunnel_pingpong(listener, connector).await - } - - #[tokio::test] - async fn wg_bench() { - let (server_cfg, client_cfg) = create_wg_config(); - let listener = WgTunnelListener::new("wg://0.0.0.0:5598".parse().unwrap(), server_cfg); - let connector = WgTunnelConnector::new("wg://127.0.0.1:5598".parse().unwrap(), client_cfg); - _tunnel_bench(listener, connector).await - } - - #[tokio::test] - async fn wg_bench_with_bind() { - let (server_cfg, client_cfg) = create_wg_config(); - let listener = WgTunnelListener::new("wg://127.0.0.1:5597".parse().unwrap(), server_cfg); - let mut connector = - WgTunnelConnector::new("wg://127.0.0.1:5597".parse().unwrap(), client_cfg); - connector.set_bind_addrs(vec!["127.0.0.1:0".parse().unwrap()]); - _tunnel_pingpong(listener, connector).await - } - - #[tokio::test] - #[should_panic] - async fn wg_bench_with_bind_fail() { - let (server_cfg, client_cfg) = create_wg_config(); - let listener = WgTunnelListener::new("wg://127.0.0.1:5596".parse().unwrap(), server_cfg); - let mut connector = - WgTunnelConnector::new("wg://127.0.0.1:5596".parse().unwrap(), client_cfg); - connector.set_bind_addrs(vec!["10.0.0.1:0".parse().unwrap()]); - _tunnel_pingpong(listener, connector).await + fn test_wg_config() -> WgConfig { + WgConfig::new_from_network_identity("test", "secret") } #[tokio::test] async fn wg_server_erase_from_map_after_close() { - let (server_cfg, client_cfg) = create_wg_config(); - let mut listener = - WgTunnelListener::new("wg://127.0.0.1:5595".parse().unwrap(), server_cfg); + let global_ctx = get_mock_global_ctx(); + let identity = global_ctx.get_network_identity(); + let server_cfg = WgConfig::new_from_network_identity( + &identity.network_name, + &identity.network_secret.unwrap_or_default(), + ); + let client = runtime_client_protocol_upgrader(global_ctx); + let mut listener = WgTunnelListener::new("wg://127.0.0.1:0".parse().unwrap(), server_cfg); listener.listen().await.unwrap(); + let remote_url = listener.local_url(); + let remote_addr = remote_url.socket_addrs(|| None).unwrap()[0]; const CONN_COUNT: usize = 10; - tokio::spawn(async move { - let mut tunnels = vec![]; + let client_task = tokio::spawn(async move { + let mut tunnels = Vec::with_capacity(CONN_COUNT); for _ in 0..CONN_COUNT { - let mut connector = WgTunnelConnector::new( - "wg://127.0.0.1:5595".parse().unwrap(), - client_cfg.clone(), - ); - let ret = connector.connect().await; - assert!(ret.is_ok()); - let t = ret.unwrap(); - let (_stream, mut sink) = t.split(); - sink.send(ZCPacket::new_with_payload("payload".as_bytes())) + let connected = connect_udp( + native_host_runtime(), + remote_addr, + Vec::new(), + UdpBindOptions::direct_connect(), + UdpSessionMode::Classified(UdpSessionProtocol::WireGuard), + ) + .await + .unwrap(); + let tunnel = client + .upgrade_client(ConnectedTransport::Udp(connected), remote_url.clone()) .await .unwrap(); - tunnels.push(t); + let (_stream, mut sink) = tunnel.split(); + sink.send(ZCPacket::new_with_payload(b"payload")) + .await + .unwrap(); + tunnels.push(tunnel); } - tokio::time::sleep(tokio::time::Duration::from_secs(1)).await; + tokio::time::sleep(Duration::from_secs(1)).await; }); for _ in 0..CONN_COUNT { - println!("accepting"); - let conn = listener.accept().await; - let (mut stream, _sink) = conn.unwrap().split(); + let tunnel = listener.accept().await.unwrap(); + let (mut stream, _sink) = tunnel.split(); let packet = stream.next().await.unwrap().unwrap(); - assert_eq!("payload".as_bytes(), packet.payload()); - println!("accepting drop"); + assert_eq!(packet.payload(), b"payload"); } - tokio::time::sleep(tokio::time::Duration::from_secs(2)).await; - - assert_eq!(0, listener.wg_peer_map.len()); + client_task.await.unwrap(); + tokio::time::sleep(Duration::from_secs(2)).await; + assert!(listener.wg_peer_map.is_empty()); } #[tokio::test] async fn bind_same_port() { - let (server_cfg, _client_cfg) = create_wg_config(); + let server_cfg = test_wg_config(); let mut listener = WgTunnelListener::new("wg://[::1]:31015".parse().unwrap(), server_cfg); - let (server_cfg, _client_cfg) = create_wg_config(); + let server_cfg = test_wg_config(); let mut listener2 = WgTunnelListener::new("wg://[::1]:31015".parse().unwrap(), server_cfg); listener.listen().await.unwrap(); listener2.listen().await.unwrap(); } - #[tokio::test] - async fn ipv6_pingpong() { - let (server_cfg, client_cfg) = create_wg_config(); - let listener = WgTunnelListener::new("wg://[::1]:31015".parse().unwrap(), server_cfg); - let connector = WgTunnelConnector::new("wg://[::1]:31015".parse().unwrap(), client_cfg); - _tunnel_pingpong(listener, connector).await - } - - #[tokio::test] - async fn ipv6_domain_pingpong() { - let (server_cfg, client_cfg) = create_wg_config(); - let listener = WgTunnelListener::new("wg://[::1]:31016".parse().unwrap(), server_cfg); - let mut connector = - WgTunnelConnector::new("wg://test.easytier.top:31016".parse().unwrap(), client_cfg); - connector.set_ip_version(IpVersion::V6); - _tunnel_pingpong(listener, connector).await; - - let (server_cfg, client_cfg) = create_wg_config(); - let listener = WgTunnelListener::new("wg://127.0.0.1:31016".parse().unwrap(), server_cfg); - let mut connector = - WgTunnelConnector::new("wg://test.easytier.top:31016".parse().unwrap(), client_cfg); - connector.set_ip_version(IpVersion::V4); - _tunnel_pingpong(listener, connector).await; - } - #[tokio::test] async fn test_alloc_port() { // v4 - let (server_cfg, _client_cfg) = create_wg_config(); + let server_cfg = test_wg_config(); let mut listener = WgTunnelListener::new("wg://0.0.0.0:0".parse().unwrap(), server_cfg); listener.listen().await.unwrap(); let port = listener.local_url().port().unwrap(); assert!(port > 0); // v6 - let (server_cfg, _client_cfg) = create_wg_config(); + let server_cfg = test_wg_config(); let mut listener = WgTunnelListener::new("wg://[::]:0".parse().unwrap(), server_cfg); listener.listen().await.unwrap(); let port = listener.local_url().port().unwrap(); diff --git a/easytier/src/utils/error.rs b/easytier/src/utils/error.rs deleted file mode 100644 index 75711d82..00000000 --- a/easytier/src/utils/error.rs +++ /dev/null @@ -1,58 +0,0 @@ -use delegate::delegate; -use derivative::Derivative; -use derive_more::{AsMut, AsRef, Deref, DerefMut, From, Into, IntoIterator}; -use std::fmt; -use std::fmt::Display; -use thiserror::Error; - -#[derive(Derivative, Debug, From, Into, Deref, DerefMut, AsRef, AsMut, IntoIterator, Error)] -#[derivative(Default(bound = ""))] -#[as_ref(forward)] -#[as_mut(forward)] -#[into_iterator(owned, ref, ref_mut)] -pub struct ErrorCollection { - pub errors: Vec, -} - -impl ErrorCollection { - delegate! { - to Vec { - #[into] - pub fn new() -> Self; - #[into] - pub fn with_capacity(capacity: usize) -> Self; - } - } -} - -impl> FromIterator for ErrorCollection { - fn from_iter>(iter: I) -> Self { - Self { - errors: iter.into_iter().map(Into::into).collect(), - } - } -} - -impl Extend for ErrorCollection { - delegate! { - to self.errors { - fn extend>(&mut self, iter: T); - } - } -} - -impl Display for ErrorCollection { - fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { - if self.errors.is_empty() { - return write!(f, "No errors"); - } - - write!(f, "{} error(s) occurred:", self.errors.len())?; - for (i, err) in self.errors.iter().enumerate() { - writeln!(f)?; - write!(f, " {}. {}", i + 1, err)?; - } - - Ok(()) - } -} diff --git a/easytier/src/utils/mod.rs b/easytier/src/utils/mod.rs index 280cd307..f9956096 100644 --- a/easytier/src/utils/mod.rs +++ b/easytier/src/utils/mod.rs @@ -1,11 +1,11 @@ -pub mod error; +#[cfg(feature = "management")] pub mod panic; pub mod string; -pub mod task; use std::net::{IpAddr, Ipv4Addr, SocketAddr, TcpListener}; use std::sync::{Arc, Weak}; +#[cfg(feature = "management")] pub type PeerRoutePair = crate::proto::api::instance::PeerRoutePair; pub fn check_tcp_available(port: u16) -> bool { diff --git a/easytier/src/utils/string.rs b/easytier/src/utils/string.rs index 8dab764f..d04ef68b 100644 --- a/easytier/src/utils/string.rs +++ b/easytier/src/utils/string.rs @@ -6,10 +6,6 @@ pub fn cost_to_str(cost: i32) -> String { } } -pub fn float_to_str(f: f64, precision: usize) -> String { - format!("{:.1$}", f, precision) -} - #[cfg(target_os = "windows")] pub fn utf8_or_gbk_to_string(s: &[u8]) -> String { use encoding::{DecoderTrap, Encoding, all::GBK}; diff --git a/easytier/src/utils/task.rs b/easytier/src/utils/task.rs deleted file mode 100644 index ce34df56..00000000 --- a/easytier/src/utils/task.rs +++ /dev/null @@ -1,142 +0,0 @@ -use crate::utils::error::ErrorCollection; -use futures::StreamExt; -use futures::stream::FuturesUnordered; -use std::future::Future; -use std::io; -use std::pin::Pin; -use std::task::{Context, Poll}; -use std::time::Duration; -use tokio::task::JoinHandle; -use tokio::time::sleep; -use tokio_util::sync::CancellationToken; -use tokio_util::task::AbortOnDropHandle; - -// region CancellableTask - -#[derive(Debug)] -pub struct CancellableTask { - handle: AbortOnDropHandle, - token: CancellationToken, -} - -impl CancellableTask { - pub fn token(&self) -> &CancellationToken { - &self.token - } - - pub fn with_handle(token: CancellationToken, handle: JoinHandle) -> Self { - Self { - handle: AbortOnDropHandle::new(handle), - token, - } - } - - pub async fn stop(mut self, timeout: Option) -> io::Result { - self.token.cancel(); - - match timeout { - Some(timeout) => tokio::time::timeout(timeout, &mut self.handle) - .await - .map_err(|e| { - tracing::warn!("task stop timeout after {:?}, aborted", timeout); - io::Error::new(io::ErrorKind::TimedOut, e) - })?, - None => self.handle.await, - } - .map_err(Into::into) - } -} - -impl CancellableTask { - pub fn new(token: CancellationToken, future: F) -> Self - where - F: Future + Send + 'static, - { - Self::with_handle(token, tokio::spawn(future)) - } - - pub fn spawn(factory: impl FnOnce(CancellationToken) -> F) -> Self - where - F: Future + Send + 'static, - { - let token = CancellationToken::new(); - Self::new(token.clone(), factory(token)) - } - - pub fn child(&self, factory: impl FnOnce(CancellationToken) -> F) -> Self - where - F: Future + Send + 'static, - { - let token = self.token.clone(); - Self::new(token.clone(), factory(token)) - } -} - -impl Future for CancellableTask { - type Output = io::Result; - fn poll(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll { - Pin::new(&mut self.handle) - .poll(cx) - .map(|result| result.map_err(Into::into)) - } -} - -// endregion - -// region HedgeExt - -pub(crate) trait HedgeExt: Iterator + Sized { - async fn hedge(self, delay: Duration) -> Result> - where - Self::Item: Future>; -} - -impl HedgeExt for I -where - I: Iterator, -{ - async fn hedge(mut self, delay: Duration) -> Result> - where - Self::Item: Future>, - { - let mut tasks = FuturesUnordered::new(); - let mut errors = ErrorCollection::new(); - let mut exhausted = false; - - macro_rules! spawn { - () => { - if let Some(fut) = self.next() { - tasks.push(fut); - } else { - exhausted = true; - } - }; - } - - spawn!(); - - while !tasks.is_empty() { - tokio::select! { - res = tasks.next() => { - match res { - Some(Ok(v)) => return Ok(v), - Some(Err(e)) => errors.push(e), - None => unreachable!(), - } - - if !exhausted { - spawn!(); - } - } - - _ = sleep(delay), if !exhausted => { - spawn!(); - } - } - } - - Err(errors) - } -} - -// endregion diff --git a/easytier/src/vpn_portal/mod.rs b/easytier/src/vpn_portal/mod.rs index 7a6caa85..afaec496 100644 --- a/easytier/src/vpn_portal/mod.rs +++ b/easytier/src/vpn_portal/mod.rs @@ -1,50 +1,2 @@ -// with vpn portal, user can use other vpn client to connect to easytier servers -// without installing easytier. -// these vpn client include: -// 1. wireguard -// 2. openvpn (TODO) -// 3. shadowsocks (TODO) - -use std::sync::Arc; - -use crate::{common::global_ctx::ArcGlobalCtx, peers::peer_manager::PeerManager}; - #[cfg(feature = "wireguard")] pub mod wireguard; - -#[async_trait::async_trait] -pub trait VpnPortal: Send + Sync { - async fn start( - &mut self, - global_ctx: ArcGlobalCtx, - peer_mgr: Arc, - ) -> anyhow::Result<()>; - async fn dump_client_config(&self, peer_mgr: Arc) -> String; - fn name(&self) -> String; - async fn list_clients(&self) -> Vec; -} - -pub struct NullVpnPortal; - -#[async_trait::async_trait] -impl VpnPortal for NullVpnPortal { - async fn start( - &mut self, - _global_ctx: ArcGlobalCtx, - _peer_mgr: Arc, - ) -> anyhow::Result<()> { - Ok(()) - } - - async fn dump_client_config(&self, _peer_mgr: Arc) -> String { - "".to_string() - } - - fn name(&self) -> String { - "null".to_string() - } - - async fn list_clients(&self) -> Vec { - vec![] - } -} diff --git a/easytier/src/vpn_portal/wireguard.rs b/easytier/src/vpn_portal/wireguard.rs index ab85ff93..629f5500 100644 --- a/easytier/src/vpn_portal/wireguard.rs +++ b/easytier/src/vpn_portal/wireguard.rs @@ -1,339 +1,92 @@ use std::{ - net::{IpAddr, Ipv4Addr, Ipv6Addr, SocketAddr, SocketAddrV6}, + net::{Ipv6Addr, SocketAddr, SocketAddrV6}, sync::Arc, }; use anyhow::Context; use base64::{Engine, prelude::BASE64_STANDARD}; -use cidr::Ipv4Inet; -use dashmap::DashMap; -use futures::StreamExt; -use pnet::packet::ipv4::Ipv4Packet; -use tokio::task::JoinSet; -use tracing::Level; - -use crate::{ - common::{ - config::NetworkIdentity, - global_ctx::{ArcGlobalCtx, GlobalCtxEvent}, - join_joinset_background, shrink_dashmap, - }, - peers::{PeerPacketFilter, peer_manager::PeerManager}, - tunnel::{ - Tunnel, TunnelListener, - mpsc::{MpscTunnel, MpscTunnelSender}, - packet_def::{PacketType, ZCPacket, ZCPacketType}, - wireguard::{WgConfig, WgTunnelListener}, - }, +use easytier_core::{ + gateway::vpn_portal::{VpnPortalClientConfigPlan, VpnPortalHost, VpnPortalListener}, + socket::SocketListener, }; -use super::VpnPortal; - -type WgPeerIpTable = Arc>>; +use crate::{ + common::{config::NetworkIdentity, global_ctx::ArcGlobalCtx}, + tunnel::wireguard::{WgConfig, WgTunnelListener}, +}; pub(crate) fn get_wg_config_for_portal(nid: &NetworkIdentity) -> WgConfig { let key_seed = format!( "{}{}", nid.network_name, - nid.network_secret.as_ref().unwrap_or(&"".to_string()) + nid.network_secret.as_ref().unwrap_or(&String::new()) ); WgConfig::new_for_portal(&key_seed, &key_seed) } -struct ClientEntry { - endpoint_addr: Option, - sink: MpscTunnelSender, +fn listener_endpoint(listener_url: &url::Url) -> &str { + &listener_url[url::Position::BeforeHost..url::Position::AfterPort] } -struct WireGuardImpl { +pub struct WireGuardPortalHost { global_ctx: ArcGlobalCtx, - peer_mgr: Arc, wg_config: WgConfig, - listener_addr: SocketAddr, - - wg_peer_ip_table: WgPeerIpTable, - - tasks: Arc>>, + listener_addr: Option, } -impl WireGuardImpl { - fn new(global_ctx: ArcGlobalCtx, peer_mgr: Arc) -> Self { - let nid = global_ctx.get_network_identity(); - let wg_config = get_wg_config_for_portal(&nid); - - let vpn_cfg = global_ctx.config.get_vpn_portal_config().unwrap(); - let listener_addr = vpn_cfg.wireguard_listen; - - Self { +impl WireGuardPortalHost { + pub fn new(global_ctx: ArcGlobalCtx, listener_addr: Option) -> Arc { + Arc::new(Self { + wg_config: get_wg_config_for_portal(&global_ctx.get_network_identity()), global_ctx, - peer_mgr, - wg_config, listener_addr, - wg_peer_ip_table: Arc::new(DashMap::new()), - tasks: Arc::new(std::sync::Mutex::new(JoinSet::new())), - } + }) } - async fn handle_incoming_conn( - t: Box, - peer_mgr: Arc, - wg_peer_ip_table: WgPeerIpTable, - ) { - let info = t.info().unwrap_or_default(); - let mut mpsc_tunnel = MpscTunnel::new(t, None); - let mut stream = mpsc_tunnel.get_stream(); - let mut ip_registered = false; - - let remote_addr = info.remote_addr.clone(); - let endpoint_addr = remote_addr.clone().map(Into::into); - peer_mgr - .get_global_ctx() - .issue_event(GlobalCtxEvent::VpnPortalClientConnected( - info.local_addr.clone().unwrap_or_default().to_string(), - info.remote_addr.clone().unwrap_or_default().to_string(), - )); - - let mut map_key = None; - - loop { - let msg = match stream.next().await { - Some(Ok(msg)) => msg, - Some(Err(err)) => { - tracing::error!(?err, "Failed to receive from wg client"); - break; - } - None => { - tracing::info!("Wireguard client disconnected"); - break; - } - }; - - assert_eq!(msg.packet_type(), ZCPacketType::WG); - let inner = msg.inner(); - let Some(i) = Ipv4Packet::new(&inner) else { - tracing::error!(?inner, "Failed to parse ipv4 packet"); - continue; - }; - if !ip_registered { - let client_entry = Arc::new(ClientEntry { - endpoint_addr: endpoint_addr.clone(), - sink: mpsc_tunnel.get_sink(), - }); - map_key = Some(i.get_source()); - // Be careful here: we may overwrite an existing entry if the client IP is reused, - // which is common when clients are behind NAT. - wg_peer_ip_table.insert(i.get_source(), client_entry.clone()); - ip_registered = true; - } - tracing::trace!(?i, "Received from wg client"); - let dst = i.get_destination(); - let _ = peer_mgr - .send_msg_by_ip( - ZCPacket::new_with_payload(inner.as_ref()), - IpAddr::V4(dst), - false, - ) - .await; - } - - if let Some(map_key) = map_key { - // Remove the client from the wg_peer_ip_table only when its endpoint address is unchanged, - // or we may break clients behind NAT. - match wg_peer_ip_table - .remove_if(&map_key, |_, entry| entry.endpoint_addr == endpoint_addr) - { - Some(_) => tracing::info!(?map_key, "Removed wg client from table"), - None => tracing::info!( - ?map_key, - "The wg client changed its endpoint address, not removing from table" - ), - } - shrink_dashmap(&wg_peer_ip_table, None); - } - - peer_mgr - .get_global_ctx() - .issue_event(GlobalCtxEvent::VpnPortalClientDisconnected( - info.local_addr.unwrap_or_default().to_string(), - info.remote_addr.unwrap_or_default().to_string(), - )); - } - - async fn start_pipeline_processor(&self) { - struct PeerPacketFilterForVpnPortal { - wg_peer_ip_table: WgPeerIpTable, - } - - #[async_trait::async_trait] - impl PeerPacketFilter for PeerPacketFilterForVpnPortal { - async fn try_process_packet_from_peer(&self, packet: ZCPacket) -> Option { - let hdr = packet.peer_manager_header().unwrap(); - if hdr.packet_type != PacketType::Data as u8 { - return Some(packet); - }; - - let payload_bytes = packet.payload(); - let ipv4 = Ipv4Packet::new(payload_bytes)?; - if ipv4.get_version() != 4 { - return Some(packet); - } - - let Some(entry) = self - .wg_peer_ip_table - .get(&ipv4.get_destination()) - .map(|f| f.clone()) - else { - return Some(packet); - }; - - tracing::trace!(?ipv4, "Packet filter for vpn portal"); - - let payload_offset = packet.packet_type().get_packet_offsets().payload_offset; - let packet = ZCPacket::new_from_buf( - packet.inner().split_off(payload_offset), - ZCPacketType::WG, - ); - - match entry.sink.try_send(packet) { - Ok(_) => { - tracing::trace!("Sent packet to wg client"); - } - Err(e) => { - tracing::debug!(?e, "Failed to send packet to wg client"); - } - } - - None - } - } - - self.peer_mgr - .add_packet_process_pipeline(Box::new(PeerPacketFilterForVpnPortal { - wg_peer_ip_table: self.wg_peer_ip_table.clone(), - })) - .await; - } - - async fn start_listener(&self, listener_addr: &SocketAddr) -> anyhow::Result<()> { + async fn start_listener(&self, listener_addr: SocketAddr) -> anyhow::Result { let mut listener_url = url::Url::parse("wg://0.0.0.0:0").unwrap(); listener_url.set_port(Some(listener_addr.port())).unwrap(); listener_url.set_ip_host(listener_addr.ip()).unwrap(); - let mut l = WgTunnelListener::new(listener_url.clone(), self.wg_config.clone()); - - tracing::info!("Wireguard VPN Portal Starting"); - + let mut listener = WgTunnelListener::new(listener_url, self.wg_config.clone()); { - let _g = self.global_ctx.net_ns.guard(); - l.listen() + let _guard = self.global_ctx.net_ns.guard(); + listener + .listen() .await - .with_context(|| "Failed to start wireguard listener for vpn portal")?; + .context("failed to start WireGuard VPN portal listener")?; } - let tasks = Arc::downgrade(&self.tasks.clone()); - let peer_mgr = self.peer_mgr.clone(); - let wg_peer_ip_table = self.wg_peer_ip_table.clone(); - self.tasks.lock().unwrap().spawn(async move { - while let Ok(t) = l.accept().await { - let Some(tasks) = tasks.upgrade() else { - break; - }; - tasks.lock().unwrap().spawn(Self::handle_incoming_conn( - t, - peer_mgr.clone(), - wg_peer_ip_table.clone(), - )); - } - }); - - self.global_ctx - .issue_event(GlobalCtxEvent::VpnPortalStarted(listener_url.to_string())); - - Ok(()) + Ok(Box::new(listener)) } +} - #[tracing::instrument(skip(self), err(level = Level::WARN))] - async fn start(&self) -> anyhow::Result<()> { - tracing::info!("Wireguard VPN Portal Starting"); - - self.start_listener(&self.listener_addr).await?; - // if binding to v4 unspecified, also start a listener on v6 unspecified - if let SocketAddr::V4(v4) = &self.listener_addr +#[async_trait::async_trait] +impl VpnPortalHost for WireGuardPortalHost { + async fn start_listeners(&self) -> anyhow::Result> { + let listener_addr = self.listener_addr.context("VPN portal config is not set")?; + let mut listeners = vec![self.start_listener(listener_addr).await?]; + if let SocketAddr::V4(v4) = listener_addr && v4.ip().is_unspecified() - { - let _ = self - .start_listener(&SocketAddr::V6(SocketAddrV6::new( + && let Ok(listener) = self + .start_listener(SocketAddr::V6(SocketAddrV6::new( Ipv6Addr::UNSPECIFIED, v4.port(), 0, 0, ))) - .await; - }; - - join_joinset_background(self.tasks.clone(), "wireguard".to_string()); - self.start_pipeline_processor().await; - - Ok(()) - } -} - -#[derive(Default)] -pub struct WireGuard { - inner: Option, -} - -#[async_trait::async_trait] -impl VpnPortal for WireGuard { - async fn start( - &mut self, - global_ctx: ArcGlobalCtx, - peer_mgr: Arc, - ) -> anyhow::Result<()> { - assert!(self.inner.is_none()); - - let vpn_cfg = global_ctx.config.get_vpn_portal_config(); - if vpn_cfg.is_none() { - anyhow::bail!("vpn cfg is not set for wireguard vpn portal"); - } - - let inner = WireGuardImpl::new(global_ctx, peer_mgr); - inner.start().await?; - self.inner = Some(inner); - Ok(()) - } - - async fn dump_client_config(&self, peer_mgr: Arc) -> String { - if self.inner.is_none() { - return "ERROR: Wireguard VPN Portal Not Started".to_string(); - } - let global_ctx = self.inner.as_ref().unwrap().global_ctx.clone(); - if global_ctx.config.get_vpn_portal_config().is_none() { - return "ERROR: VPN Portal Config Not Set".to_string(); - } - - let routes = peer_mgr.list_routes().await; - let mut allow_ips = routes - .iter() - .flat_map(|x| x.proxy_cidrs.iter().map(String::to_string)) - .collect::>(); - if let Some(ipv4) = routes - .iter() - .filter_map(|x| x.ipv4_addr) - .chain(global_ctx.get_ipv4().into_iter().map(Into::into)) - .next() + .await { - let inet = Ipv4Inet::from(ipv4); - allow_ips.push(inet.network().to_string()); + listeners.push(listener); } + Ok(listeners) + } - let vpn_cfg = global_ctx.config.get_vpn_portal_config().unwrap(); - let client_cidr = vpn_cfg.client_cidr; + fn name(&self) -> String { + "wireguard".to_owned() + } - allow_ips.push(client_cidr.to_string()); - - let allow_ips = allow_ips.into_iter().collect::>().join(","); - - let cfg = self.inner.as_ref().unwrap().wg_config.clone(); - let cfg_str = format!( + fn render_client_config(&self, plan: &VpnPortalClientConfigPlan) -> String { + let listener_addr = listener_endpoint(&plan.listener_url); + format!( r#" [Interface] PrivateKey = {peer_secret_key} @@ -341,39 +94,36 @@ Address = {address} # should assign an ip from this cidr manually [Peer] PublicKey = {my_public_key} -AllowedIPs = {allow_ips} -Endpoint = {listenr_addr} # should be the public ip(or domain) of the vpn server +AllowedIPs = {allowed_ips} +Endpoint = {listener_addr} # should be the public ip(or domain) of the vpn server PersistentKeepalive = 25 "#, - peer_secret_key = BASE64_STANDARD.encode(cfg.peer_secret_key()), - my_public_key = BASE64_STANDARD.encode(cfg.my_public_key()), - listenr_addr = self.inner.as_ref().unwrap().listener_addr, - allow_ips = allow_ips, - address = client_cidr.first_address().to_string() + "/32", - ); - - cfg_str + peer_secret_key = BASE64_STANDARD.encode(self.wg_config.peer_secret_key()), + my_public_key = BASE64_STANDARD.encode(self.wg_config.my_public_key()), + listener_addr = listener_addr, + allowed_ips = plan.allowed_ips.join(","), + address = plan.client_cidr.first_address().to_string() + "/32", + ) } - fn name(&self) -> String { - "wireguard".to_string() - } - - async fn list_clients(&self) -> Vec { - self.inner - .as_ref() - .map(|w| { - w.wg_peer_ip_table - .iter() - .map(|x| { - x.value() - .endpoint_addr - .as_ref() - .map(|x| x.to_string()) - .unwrap_or_default() - }) - .collect() - }) - .unwrap_or_default() + fn not_started_client_config(&self) -> String { + "ERROR: Wireguard VPN Portal Not Started".to_owned() + } +} + +#[cfg(test)] +mod tests { + use super::listener_endpoint; + + #[test] + fn listener_endpoint_uses_the_active_listener_url() { + assert_eq!( + listener_endpoint(&"wg://192.0.2.10:51820".parse().unwrap()), + "192.0.2.10:51820" + ); + assert_eq!( + listener_endpoint(&"wg://[2001:db8::10]:51820".parse().unwrap()), + "[2001:db8::10]:51820" + ); } } diff --git a/easytier/src/web_client/controller.rs b/easytier/src/web_client/controller.rs deleted file mode 100644 index 7e7e4cfa..00000000 --- a/easytier/src/web_client/controller.rs +++ /dev/null @@ -1,65 +0,0 @@ -use std::sync::Arc; - -use crate::{ - instance_manager::NetworkInstanceManager, - proto::{rpc_impl::service_registry::ServiceRegistry, web::DeviceOsInfo}, - rpc_service::api::register_api_rpc_service, - web_client::WebClientHooks, -}; - -pub struct Controller { - token: String, - machine_id: uuid::Uuid, - hostname: String, - device_os: DeviceOsInfo, - manager: Arc, - hooks: Arc, -} - -impl Controller { - pub fn new( - token: String, - machine_id: uuid::Uuid, - hostname: String, - device_os: DeviceOsInfo, - manager: Arc, - hooks: Arc, - ) -> Self { - Controller { - token, - machine_id, - hostname, - device_os, - manager, - hooks, - } - } - - pub fn list_network_instance_ids(&self) -> Vec { - self.manager.list_network_instance_ids() - } - - pub fn token(&self) -> String { - self.token.clone() - } - - pub fn hostname(&self) -> String { - self.hostname.clone() - } - - pub fn machine_id(&self) -> uuid::Uuid { - self.machine_id - } - - pub fn device_os(&self) -> DeviceOsInfo { - self.device_os.clone() - } - - pub fn register_api_rpc_service(&self, registry: &ServiceRegistry) { - register_api_rpc_service(&self.manager, registry, Some(self.hooks.clone())); - } - - pub(super) fn notify_manager_stopping(&self) { - self.manager.notify_stop_check(); - } -} diff --git a/easytier/src/web_client/mod.rs b/easytier/src/web_client/mod.rs index 926150af..02699a2a 100644 --- a/easytier/src/web_client/mod.rs +++ b/easytier/src/web_client/mod.rs @@ -1,42 +1,72 @@ use std::sync::Arc; +use anyhow::{Context as _, Result}; +use async_trait::async_trait; +use easytier_core::{ + connectivity::{manual::ManualTunnelConnector, protocol::raw::TunnelDialer}, + management::{ConfigServerEndpoint, WebClientConfig}, + socket::IpVersion, + tunnel::Tunnel, +}; +use url::Url; + use crate::{ common::{ - MachineIdOptions, - config::TomlConfigLoader, - global_ctx::{ArcGlobalCtx, GlobalCtx}, - log, - os_info::collect_device_os_info, - resolve_machine_id, - stun::MockStunInfoCollector, + MachineIdOptions, config::TomlConfigLoader, constants::EASYTIER_VERSION, + global_ctx::GlobalCtx, os_info::collect_device_os_info, resolve_machine_id, }, - connector::create_connector_by_url, - instance_manager::{DaemonGuard, NetworkInstanceManager}, - proto::common::NatType, - tunnel::{IpVersion, Tunnel, TunnelConnector, TunnelError, TunnelScheme}, + instance::{ + composition::runtime_one_shot_manual_connector, + config_storage::NativeConfigFileStorage, + factory::{NativeInstanceFactory, NativeInstanceManager}, + host::NativeInstanceHost, + }, + rpc_service::logger::NativeLoggerControl, + tunnel::TunnelScheme, }; -use anyhow::{Context as _, Result}; -use async_trait::async_trait; -use tokio_util::task::AbortOnDropHandle; -use url::Url; -use uuid::Uuid; -#[async_trait] -pub trait WebClientHooks: Send + Sync { - fn manages_remote_config_instances(&self) -> bool { - false +pub use easytier_core::management::InstanceMutationHooks as WebClientHooks; + +pub struct WebClient { + inner: easytier_core::management::WebClient, +} + +impl WebClient { + pub fn new( + connector: T, + token: S, + machine_id: uuid::Uuid, + hostname: H, + secure_mode: bool, + manager: Arc, + hooks: Option>, + ) -> Self + where + T: TunnelDialer + 'static, + S: ToString, + H: ToString, + { + Self { + inner: easytier_core::management::WebClient::new( + connector, + WebClientConfig { + token: token.to_string(), + machine_id, + hostname: hostname.to_string(), + device_os: collect_device_os_info(), + easytier_version: EASYTIER_VERSION.to_owned(), + secure_mode, + }, + manager, + hooks.unwrap_or_else(|| Arc::new(DefaultHooks)), + Arc::new(NativeConfigFileStorage), + Arc::new(NativeLoggerControl), + ), + } } - async fn pre_run_network_instance(&self, _cfg: &TomlConfigLoader) -> Result<(), String> { - Ok(()) - } - - async fn post_run_network_instance(&self, _id: &Uuid) -> Result<(), String> { - Ok(()) - } - - async fn post_remove_network_instances(&self, _ids: &[Uuid]) -> Result<(), String> { - Ok(()) + pub fn is_connected(&self) -> bool { + self.inner.is_connected() } } @@ -45,36 +75,17 @@ pub struct DefaultHooks; #[async_trait] impl WebClientHooks for DefaultHooks {} -pub mod controller; -pub mod security; -pub mod session; - -use std::sync::atomic::{AtomicBool, Ordering}; - -pub struct WebClient { - controller: Arc, - tasks: AbortOnDropHandle<()>, - manager_guard: DaemonGuard, - connected: Arc, -} - struct ConfigServerConnector { url: Url, - global_ctx: ArcGlobalCtx, + connector: ManualTunnelConnector, } #[async_trait] -impl TunnelConnector for ConfigServerConnector { - async fn connect(&mut self) -> std::result::Result, TunnelError> { - let mut connector = - create_connector_by_url(self.url.as_str(), &self.global_ctx, IpVersion::Both) - .await - .map_err(|err| match err { - crate::common::error::Error::TunnelError(err) => err, - err => TunnelError::Anyhow(err.into()), - })?; - - connector.connect().await +impl TunnelDialer for ConfigServerConnector { + async fn connect(&self) -> anyhow::Result> { + self.connector + .connect(self.url.clone(), IpVersion::Both) + .await } fn remote_url(&self) -> Url { @@ -82,221 +93,38 @@ impl TunnelConnector for ConfigServerConnector { } } -impl WebClient { - pub fn new( - connector: T, - token: S, - machine_id: Uuid, - hostname: H, - secure_mode: bool, - manager: Arc, - hooks: Option>, - ) -> Self { - let manager_guard = manager.register_daemon(); - let hooks = hooks.unwrap_or_else(|| Arc::new(DefaultHooks)); - let controller = Arc::new(controller::Controller::new( - token.to_string(), - machine_id, - hostname.to_string(), - collect_device_os_info(), - manager, - hooks, - )); - let connected = Arc::new(AtomicBool::new(false)); - - let controller_clone = controller.clone(); - let connected_clone = connected.clone(); - let tasks = AbortOnDropHandle::new(tokio::spawn(async move { - Self::routine( - controller_clone, - connected_clone, - secure_mode, - Box::new(connector), - ) - .await; - })); - - WebClient { - controller, - tasks, - manager_guard, - connected, - } - } - - async fn routine( - controller: Arc, - connected: Arc, - secure_mode: bool, - mut connector: Box, - ) { - loop { - let conn = match connector.connect().await { - Ok(conn) => conn, - Err(error) => { - let wait = 1; - log::warn!(%error, "Failed to connect to the server, retrying in {} seconds...", wait); - tokio::time::sleep(std::time::Duration::from_secs(wait)).await; - continue; - } - }; - - connected.store(true, Ordering::Release); - log::info!("Successfully connected to {:?}", conn.info()); - - let mut session = session::Session::new(conn, controller.clone()); - let support_encryption = match tokio::time::timeout( - std::time::Duration::from_secs(3), - session.get_feature(), - ) - .await - { - Ok(Ok(feature)) => feature.support_encryption, - Ok(Err(error)) => { - log::warn!(%error, "GetFeature rpc failed, fallback to legacy tunnel"); - false - } - Err(_) => { - log::warn!("GetFeature rpc timeout, fallback to legacy tunnel"); - false - } - }; - - if support_encryption && security::web_secure_tunnel_supported() { - log::info!("Server supports encryption, reconnecting with secure tunnel"); - drop(session); - - let conn = match connector.connect().await { - Ok(conn) => conn, - Err(error) => { - connected.store(false, Ordering::Release); - let wait = 1; - log::warn!(%error, "Failed to reconnect secure tunnel, retrying in {} seconds...", wait); - tokio::time::sleep(std::time::Duration::from_secs(wait)).await; - continue; - } - }; - - let conn = match security::upgrade_client_tunnel(conn).await { - Ok(conn) => conn, - Err(error) => { - connected.store(false, Ordering::Release); - let wait = 1; - log::warn!(%error, "Noise handshake failed, retrying in {} seconds...", wait); - tokio::time::sleep(std::time::Duration::from_secs(wait)).await; - continue; - } - }; - - let mut session = session::Session::new(conn, controller.clone()); - session.start_heartbeat().await; - session.wait().await; - connected.store(false, Ordering::Release); - continue; - } - - if support_encryption { - if secure_mode { - connected.store(false, Ordering::Release); - let wait = 1; - log::warn!( - "secure-mode enabled but local build lacks aes-gcm support for web secure tunnel, retrying in {} seconds...", - wait - ); - tokio::time::sleep(std::time::Duration::from_secs(wait)).await; - continue; - } - - log::warn!( - "Server supports encryption but local build lacks aes-gcm support for web secure tunnel, falling back to legacy tunnel" - ); - } - - if secure_mode { - connected.store(false, Ordering::Release); - let wait = 1; - log::warn!( - "secure-mode enabled but server does not support encryption, retrying in {} seconds...", - wait - ); - tokio::time::sleep(std::time::Duration::from_secs(wait)).await; - continue; - } - - session.start_heartbeat().await; - session.wait().await; - connected.store(false, Ordering::Release); - } - } - - pub fn is_connected(&self) -> bool { - self.connected.load(Ordering::Acquire) - } +pub fn parse_config_server_endpoint(input: &str) -> anyhow::Result { + ConfigServerEndpoint::parse(input, |url| TunnelScheme::try_from(url).is_ok()) } pub async fn run_web_client( - config_server_url_s: &str, - machine_id_opts: MachineIdOptions, + config_server_url: &str, + machine_id_options: MachineIdOptions, hostname: Option, secure_mode: bool, - manager: Arc, + manager: Arc, hooks: Option>, ) -> Result { - let machine_id = resolve_machine_id(&machine_id_opts) + let machine_id = resolve_machine_id(&machine_id_options) .with_context(|| "failed to resolve machine id for web client")?; - let config_server_url = match Url::parse(config_server_url_s) { - Ok(u) => u, - Err(_) => format!( - "udp://config-server.easytier.cn:22020/{}", - config_server_url_s - ) - .parse() - .with_context(|| "failed to parse config server URL")?, - }; - - TunnelScheme::try_from(&config_server_url).map_err(|_| { - anyhow::anyhow!( - "unsupported config server scheme: {}", - config_server_url.scheme() - ) - })?; - - let mut c_url = config_server_url.clone(); - if !matches!(c_url.scheme(), "ws" | "wss") { - c_url.set_path(""); - } - let token = config_server_url - .path_segments() - .and_then(|mut x| x.next_back()) - .map(|x| percent_encoding::percent_decode_str(x).decode_utf8()) - .transpose() - .with_context(|| "failed to decode config server token")? - .map(|x| x.to_string()) - .unwrap_or_default(); - - if token.is_empty() { - return Err(anyhow::anyhow!("empty token")); - } + let endpoint = parse_config_server_endpoint(config_server_url)?; let config = TomlConfigLoader::default(); - let global_ctx = Arc::new(GlobalCtx::new(config)); - global_ctx.replace_stun_info_collector(Box::new(MockStunInfoCollector { - udp_nat_type: NatType::Unknown, - })); + let global_ctx = Arc::new(GlobalCtx::new(config.clone())); let mut flags = global_ctx.get_flags(); flags.bind_device = false; global_ctx.set_flags(flags); + let hostname = + hostname.unwrap_or_else(|| gethostname::gethostname().to_string_lossy().to_string()); + let connector = + runtime_one_shot_manual_connector(global_ctx, &config, manager.process_runtime())?; - let hostname = match hostname { - None => gethostname::gethostname().to_string_lossy().to_string(), - Some(hostname) => hostname, - }; Ok(WebClient::new( ConfigServerConnector { - url: c_url, - global_ctx, + url: endpoint.connect_url().clone(), + connector, }, - token.to_string(), + endpoint.token(), machine_id, hostname, secure_mode, @@ -309,11 +137,11 @@ pub async fn run_web_client( mod tests { use std::sync::{Arc, atomic::AtomicBool}; - use crate::{common::MachineIdOptions, instance_manager::NetworkInstanceManager}; + use crate::{common::MachineIdOptions, instance::factory::native_instance_manager}; #[tokio::test] async fn test_manager_wait() { - let manager = Arc::new(NetworkInstanceManager::new()); + let manager = Arc::new(native_instance_manager()); let temp_dir = tempfile::tempdir().unwrap(); let client = super::run_web_client( format!("ring://{}/test", uuid::Uuid::new_v4()).as_str(), @@ -333,21 +161,17 @@ mod tests { tokio::spawn(async move { tokio::time::sleep(std::time::Duration::from_secs(3)).await; - println!("Dropping client..."); sleep_finish_clone.store(true, std::sync::atomic::Ordering::Relaxed); drop(client); - println!("Client dropped."); }); - println!("Waiting for manager..."); manager.wait().await; assert!(sleep_finish.load(std::sync::atomic::Ordering::Relaxed)); - println!("Manager stopped."); } #[tokio::test] async fn test_run_web_client_with_unreachable_config_server() { - let manager = Arc::new(NetworkInstanceManager::new()); + let manager = Arc::new(native_instance_manager()); let temp_dir = tempfile::tempdir().unwrap(); let client = super::run_web_client( "udp://config-server.invalid:22020/test", @@ -365,6 +189,5 @@ mod tests { tokio::time::sleep(std::time::Duration::from_millis(100)).await; assert!(!client.is_connected()); - drop(client); } } diff --git a/easytier/src/web_client/session.rs b/easytier/src/web_client/session.rs deleted file mode 100644 index dbc6505f..00000000 --- a/easytier/src/web_client/session.rs +++ /dev/null @@ -1,170 +0,0 @@ -use std::sync::{Arc, Weak}; - -use tokio::{ - sync::{Mutex, broadcast}, - task::JoinSet, - time::interval, -}; - -use crate::{ - common::constants::EASYTIER_VERSION, - proto::{ - rpc_impl::bidirect::BidirectRpcManager, - rpc_types::controller::BaseController, - web::{ - GetFeatureRequest, GetFeatureResponse, HeartbeatRequest, HeartbeatResponse, - WebServerServiceClientFactory, - }, - }, - tunnel::Tunnel, -}; - -use super::controller::Controller; - -#[derive(Debug, Clone)] -struct HeartbeatCtx { - notifier: Arc>, - resp: Arc>>, -} - -pub struct Session { - rpc_mgr: BidirectRpcManager, - controller: Arc, - - heartbeat_ctx: HeartbeatCtx, - heartbeat_started: std::sync::atomic::AtomicBool, - - tasks: Mutex>, -} - -impl Session { - pub fn new(tunnel: Box, controller: Arc) -> Self { - let rpc_mgr = BidirectRpcManager::new(); - rpc_mgr.run_with_tunnel(tunnel); - - controller.register_api_rpc_service(rpc_mgr.rpc_server().registry()); - - let (tx, _rx1) = broadcast::channel(2); - let heartbeat_ctx = HeartbeatCtx { - notifier: Arc::new(tx), - resp: Arc::new(Mutex::new(None)), - }; - - Session { - rpc_mgr, - controller, - heartbeat_ctx, - heartbeat_started: std::sync::atomic::AtomicBool::new(false), - tasks: Mutex::new(JoinSet::new()), - } - } - - fn heartbeat_routine( - rpc_mgr: &BidirectRpcManager, - controller: Weak, - tasks: &mut JoinSet<()>, - ctx: HeartbeatCtx, - ) { - let controller = controller.upgrade().unwrap(); - let mid = controller.machine_id(); - let inst_id = uuid::Uuid::new_v4(); - let token = controller.token(); - let hostname = controller.hostname(); - let device_os = controller.device_os(); - let controller = Arc::downgrade(&controller); - - let ctx_clone = ctx.clone(); - let mut tick = interval(std::time::Duration::from_secs(1)); - let client = rpc_mgr - .rpc_client() - .scoped_client::>(1, 1, "".to_string()); - tasks.spawn(async move { - loop { - tick.tick().await; - - let Some(controller) = controller.upgrade() else { - break; - }; - - let req = HeartbeatRequest { - machine_id: Some(mid.into()), - inst_id: Some(inst_id.into()), - user_token: token.to_string(), - - easytier_version: EASYTIER_VERSION.to_string(), - hostname: hostname.clone(), - report_time: chrono::Local::now().to_rfc3339(), - device_os: Some(device_os.clone()), - support_config_source: true, - - running_network_instances: controller - .list_network_instance_ids() - .into_iter() - .map(Into::into) - .collect(), - }; - - match client - .heartbeat(BaseController::default(), req.clone()) - .await - { - Err(e) => { - tracing::error!("heartbeat failed: {:?}", e); - break; - } - Ok(resp) => { - tracing::debug!("heartbeat response: {:?}", resp); - let _ = ctx_clone.notifier.send(resp); - ctx_clone.resp.lock().await.replace(resp); - } - } - } - }); - } - - pub async fn start_heartbeat(&self) { - if self - .heartbeat_started - .swap(true, std::sync::atomic::Ordering::AcqRel) - { - return; - } - let mut tasks = self.tasks.lock().await; - Self::heartbeat_routine( - &self.rpc_mgr, - Arc::downgrade(&self.controller), - &mut tasks, - self.heartbeat_ctx.clone(), - ); - } - - async fn wait_routines(&self) { - self.tasks.lock().await.join_next().await; - // if any task failed, we should abort all tasks - self.tasks.lock().await.abort_all(); - } - - pub async fn wait(&mut self) { - tokio::select! { - _ = self.rpc_mgr.wait() => {} - _ = self.wait_routines() => {} - } - } - - pub async fn get_feature( - &self, - ) -> Result { - let client = self - .rpc_mgr - .rpc_client() - .scoped_client::>(1, 1, "".to_string()); - client - .get_feature(BaseController::default(), GetFeatureRequest {}) - .await - } - - pub async fn wait_next_heartbeat(&self) -> Option { - let mut rx = self.heartbeat_ctx.notifier.subscribe(); - rx.recv().await.ok() - } -}