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() - } -}